Evaluation

class medtokenizers.TokenizerEvaluator(model, device=None, data_range=1.0, compute_lpips=False, use_amp=False)[source][source]

Bases: object

Comprehensive evaluator for medical image tokenizers.

This class provides a high-level interface for evaluating tokenizers on test datasets with detailed metrics and reporting. It handles both continuous (VAE) and discrete (VQ-VAE, FSQ) tokenizers.

The Evaluation Pipeline

For each batch: 1. Forward pass through tokenizer (encode -> [quantize] -> decode) 2. Compute reconstruction metrics (PSNR, SSIM, optional LPIPS) 3. For discrete: compute codebook metrics (perplexity, usage) 4. Aggregate statistics across all batches

Thread Safety

The evaluator maintains state (model, device) and should not be used concurrently from multiple threads. Create separate evaluators for parallel evaluation.

type model:

Union[ContinuousTokenizer, DiscreteTokenizer]

param model:

Trained tokenizer model (ContinuousTokenizer or DiscreteTokenizer)

type device:

str | None, default: None

param device:

Device for evaluation (‘cuda’, ‘cpu’, or specific GPU like ‘cuda:0’) Auto-detects if None.

type data_range:

float, default: 1.0

param data_range:

Maximum pixel value in images (default: 1.0 for normalized data)

type compute_lpips:

bool, default: False

param compute_lpips:

Whether to compute LPIPS metric. Requires lpips package and adds ~10% overhead.

type use_amp:

bool, default: False

param use_amp:

Whether to use automatic mixed precision for faster evaluation.

Example

>>> model = DiscreteTokenizer.from_pretrained("./my-vqvae")
>>> evaluator = TokenizerEvaluator(model, device='cuda', compute_lpips=True)
>>>
>>> # Full dataset evaluation
>>> results = evaluator.evaluate(test_loader)
>>> print(f"PSNR: {results['avg_metrics'].psnr:.2f} dB")
>>> print(f"SSIM: {results['avg_metrics'].ssim:.4f}")
>>>
>>> # Quick sanity check on subset
>>> quick_results = evaluator.evaluate(test_loader, num_samples=100)
__init__(model, device=None, data_range=1.0, compute_lpips=False, use_amp=False)[source][source]
Parameters:
evaluate_batch(images, mask=None)[source][source]

Evaluate a single batch of images.

Processes one batch through the tokenizer and computes all metrics. Uses inference mode and optional AMP for efficiency.

Parameters:
  • images (Tensor) – Input images of shape (B, C, H, W) or (B, C, H, W, D)

  • mask (Tensor | None, default: None) – Optional binary mask for masked metrics. Same spatial shape.

Returns:

  • ‘metrics’: EvaluationMetrics object with all computed metrics

  • ’compression_ratio’: Input/latent size ratio

  • ’reconstructions’: Reconstructed images (on CPU)

  • ’indices’: Quantization indices (discrete only, on CPU)

Return type:

Dictionary containing

Example

>>> batch = next(iter(test_loader))
>>> results = evaluator.evaluate_batch(batch)
>>> print(f"Batch PSNR: {results['metrics'].psnr:.2f}")
evaluate(data_loader, num_samples=None, save_reconstructions=False)[source][source]

Evaluate model on a complete dataset.

Iterates through the data loader, computing metrics for each batch and aggregating into summary statistics.

Parameters:
  • data_loader (DataLoader) – DataLoader yielding test batches. Can yield: - Tensor: images only - Tuple: (images, masks) - Dict: {‘image’: images, ‘mask’: masks}

  • num_samples (int | None, default: None) – Maximum samples to evaluate (None = all)

  • save_reconstructions (bool, default: False) – Whether to keep first 10 reconstruction samples

Returns:

  • ‘avg_metrics’: Aggregated EvaluationMetrics

  • ’num_samples’: Total samples evaluated

  • ’model_config’: Model configuration dict

  • ’is_discrete’: Whether discrete tokenizer

  • ’reconstruction_samples’: List of sample dicts (if save_reconstructions)

Return type:

Dictionary containing

Example

>>> # Full evaluation
>>> results = evaluator.evaluate(test_loader)
>>>
>>> # Quick subset evaluation
>>> results = evaluator.evaluate(test_loader, num_samples=100)
>>>
>>> # With reconstruction samples
>>> results = evaluator.evaluate(test_loader, save_reconstructions=True)
>>> for i, sample in enumerate(results['reconstruction_samples']):
...     visualize(sample['original'], sample['reconstruction'])
static save_results(results, save_path, save_samples=False)[source][source]

Save evaluation results to disk.

Saves metrics as JSON for easy parsing. Optionally saves reconstruction samples as compressed numpy archive.

Parameters:
  • results (dict[str, Any]) – Results dictionary from evaluate()

  • save_path (str | Path) – Path for JSON results

  • save_samples (bool, default: False) – Whether to also save reconstruction samples as .npz

Return type:

None

Example

>>> results = evaluator.evaluate(test_loader)
>>> TokenizerEvaluator.save_results(results, "./results.json")
>>> # Creates: ./results.json and optionally ./results.npz
static load_results(load_path)[source][source]

Load evaluation results from disk.

Parameters:

load_path (str | Path) – Path to results JSON file

Return type:

dict[str, Any]

Returns:

Results dictionary with EvaluationMetrics object

Example

>>> results = TokenizerEvaluator.load_results("./results.json")
>>> print(f"Loaded results: PSNR={results['avg_metrics'].psnr:.2f}")
print_results(results)[source][source]

Pretty print evaluation results to console.

Displays a formatted summary of evaluation metrics and model info.

Parameters:

results (dict[str, Any]) – Results dictionary from evaluate()

Return type:

None

Example

>>> results = evaluator.evaluate(test_loader)
>>> evaluator.print_results(results)
# Prints formatted table of metrics

Metrics

medtokenizers.compute_psnr(reconstruction, target, data_range=1.0, mask=None)[source][source]

Compute Peak Signal-to-Noise Ratio (PSNR).

Parameters:
  • reconstruction (Union[Tensor, ndarray]) – Reconstructed images (B, C, H, W) or (B, C, H, W, D)

  • target (Union[Tensor, ndarray]) – Target images (B, C, H, W) or (B, C, H, W, D)

  • data_range (float, default: 1.0) – Maximum possible pixel value (default: 1.0 for normalized images)

  • mask (Union[Tensor, ndarray, None], default: None) – Optional binary mask to compute metric only on masked region

Return type:

float

Returns:

PSNR value in dB

Example

>>> recon = torch.randn(8, 1, 128, 128, 128)
>>> target = torch.randn(8, 1, 128, 128, 128)
>>> psnr = compute_psnr(recon, target, data_range=1.0)
medtokenizers.compute_ssim(reconstruction, target, data_range=1.0, window_size=11, mask=None)[source][source]

Compute Structural Similarity Index (SSIM).

Uses a Gaussian window-based approach to measure structural similarity between images. Works for both 2D and 3D images.

Parameters:
  • reconstruction (Union[Tensor, ndarray]) – Reconstructed images (B, C, H, W) or (B, C, H, W, D)

  • target (Union[Tensor, ndarray]) – Target images (B, C, H, W) or (B, C, H, W, D)

  • data_range (float, default: 1.0) – Maximum possible pixel value (default: 1.0)

  • window_size (int, default: 11) – Size of the Gaussian window (default: 11)

  • mask (Union[Tensor, ndarray, None], default: None) – Optional binary mask to compute metric only on masked region

Return type:

float

Returns:

SSIM value (between -1 and 1, where 1 is perfect similarity)

Example

>>> recon = torch.randn(8, 1, 128, 128, 128)
>>> target = torch.randn(8, 1, 128, 128, 128)
>>> ssim = compute_ssim(recon, target, data_range=1.0)
medtokenizers.compute_lpips(reconstruction, target, net='alex', device=None)[source][source]

Compute Learned Perceptual Image Patch Similarity (LPIPS).

Note: This requires the lpips package to be installed:

pip install lpips

Parameters:
  • reconstruction (Union[Tensor, ndarray]) – Reconstructed images (B, C, H, W)

  • target (Union[Tensor, ndarray]) – Target images (B, C, H, W)

  • net (str, default: 'alex') – Network to use (‘alex’, ‘vgg’, ‘squeeze’)

  • device (Optional[str], default: None) – Device to use for computation

Return type:

float

Returns:

LPIPS value (lower is better, typically 0-1)

Example

>>> recon = torch.randn(8, 1, 128, 128)
>>> target = torch.randn(8, 1, 128, 128)
>>> lpips_val = compute_lpips(recon, target, net='alex')
medtokenizers.compute_mse(reconstruction, target, mask=None)[source][source]

Compute Mean Squared Error.

Parameters:
  • reconstruction (Union[Tensor, ndarray]) – Reconstructed images (B, C, H, W) or (B, C, H, W, D)

  • target (Union[Tensor, ndarray]) – Target images (B, C, H, W) or (B, C, H, W, D)

  • mask (Union[Tensor, ndarray, None], default: None) – Optional binary mask to compute metric only on masked region

Return type:

float

Returns:

MSE value

Example

>>> recon = torch.randn(8, 1, 128, 128, 128)
>>> target = torch.randn(8, 1, 128, 128, 128)
>>> mse = compute_mse(recon, target)
medtokenizers.compute_mae(reconstruction, target, mask=None)[source][source]

Compute Mean Absolute Error.

Parameters:
  • reconstruction (Union[Tensor, ndarray]) – Reconstructed images (B, C, H, W) or (B, C, H, W, D)

  • target (Union[Tensor, ndarray]) – Target images (B, C, H, W) or (B, C, H, W, D)

  • mask (Union[Tensor, ndarray, None], default: None) – Optional binary mask to compute metric only on masked region

Return type:

float

Returns:

MAE value

Example

>>> recon = torch.randn(8, 1, 128, 128, 128)
>>> target = torch.randn(8, 1, 128, 128, 128)
>>> mae = compute_mae(recon, target)
medtokenizers.compute_perplexity(indices, codebook_size)[source][source]

Compute codebook perplexity for discrete tokenizers.

Perplexity measures how well the codebook is being utilized. Higher perplexity indicates better codebook usage.

Parameters:
  • indices (Union[Tensor, ndarray]) – Discrete token indices (B, H, W) or (B, H, W, D)

  • codebook_size (int) – Size of the codebook

Return type:

float

Returns:

Perplexity value

Example

>>> indices = torch.randint(0, 1024, (8, 32, 32, 32))
>>> perplexity = compute_perplexity(indices, codebook_size=1024)
medtokenizers.compute_codebook_usage(indices, codebook_size)[source][source]

Compute codebook usage percentage for discrete tokenizers.

Measures what fraction of the codebook is actually used.

Parameters:
  • indices (Union[Tensor, ndarray]) – Discrete token indices (B, H, W) or (B, H, W, D)

  • codebook_size (int) – Size of the codebook

Return type:

float

Returns:

Usage percentage (0-100)

Example

>>> indices = torch.randint(0, 1024, (8, 32, 32, 32))
>>> usage = compute_codebook_usage(indices, codebook_size=1024)
>>> print(f"Codebook usage: {usage:.1f}%")
medtokenizers.clear_metric_caches()[source][source]

Clear all cached metric computation resources.

Useful for freeing GPU memory after evaluation or when switching devices.

Return type:

None