| import argparse |
| import json |
| from pathlib import Path |
|
|
| import numpy as np |
| import yaml |
| import matplotlib.pyplot as plt |
|
|
|
|
| def rgb(image): |
| array = np.clip(image[[3, 2, 1]], 0, 1).transpose(1, 2, 0) |
| return (array * 255).astype(np.uint8) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="Evaluate and visualize reconstruction") |
| parser.add_argument("--config", default="conf/config.yaml") |
| args = parser.parse_args() |
| with open(args.config, encoding="utf-8") as handle: |
| config = yaml.safe_load(handle) |
| output_dir = Path(config["runtime"]["output_dir"]) |
| reconstruction_path = output_dir / "reconstruction.npz" |
| if not reconstruction_path.exists(): |
| raise FileNotFoundError( |
| f"Missing inference output: {reconstruction_path}. " |
| "Run `python scripts/inference.py` first." |
| ) |
| with np.load(reconstruction_path) as data: |
| inputs = data["inputs"] |
| reconstructions = data["reconstructions"] |
| predictions = data["predictions"] |
| mask_images = data["mask_images"] |
| data_source = str(data["data_source"]) if "data_source" in data.files else "unknown" |
| protocol = str(data["protocol"]) if "protocol" in data.files else "unknown" |
| normalization = str(data["normalization"]) if "normalization" in data.files else "unknown" |
| denominator = max(float(mask_images.sum()), 1.0) |
| mse = float((((inputs - predictions) ** 2) * mask_images).sum() / denominator) |
| mae = float((np.abs(inputs - predictions) * mask_images).sum() / denominator) |
| psnr = float(-10 * np.log10(max(mse, 1e-12))) |
| per_band_denominator = np.maximum(mask_images.sum(axis=(0, 2, 3)), 1) |
| spectral_rmse = np.sqrt((((inputs - predictions) ** 2) * mask_images).sum(axis=(0, 2, 3)) / per_band_denominator) |
| dot = (inputs * reconstructions).sum(axis=1) |
| norms = np.linalg.norm(inputs, axis=1) * np.linalg.norm(reconstructions, axis=1) |
| pixel_mask = (mask_images > 0).any(axis=1) & (norms > 1e-8) |
| sam = np.arccos(np.clip(dot / np.maximum(norms, 1e-8), -1, 1)) |
| metrics = {"masked_mse": mse, "masked_mae": mae, "masked_psnr_db": psnr, |
| "masked_spectral_angle_deg": float(np.degrees(sam[pixel_mask]).mean()), |
| "data_source": data_source, "protocol": protocol, |
| "normalization": normalization, |
| "per_band_rmse": spectral_rmse.tolist()} |
| with open(output_dir / "metrics.json", "w", encoding="utf-8") as handle: |
| json.dump(metrics, handle, indent=2) |
| masked = inputs[0] * (1.0 - mask_images[0]) |
| figure, axes = plt.subplots(1, 4, figsize=(13, 3.5)) |
| for axis, image, title in zip(axes, [inputs[0], masked, predictions[0], reconstructions[0]], |
| ["Input", "Visible tokens", "MAE prediction", "Composite"]): |
| axis.imshow(rgb(image)) |
| axis.set_title(title) |
| axis.axis("off") |
| figure.tight_layout() |
| figure.savefig(output_dir / "reconstruction.png", dpi=140) |
| plt.close(figure) |
| print(json.dumps(metrics, indent=2)) |
| print(f"saved: {output_dir / 'metrics.json'}") |
| print(f"saved: {output_dir / 'reconstruction.png'}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|