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