Download scripts/inference.py from OneScience-Group/TerraMind: direct link, hf CLI and curl.
- Browser
- Download file 1.86 kB
-
https://huggingface.co/OneScience-Group/TerraMind/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/TerraMind/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/TerraMind/resolve/main/scripts/inference.py
1.86 kB
| """Run TerraMind conditional any-to-any token generation.""" | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import yaml | |
| from torch.utils.data import DataLoader | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| from model.terramind import TerraMind | |
| from train import TerraMindDataset, device_from_config, unpack | |
| def main(): | |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) | |
| device = device_from_config(config) | |
| checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location=device, weights_only=True) | |
| model = TerraMind(checkpoint["pixel_modalities"], checkpoint["token_modalities"], checkpoint["model_config"]).to(device) | |
| model.load_state_dict(checkpoint["model"]) | |
| model.eval() | |
| loader = DataLoader(TerraMindDataset(ROOT / config["data"]["root"] / "test.npz", config), batch_size=2) | |
| pixels, tokens = unpack(next(iter(loader)), config, device) | |
| with torch.no_grad(): | |
| conditioning_pixels = {"s2l2a": pixels["s2l2a"]} | |
| generated, embedding = model.generate(conditioning_pixels, tokens, ["lulc", "ndvi", "s1grd"], | |
| input_token_modalities=["coords", "caption"]) | |
| payload = {"embedding": embedding.cpu().numpy(), "pixel_s2l2a": pixels["s2l2a"].cpu().numpy(), | |
| "conditioning_modalities": np.asarray(["pixel_s2l2a", "token_coords", "token_caption"])} | |
| for name, values in generated.items(): | |
| payload[f"generated_{name}"] = values.cpu().numpy() | |
| payload[f"target_{name}"] = tokens[name].cpu().numpy() | |
| output = ROOT / config["paths"]["inference_dir"] / "predictions.npz" | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| np.savez_compressed(output, **payload) | |
| print(f"predictions={output.relative_to(ROOT)}") | |
| if __name__ == "__main__": | |
| main() | |