zhangrenchao's picture
Add engineering reproduction package
3549cf5 verified
Raw
History Blame Contribute Delete
3.18 kB
"""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()