Evaluation
- class medtokenizers.TokenizerEvaluator(model, device=None, data_range=1.0, compute_lpips=False, use_amp=False)[source][source]
Bases:
objectComprehensive 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:
- 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:
model (
Union[ContinuousTokenizer,DiscreteTokenizer])data_range (
float, default:1.0)compute_lpips (
bool, default:False)use_amp (
bool, default:False)
- 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:
- 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:
- Return type:
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:
- Return type:
- Returns:
Results dictionary with EvaluationMetrics object
Example
>>> results = TokenizerEvaluator.load_results("./results.json") >>> print(f"Loaded results: PSNR={results['avg_metrics'].psnr:.2f}")
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:
- 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:
- 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:
- Return type:
- 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:
- Return type:
- 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:
- Return type:
- 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:
- Return type:
- 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:
- Return type:
- 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}%")