"""Generate float and paper-style signed-int8 annual embedding fields.""" import sys from pathlib import Path import numpy as np import torch import yaml ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) sys.path.insert(0, str(ROOT / "scripts")) from model.alphaearthfoundations import AlphaEarthFoundations, dequantize_embeddings, quantize_embeddings from train import AEFDataset, device_from_config, unpack def main(): config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) torch.manual_seed(config["seed"]) device = device_from_config(config) checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location=device, weights_only=True) if checkpoint["format_version"] != config["data"]["format_version"]: raise ValueError("Checkpoint and data formats do not match") model = AlphaEarthFoundations(checkpoint["input_sources"], checkpoint["target_sources"], checkpoint["model_config"]).to(device) model.load_state_dict(checkpoint["model"]) model.eval() dataset = AEFDataset(ROOT / config["data"]["root"] / "test.npz", config) embeddings, quantized, restored = [], [], [] reconstruction = {name: [] for name in config["data"]["target_sources"]} selected_targets = {name: [] for name in config["data"]["target_sources"]} selected_masks = {name: [] for name in config["data"]["target_sources"]} with torch.no_grad(): for index in range(len(dataset)): batch = {key: value.unsqueeze(0) for key, value in dataset[index].items()} (sources, timestamps, frame_available, targets, masks, target_times, target_periods, geometry) = unpack(batch, config, device) output = model(sources, timestamps, batch["valid_period"].to(device), frame_available, target_times, geometry, target_periods) q = quantize_embeddings(output["embedding"]) embeddings.append(output["embedding"].cpu().numpy()) quantized.append(q.cpu().numpy()) restored.append(dequantize_embeddings(q).cpu().numpy()) for name, values in output["reconstructions"].items(): reconstruction[name].append(values.cpu().numpy()) selected_targets[name].append(targets[name].cpu().numpy()) selected_masks[name].append(masks[name].cpu().numpy()) output_dir = ROOT / config["paths"]["inference_dir"] output_dir.mkdir(parents=True, exist_ok=True) payload = {"embedding": np.concatenate(embeddings), "embedding_s8_power2": np.concatenate(quantized), "embedding_dequantized": np.concatenate(restored)} payload.update({f"reconstruction_{name}": np.concatenate(values) for name, values in reconstruction.items()}) payload.update({f"target_{name}": np.concatenate(values) for name, values in selected_targets.items()}) payload.update({f"mask_{name}": np.concatenate(values) for name, values in selected_masks.items()}) np.savez_compressed(output_dir / "predictions.npz", **payload) print(f"predictions={(output_dir / 'predictions.npz').relative_to(ROOT)}") if __name__ == "__main__": main()