"""Train AEF reconstruction, uniformity, consistency and text objectives.""" 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.alphaearthfoundations import AlphaEarthFoundations, compute_losses class AEFDataset(Dataset): def __init__(self, path, config): self.data = np.load(path) self.input_sources = config["data"]["input_sources"] self.target_sources = config["data"]["target_sources"] format_version = str(self.data["format_version"]) if format_version != config["data"]["format_version"]: raise ValueError(f"Expected format {config['data']['format_version']}, got {format_version}") for name, spec in self.input_sources.items(): expected = (spec["timesteps"], spec["channels"], config["data"]["image_size"], config["data"]["image_size"]) if self.data[name].shape[1:] != expected: raise ValueError(f"{name} has shape {self.data[name].shape[1:]}, expected {expected}") def __len__(self): return len(self.data["valid_period"]) def __getitem__(self, index): item = {name: torch.from_numpy(self.data[name][index]).float() for name in self.input_sources} item.update({f"timestamps_{name}": torch.from_numpy(self.data[f"timestamps_{name}"][index]).long() for name in self.input_sources}) for name in self.input_sources: item[f"frame_available_{name}"] = torch.from_numpy(self.data[f"frame_available_{name}"][index]) item[f"channel_available_{name}"] = torch.from_numpy(self.data[f"channel_available_{name}"][index]) item[f"pixel_valid_{name}"] = torch.from_numpy(self.data[f"pixel_valid_{name}"][index]).float() item[f"geometry_{name}"] = torch.from_numpy(self.data[f"geometry_{name}"][index]).float() for name, spec in self.target_sources.items(): if name in self.input_sources: continue values = torch.from_numpy(self.data[f"target_sequence_{name}"][index]) item[f"target_sequence_{name}"] = values.long() if spec["type"] == "categorical" else values.float() item[f"target_timestamps_{name}"] = torch.from_numpy(self.data[f"target_timestamps_{name}"][index]).long() item[f"target_pixel_valid_{name}"] = torch.from_numpy(self.data[f"target_pixel_valid_{name}"][index]).float() item[f"target_geometry_{name}"] = torch.from_numpy(self.data[f"target_geometry_{name}"][index]).float() item[f"target_frame_available_{name}"] = torch.from_numpy(self.data[f"target_frame_available_{name}"][index]) item["support_period"] = torch.from_numpy(self.data["support_period"][index]).long() item["valid_period"] = torch.from_numpy(self.data["valid_period"][index]).long() item["text_target"] = torch.from_numpy(self.data["text_target"][index]).float() 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 _select_indices(available): selected = [] for row in available: candidates = torch.nonzero(row, as_tuple=False).flatten() choice = torch.randint(len(candidates), (), device=candidates.device) selected.append(candidates[choice]) return torch.stack(selected).to(available.device) def _gather(values, indices): return values[torch.arange(len(values), device=values.device), indices] def unpack(batch, config, device, remove_input_targets=True): input_names, target_specs = config["data"]["input_sources"], config["data"]["target_sources"] sources, frame_available = {}, {} timestamps = {name: batch[f"timestamps_{name}"].to(device) for name in input_names} targets, masks, target_times, target_periods, geometry = {}, {}, {}, {}, {} for name in input_names: values = batch[name].to(device) available = batch[f"frame_available_{name}"].to(device).clone() channel_available = batch[f"channel_available_{name}"].to(device) values = values * channel_available[:, :, :, None, None] indices = _select_indices(available) targets[name] = _gather(values, indices) masks[name] = _gather(batch[f"pixel_valid_{name}"].to(device), indices) * _gather(channel_available, indices)[:, :, None, None] target_times[name] = _gather(timestamps[name], indices) target_periods[name] = torch.stack([ target_times[name] - 5 * 86_400_000, target_times[name] + 5 * 86_400_000 ], dim=1) geometry[name] = _gather(batch[f"geometry_{name}"].to(device), indices) if remove_input_targets: available[torch.arange(len(available), device=device), indices] = False sources[name] = values * available[:, :, None, None, None] frame_available[name] = available for name in target_specs: if name in input_names: continue available = batch[f"target_frame_available_{name}"].to(device) indices = _select_indices(available) targets[name] = _gather(batch[f"target_sequence_{name}"].to(device), indices) masks[name] = _gather(batch[f"target_pixel_valid_{name}"].to(device), indices) target_times[name] = _gather(batch[f"target_timestamps_{name}"].to(device), indices) target_periods[name] = torch.stack([ target_times[name] - 5 * 86_400_000, target_times[name] + 5 * 86_400_000 ], dim=1) geometry[name] = _gather(batch[f"target_geometry_{name}"].to(device), indices) return sources, timestamps, frame_available, targets, masks, target_times, target_periods, geometry def perturb_sources(sources, frame_available, timestamps, support_period, config): perturbed, perturbed_available = {}, {} source_dropout, frame_dropout = config["train"]["source_dropout"], config["train"]["frame_dropout"] for name, values in sources.items(): output, available = values + 0.01 * torch.randn_like(values), frame_available[name].clone() source_mask = torch.rand(values.shape[0], 1, 1, 1, 1, device=values.device) < source_dropout[name] strategy = torch.randint(3, (1,), device=values.device).item() if strategy == 0: dropped = torch.rand_like(available.float()) < frame_dropout[name] else: midpoint = support_period[:, :1] + (support_period[:, 1:] - support_period[:, :1]) // 2 dropped = timestamps[name] >= midpoint if strategy == 1 else timestamps[name] < midpoint available = available & ~dropped & ~source_mask[:, 0, 0, 0] perturbed[name] = output * available[:, :, None, None, None] perturbed_available[name] = available return perturbed, perturbed_available def main(): config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) torch.manual_seed(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 = AEFDataset(ROOT / config["data"]["root"] / "train.npz", config) sampler = DistributedSampler(dataset, shuffle=True) if distributed else None loader = DataLoader(dataset, batch_size=config["train"]["batch_size"], sampler=sampler, shuffle=sampler is None, num_workers=config["train"]["num_workers"]) model = AlphaEarthFoundations(config["data"]["input_sources"], config["data"]["target_sources"], config["model"]).to(device) if distributed: model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) optimizer = torch.optim.Adam(model.parameters(), lr=config["train"]["learning_rate"], weight_decay=config["train"]["weight_decay"]) weights = {name: config["train"][f"{name}_weight"] for name in ("reconstruction", "uniformity", "consistency", "text")} history = [] for epoch in range(config["train"]["epochs"]): if sampler: sampler.set_epoch(epoch) model.train() totals = {} for batch in loader: sources, timestamps, frame_available, targets, masks, target_times, target_periods, geometry = unpack(batch, config, device) valid_period = batch["valid_period"].to(device) teacher = model(sources, timestamps, valid_period, frame_available, target_times, geometry, target_periods) perturbed, perturbed_available = perturb_sources( sources, frame_available, timestamps, batch["support_period"].to(device), config ) student = model(perturbed, timestamps, valid_period, perturbed_available) loss, components = compute_losses(teacher, student, targets, masks, batch["text_target"].to(device), config["data"]["target_sources"], weights) optimizer.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() for name, value in components.items(): totals[name] = totals.get(name, 0.0) + float(value.detach()) metrics = {name: value / len(loader) for name, 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"], "input_sources": config["data"]["input_sources"], "target_sources": config["data"]["target_sources"], "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()