| """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() |
|
|