Download scripts/result.py from OneScience-Group/SatMAE: direct link, hf CLI and curl.
- Browser
- Download file 6.25 kB
-
https://huggingface.co/OneScience-Group/SatMAE/resolve/main/scripts/result.py
- Command line
-
hf download hf://OneScience-Group/SatMAE/scripts/result.py
-
curl -L -o result.py https://huggingface.co/OneScience-Group/SatMAE/resolve/main/scripts/result.py
6.25 kB
| """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() | |