"""Evaluate SatMAE masked reconstruction across time and channels.""" import argparse import json from pathlib import Path import matplotlib.pyplot as plt import numpy as np, yaml ROOT = Path(__file__).resolve().parents[1] def parse_args(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml") parser.add_argument("--input", type=Path, default=None) parser.add_argument("--output-dir", type=Path, default=None) return parser.parse_args() def unpatchify(patches, image_size, patch_size, channels): side = image_size // patch_size image = patches.reshape(side, side, channels, patch_size, patch_size) return image.transpose(2, 0, 3, 1, 4).reshape(channels, image_size, image_size) def display_image(image): image = image[:3].transpose(1, 2, 0) low, high = float(image.min()), float(image.max()) return np.clip((image - low) / max(high - low, 1e-8), 0.0, 1.0) def main(): args = parse_args() cfg = yaml.safe_load(args.config.read_text()) source = args.input or ROOT / cfg["paths"]["inference_dir"] / "reconstruction.npz" if not source.exists(): raise FileNotFoundError("Run inference before evaluation") a = np.load(source) masked = a["mask"].astype(bool) out = args.output_dir or ROOT / cfg["paths"]["evaluation_dir"]; out.mkdir(parents=True, exist_ok=True) if cfg["model"]["mode"] == "multispectral": groups = cfg["model"]["spectral_groups"] group_mask = masked.reshape(masked.shape[0], len(groups), -1) group_mse, masked_group_mse = [], [] weighted_error = 0.0 weighted_count = 0 masked_error_sum = 0.0 masked_count = 0 for index, group in enumerate(groups): target = a[f"target_group_{index}"] prediction = a[f"prediction_group_{index}"] squared = (prediction - target) ** 2 patch_error = squared.mean(axis=-1) group_mse.append(float(squared.mean())) selected = group_mask[:, index] masked_group_mse.append(float(patch_error[selected].mean())) weighted_error += float(squared.sum()) weighted_count += squared.size masked_error_sum += float(patch_error[selected].sum()) masked_count += int(selected.sum()) result = { "masked_mse": masked_error_sum / max(masked_count, 1), "reconstruction_mse": weighted_error / max(weighted_count, 1), "group_mse": group_mse, "masked_group_mse": masked_group_mse, "data_source": "synthetic", "protocol": cfg["data"]["protocol"], } (out / "metrics.json").write_text(json.dumps(result, indent=2) + "\n") print(json.dumps(result, indent=2)); print("evaluation=", out) return squared_error = (a["prediction"] - a["target"]) ** 2 patch_error = squared_error.mean(axis=-1) error = float(squared_error.mean()) masked_error = float(patch_error[masked].mean()) if masked.any() else error result = {"masked_mse": masked_error, "reconstruction_mse": error, "data_source": "synthetic", "protocol": cfg["data"]["protocol"]} size = cfg["model"]["image_size"]; patch = cfg["model"]["patch_size"]; channels = cfg["model"]["in_channels"] patch_count = (size // patch) ** 2 target = a["target"][0, :patch_count] prediction = a["prediction"][0, :patch_count] patch_mask = a["mask"][0, :patch_count] masked_target = target.copy(); masked_target[patch_mask] = 0.0 panels = [ ("Original", unpatchify(target, size, patch, channels)), ("Masked input", unpatchify(masked_target, size, patch, channels)), ("Reconstruction", unpatchify(prediction, size, patch, channels)), ] figure, axes = plt.subplots(1, 3, figsize=(10, 3.4)) for axis, (title, image) in zip(axes, panels): axis.imshow(display_image(image)); axis.set_title(title); axis.axis("off") figure.tight_layout(); figure.savefig(out / "temporal_frame_reconstruction.png", dpi=160, bbox_inches="tight"); plt.close(figure) frames = cfg["model"]["frames"] if cfg["model"]["mode"] == "temporal" else 1 frame_mse, masked_frame_mse = [], [] channel_mse = np.zeros(channels, dtype=np.float64) for frame in range(frames): start, end = frame * patch_count, (frame + 1) * patch_count frame_target = a["target"][:, start:end] frame_prediction = a["prediction"][:, start:end] mse = float(np.mean((frame_prediction - frame_target) ** 2)) frame_mse.append(mse) frame_mask = masked[:, start:end] frame_patch_error = patch_error[:, start:end] masked_frame_mse.append(float(frame_patch_error[frame_mask].mean()) if frame_mask.any() else mse) shaped_error = ((frame_prediction - frame_target) ** 2).reshape(-1, channels, patch * patch).mean(axis=(0, 2)) channel_mse += shaped_error channel_mse /= frames figure, axis = plt.subplots(figsize=(6.2, 3.8)) frame_index = np.arange(1, frames + 1) axis.plot(frame_index, frame_mse, marker="o", linewidth=2, label="All patches") axis.plot(frame_index, masked_frame_mse, marker="s", linewidth=2, label="Masked patches") axis.set(xlabel="Time frame", ylabel="MSE", title="Temporal Reconstruction Error") axis.legend(); axis.grid(alpha=0.25); figure.tight_layout(); figure.savefig(out / "temporal_reconstruction_error.png", dpi=160); plt.close(figure) figure, axis = plt.subplots(figsize=(6.2, 3.8)) axis.bar(np.arange(channels), channel_mse, color="#287271") axis.set_xticks(np.arange(channels), [f"C{i + 1}" for i in range(channels)]) axis.set(xlabel="Input channel", ylabel="MSE", title="Channel Reconstruction Error") figure.tight_layout(); figure.savefig(out / "spectral_band_reconstruction.png", dpi=160); plt.close(figure) result["frame_mse"] = frame_mse result["masked_frame_mse"] = masked_frame_mse result["channel_mse"] = channel_mse.tolist() (out / "metrics.json").write_text(json.dumps(result, indent=2) + "\n") print(json.dumps(result, indent=2)); print("evaluation=", out) if __name__ == "__main__": main()