TerraMind / scripts /result.py
zhangrenchao's picture
Add engineering reproduction package
3571a70 verified
Raw
History Blame Contribute Delete
1.78 kB
"""Evaluate conditional token generation and visualize patch predictions."""
import json
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import yaml
ROOT = Path(__file__).resolve().parents[1]
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
data = np.load(ROOT / config["paths"]["inference_dir"] / "predictions.npz")
metrics = {"samples": int(len(data["embedding"])), "mean_embedding_norm": float(np.linalg.norm(data["embedding"], axis=1).mean()),
"conditioning_modalities": data["conditioning_modalities"].tolist(), "token_accuracy": {},
"random_token_accuracy": 1.0 / int(config["model"]["engineering_vocab_size"])}
for name in ("lulc", "ndvi", "s1grd"):
metrics["token_accuracy"][name] = float((data[f"generated_{name}"] == data[f"target_{name}"]).mean())
output = ROOT / config["paths"]["evaluation_dir"]
output.mkdir(parents=True, exist_ok=True)
(output / "metrics.json").write_text(json.dumps(metrics, indent=2) + "\n")
figure, axes = plt.subplots(2, 3, figsize=(9, 6))
grid = int(np.sqrt(data["generated_lulc"].shape[1]))
for column, name in enumerate(("lulc", "ndvi", "s1grd")):
axes[0, column].imshow(data[f"target_{name}"][0].reshape(grid, grid), cmap="viridis")
axes[1, column].imshow(data[f"generated_{name}"][0].reshape(grid, grid), cmap="viridis")
axes[0, column].set_title(f"{name} target")
axes[1, column].set_title(f"{name} generated")
axes[0, column].axis("off")
axes[1, column].axis("off")
figure.tight_layout()
figure.savefig(output / "comparison.png", dpi=150)
plt.close(figure)
if __name__ == "__main__":
main()