Skip to main content

An sklearn-style FID metric class for Jax

Project description

FIDax

A JAX implementation of the Fréchet Inception Distance (FID) metric for evaluating generative models in form of a sklearn-compatible metric.

Features

  • Pure JAX Implementation: Leverages JAX's JIT compilation for fast computation
  • Pre-computed Statistics: Supports using pre-computed real image statistics for faster evaluation
  • GPU Accelerated: Optimized for CUDA-enabled GPUs
  • Torchmetrics Compatible: Results match torchmetrics implementation up to 1e-1 absolute tolerance with FP32 execution of the InceptionV3 model and FP64 for the metric computation on CIFAR10 tests

Installation

# Clone the repository
git clone git@github.com:wittenator/fidax.git
cd fidax

# Install dependencies using uv
uv sync --frozen

or install it directly as a dependency with e.g. uv:

uv add fidax

Quick Start

import jax 
jax.config.update("jax_enable_x64", True)
import jax.numpy as jnp
from fidax.fid import FrechetInceptionDistance

# Initialize FID metric
fid = FrechetInceptionDistance(max_samples=100)

# Update with real images (shape: [N, 299, 299, 3], range: [-1, 1])
real_images = jnp.random.uniform(-1, 1, (100, 299, 299, 3))
fid.update(real_images, real=True)

# Update with generated/fake images
fake_images = jnp.random.uniform(-1, 1, (100, 299, 299, 3))
fid.update(fake_images, real=False)

# Compute FID score
fid_score = fid.compute()
print(f"FID Score: {fid_score}")

Advanced Usage

Pre-computed Statistics

# Use pre-computed real statistics for faster evaluation
real_stats = {
    "mu": mu_real,      # Mean of real activations
    "sigma": sigma_real # Covariance of real activations
}

fid = FrechetInceptionDistance(max_samples=1000, real_stats=real_stats)
# Only need to update with fake images
fid.update(fake_images, real=False)

Requirements

  • Python ≥ 3.12
  • JAX with CUDA support
  • Flax
  • NumPy

See pyproject.toml for complete dependency list.

Development

This project uses a development container with GPU support. To set up the development environment:

# The dev container will automatically install dependencies
# Run tests
uv run pytest src/fidax/test_fid_metric.py

Testing

The implementation includes tests against torchmetrics:

uv run pytest src/fidax/test_fid_metric.py -v

Tests verify:

  • Equivalence with torchmetrics implementation
  • Pre-computed statistics functionality
  • Real-world performance on CIFAR-10 dataset

License

Apache 2.0 License - see LICENSE for details.

Related Projects

  • jax-fid-parallel - Parallel implementation of FID computation in JAX
  • jax-fid - Original JAX implementation of FID that inspired this project

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

fidax-0.2.tar.gz (117.8 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

fidax-0.2-py3-none-any.whl (18.5 kB view details)

Uploaded Python 3

File details

Details for the file fidax-0.2.tar.gz.

File metadata

  • Download URL: fidax-0.2.tar.gz
  • Upload date:
  • Size: 117.8 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: uv/0.9.5

File hashes

Hashes for fidax-0.2.tar.gz
Algorithm Hash digest
SHA256 301bcdff9b5e4105d229982a9679d7767e1dc9bc817f6c9cc4f43f4d509adc2e
MD5 0dcedf04ad8fe2613cffa11e59f6eb26
BLAKE2b-256 158c39b8ccf6b1d18a48d5339bc880d5edd88ff9e7cf8ccf0242884de1ee6b43

See more details on using hashes here.

File details

Details for the file fidax-0.2-py3-none-any.whl.

File metadata

  • Download URL: fidax-0.2-py3-none-any.whl
  • Upload date:
  • Size: 18.5 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: uv/0.9.5

File hashes

Hashes for fidax-0.2-py3-none-any.whl
Algorithm Hash digest
SHA256 1de247d9a06db697b74684e94b7ea6c0751a1639c3e2f73be4af0e059ccc7dea
MD5 19b37537078b0867f82687f5dff71271
BLAKE2b-256 e38d07a31616c14ea040ba5440b2b9e2a001f008c368091c2738a41e44a98d9a

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page