| """Evaluate DOFA on masked pixels with explicit PSNR data ranges.""" |
|
|
| import json |
| from pathlib import Path |
|
|
| import matplotlib |
| matplotlib.use("Agg") |
| import matplotlib.pyplot as plt |
| import numpy as np |
| import yaml |
|
|
| ROOT = Path(__file__).resolve().parents[1] |
|
|
|
|
| def scalar(archive, key): |
| if key not in archive or archive[key].ndim != 0: |
| raise ValueError(f"Inference metadata {key} must be present as a scalar") |
| return archive[key].item() |
|
|
|
|
| def display(array): |
| selected = array[:3] if len(array) >= 3 else np.repeat(array[:1], 3, 0) |
| selected = selected.transpose(1, 2, 0) |
| return np.clip((selected - selected.min()) / max(np.ptp(selected), 1e-6), 0, 1) |
|
|
|
|
| def pixel_mask(patch_mask, image_size, patch_size): |
| side = image_size // patch_size |
| return np.repeat(np.repeat(patch_mask.reshape(side, side), patch_size, 0), patch_size, 1) |
|
|
|
|
| def main(): |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text(encoding="utf-8")) |
| input_dir = ROOT / config["paths"]["inference_dir"] |
| output_dir = ROOT / config["paths"]["evaluation_dir"] |
| output_dir.mkdir(parents=True, exist_ok=True) |
| metrics, sources, protocols = {}, set(), set() |
| for modality, modality_config in config["data"]["modalities"].items(): |
| path = input_dir / f"{modality}_reconstruction.npz" |
| if not path.is_file(): |
| raise FileNotFoundError("Run inference before result evaluation") |
| archive = np.load(path) |
| inputs, reconstructions, masks = archive["inputs"], archive["reconstructions"], archive["masks"] |
| protocol = str(scalar(archive, "protocol")) |
| source = str(scalar(archive, "data_source")) |
| output_modality = str(scalar(archive, "modality")) |
| data_range = float(scalar(archive, "data_range")) |
| expected_shape = (modality_config["channels"], config["data"]["image_size"], config["data"]["image_size"]) |
| if inputs.ndim != 4 or inputs.shape[1:] != expected_shape or reconstructions.shape != inputs.shape: |
| raise ValueError(f"{path}: invalid reconstruction NCHW shape") |
| if protocol != config["data"]["protocol"] or output_modality != modality: |
| raise ValueError(f"{path}: protocol/modality metadata does not match config") |
| if not np.isfinite(data_range) or data_range <= 0: |
| raise ValueError(f"{path}: PSNR requires a positive data_range") |
| sources.add(source) |
| protocols.add(protocol) |
| expanded = np.stack([pixel_mask(mask, config["data"]["image_size"], |
| config["model"]["patch_size"]) for mask in masks])[:, None] |
| errors = (inputs - reconstructions)[np.broadcast_to(expanded, inputs.shape)] |
| mse, mae = float(np.mean(errors**2)), float(np.mean(np.abs(errors))) |
| metrics[modality] = {"masked_mse": mse, "masked_mae": mae, |
| "masked_psnr_db": float(10 * np.log10(data_range**2 / max(mse, 1e-12))), |
| "psnr_data_range": data_range, "channels": int(inputs.shape[1]), |
| "masked_fraction": float(expanded.mean()), "data_source": source, |
| "protocol": protocol} |
| masked_input = inputs[0].copy() |
| masked_input[:, expanded[0, 0].astype(bool)] = 0 |
| figure, axes = plt.subplots(1, 3, figsize=(10, 3)) |
| for axis, image, title in zip(axes, (inputs[0], masked_input, reconstructions[0]), |
| ("target", "masked input", "reconstruction")): |
| axis.imshow(display(image)); axis.set_title(f"{modality}: {title}"); axis.axis("off") |
| figure.tight_layout() |
| figure.savefig(output_dir / f"{modality}_comparison.png", dpi=120) |
| plt.close(figure) |
| payload = {"modalities": metrics, "data_sources": sorted(sources), |
| "protocols": sorted(protocols), "protocol": config["data"]["protocol"], |
| "metric_scope": "masked_pixels_only"} |
| (output_dir / "metrics.json").write_text(json.dumps(payload, indent=2) + "\n") |
| print(json.dumps(payload, indent=2)); print(f"evaluation={output_dir}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|