DOFA / scripts /result.py
zhangrenchao's picture
Upload DOFA model package
1d4cac8 verified
Raw
History Blame Contribute Delete
4.15 kB
"""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()