ClayFoundation / scripts /inference.py
zhangrenchao's picture
Add engineering reproduction package
53becf5 verified
Raw
History Blame Contribute Delete
2.44 kB
"""Generate Clay embeddings and MAE reconstructions for every configured sensor."""
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))
from model.clayfoundation import ClayFoundation
from train import ClayDataset, device_from_config
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)
if checkpoint["format_version"] != config["data"]["format_version"]:
raise ValueError("checkpoint and data formats are incompatible")
model = ClayFoundation(checkpoint["model_config"]).to(device)
model.load_state_dict(checkpoint["model"])
model.eval()
dataset = ClayDataset(ROOT / config["data"]["root"] / "test.npz", config)
raw = dataset.data
payload = {
"format_version": np.asarray(config["data"]["format_version"]),
"class_target": raw["class_target"],
"regression_target": raw["regression_target"],
}
time = torch.from_numpy(raw["time"]).to(device)
latlon = torch.from_numpy(raw["latlon"]).to(device)
teacher = torch.from_numpy(raw["teacher_target"]).to(device)
with torch.no_grad():
for name, spec in config["data"]["sensors"].items():
pixels = torch.from_numpy(raw[f"pixels_{name}"]).to(device)
waves = torch.from_numpy(raw[f"wavelengths_{name}"]).to(device)
outputs = model(pixels, time, latlon, float(spec["gsd"]), waves, teacher, mask_ratio=0.0)
payload[f"embedding_{name}"] = outputs["embedding"].cpu().numpy()
payload[f"projected_embedding_{name}"] = outputs["projected_embedding"].cpu().numpy()
payload[f"reconstruction_{name}"] = outputs["reconstruction"].cpu().numpy()
payload[f"pixels_{name}"] = raw[f"pixels_{name}"]
payload[f"reconstruction_loss_{name}"] = np.asarray(float(outputs["reconstruction_loss"]))
payload[f"representation_loss_{name}"] = np.asarray(float(outputs["representation_loss"]))
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()