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