SatMAE / scripts /result.py
zhangrenchao's picture
Upload SatMAE model package
355f250 verified
Raw
History Blame Contribute Delete
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()