Calculate FID score based on accumulated extracted features from the two distributions.
(self)
| 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.""" |
no test coverage detected