| """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() |
|
|