File size: 1,857 Bytes
3571a70 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 | """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()
|