MCPcopy Create free account
hub / github.com/CompVis/zigma / compute

Method compute

utils/torchmetric_fvd.py:391–418  ·  view source on GitHub ↗

Calculate FID score based on accumulated extracted features from the two distributions.

(self)

Source from the content-addressed store, hash-verified

389 self.fake_features_num_samples += videos.shape[0]
390
391 def compute(self) -> Tensor:
392 """Calculate FID score based on accumulated extracted features from the two distributions."""
393 if self.real_features_num_samples < 2 or self.fake_features_num_samples < 2:
394 raise RuntimeError(
395 "More than one sample is required for both the real and fake distributed to compute FID"
396 )
397 mean_real = (self.real_features_sum / self.real_features_num_samples).unsqueeze(
398 0
399 )
400 mean_fake = (self.fake_features_sum / self.fake_features_num_samples).unsqueeze(
401 0
402 )
403
404 cov_real_num = (
405 self.real_features_cov_sum
406 - self.real_features_num_samples * mean_real.t().mm(mean_real)
407 )
408 cov_real = cov_real_num / (self.real_features_num_samples - 1)
409 cov_fake_num = (
410 self.fake_features_cov_sum
411 - self.fake_features_num_samples * mean_fake.t().mm(mean_fake)
412 )
413 cov_fake = cov_fake_num / (self.fake_features_num_samples - 1)
414 return (
415 _compute_fid(mean_real.squeeze(0), cov_real, mean_fake.squeeze(0), cov_fake)
416 .to(self.orig_dtype)
417 .item()
418 )
419
420 def reset(self) -> None:
421 """Reset metric states."""

Callers 1

torchmetric_fvd.pyFile · 0.45

Calls 2

toMethod · 0.80
_compute_fidFunction · 0.70

Tested by

no test coverage detected