Download scripts/inference.py from OneScience-Group/AlphaEarthFoundations: direct link, hf CLI and curl.
- Browser
- Download file 3.18 kB
-
https://huggingface.co/OneScience-Group/AlphaEarthFoundations/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/AlphaEarthFoundations/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/AlphaEarthFoundations/resolve/main/scripts/inference.py
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() | |