| """Train the reduced Clay MAE on deterministic multi-sensor synthetic chips.""" |
|
|
| import json |
| import os |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import yaml |
| from torch.nn.parallel import DistributedDataParallel |
| from torch.utils.data import DataLoader, Dataset, DistributedSampler |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(ROOT)) |
| from model.clayfoundation import ClayFoundation, compute_loss |
|
|
|
|
| class ClayDataset(Dataset): |
| def __init__(self, path, config): |
| self.data = np.load(path) |
| self.sensors = config["data"]["sensors"] |
| if str(self.data["format_version"]) != config["data"]["format_version"]: |
| raise ValueError("incompatible synthetic data format") |
| size = int(config["data"]["image_size"]) |
| for name, spec in self.sensors.items(): |
| expected = (int(spec["channels"]), size, size) |
| if self.data[f"pixels_{name}"].shape[1:] != expected: |
| raise ValueError(f"{name} shape does not match {expected}") |
|
|
| def __len__(self): |
| return len(self.data["time"]) |
|
|
| def __getitem__(self, index): |
| item = { |
| "time": torch.from_numpy(self.data["time"][index]), |
| "latlon": torch.from_numpy(self.data["latlon"][index]), |
| "teacher_target": torch.from_numpy(self.data["teacher_target"][index]), |
| } |
| for name in self.sensors: |
| item[f"pixels_{name}"] = torch.from_numpy(self.data[f"pixels_{name}"][index]) |
| item[f"valid_{name}"] = torch.from_numpy(self.data[f"valid_{name}"][index]) |
| item[f"waves_{name}"] = torch.from_numpy(self.data[f"wavelengths_{name}"][index]) |
| return item |
|
|
|
|
| def device_from_config(config, local_rank=0): |
| requested = config["runtime"]["device"] |
| if requested == "auto": |
| return torch.device("cuda", local_rank) if torch.cuda.is_available() else torch.device("cpu") |
| return torch.device(requested) |
|
|
|
|
| def main(): |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) |
| torch.manual_seed(int(config["seed"])) |
| distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1 |
| local_rank = int(os.environ.get("LOCAL_RANK", "0")) |
| if distributed: |
| torch.distributed.init_process_group("nccl" if torch.cuda.is_available() else "gloo") |
| rank = torch.distributed.get_rank() if distributed else 0 |
| device = device_from_config(config, local_rank) |
| if device.type == "cuda": |
| torch.cuda.set_device(device) |
| dataset = ClayDataset(ROOT / config["data"]["root"] / "train.npz", config) |
| sampler = DistributedSampler(dataset, shuffle=True) if distributed else None |
| loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), sampler=sampler, |
| shuffle=sampler is None, num_workers=int(config["train"]["num_workers"])) |
| model = ClayFoundation(config["model"]).to(device) |
| if distributed: |
| model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) |
| optimizer = torch.optim.AdamW(model.parameters(), lr=float(config["train"]["learning_rate"]), |
| weight_decay=float(config["train"]["weight_decay"]), betas=(0.9, 0.95)) |
| history = [] |
| for epoch in range(int(config["train"]["epochs"])): |
| if sampler: |
| sampler.set_epoch(epoch) |
| model.train() |
| totals = {"reconstruction": 0.0, "representation": 0.0, "total": 0.0} |
| steps = 0 |
| for batch in loader: |
| optimizer.zero_grad(set_to_none=True) |
| sensor_losses = [] |
| components = [] |
| for name, spec in config["data"]["sensors"].items(): |
| pixels = batch[f"pixels_{name}"].to(device) * batch[f"valid_{name}"].to(device) |
| outputs = model(pixels, batch["time"].to(device), batch["latlon"].to(device), |
| float(spec["gsd"]), batch[f"waves_{name}"].to(device), |
| batch["teacher_target"].to(device)) |
| loss, values = compute_loss(outputs, float(config["train"]["reconstruction_weight"]), |
| float(config["train"]["representation_weight"])) |
| sensor_losses.append(loss) |
| components.append(values) |
| loss = torch.stack(sensor_losses).mean() |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) |
| optimizer.step() |
| for key in totals: |
| totals[key] += sum(float(item[key].detach()) for item in components) / len(components) |
| steps += 1 |
| metrics = {key: value / max(steps, 1) for key, value in totals.items()} |
| history.append({"epoch": epoch + 1, **metrics}) |
| if rank == 0: |
| print(f"epoch={epoch + 1} total_loss={metrics['total']:.6f} reconstruction={metrics['reconstruction']:.6f}") |
| if rank == 0: |
| checkpoint = ROOT / config["paths"]["checkpoint"] |
| metrics_path = ROOT / config["paths"]["training_metrics"] |
| checkpoint.parent.mkdir(parents=True, exist_ok=True) |
| metrics_path.parent.mkdir(parents=True, exist_ok=True) |
| state = model.module.state_dict() if distributed else model.state_dict() |
| torch.save({"model": state, "model_config": config["model"], "sensors": config["data"]["sensors"], |
| "format_version": config["data"]["format_version"]}, checkpoint) |
| metrics_path.write_text(json.dumps({"history": history}, indent=2) + "\n") |
| print(f"checkpoint={checkpoint.relative_to(ROOT)}") |
| if distributed: |
| torch.distributed.destroy_process_group() |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|