| """Train FuXi-Ocean on globally indexed tiles, with optional torchrun DDP.""" |
|
|
| import json |
| import os |
| from pathlib import Path |
| import random |
| import sys |
|
|
| import numpy as np |
| import torch |
| from torch.nn.parallel import DistributedDataParallel |
| from torch.utils.data import DataLoader, Dataset |
| from torch.utils.data.distributed import DistributedSampler |
| import yaml |
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(ROOT)) |
| from model.fuxi_ocean import FORMAT_VERSION, FuXiOcean, channel_mask, latitude_weighted_charbonnier |
|
|
|
|
| class TileDataset(Dataset): |
| def __init__(self, data, indices): |
| self.data, self.indices = data, list(indices) |
|
|
| def __len__(self): |
| return len(self.indices) |
|
|
| def __getitem__(self, item): |
| index = self.indices[item] |
| lat, lon = self.data["latitude_deg"][index], self.data["longitude_deg"][index] |
| coordinates = np.stack(np.meshgrid(lon / 180 - 1, lat / 90, indexing="xy")) |
| names = ("ocean", "atmosphere", "bathymetry_m", "depth_mask", "time_features", "targets") |
| values = [self.data[name][index] for name in names] |
| values[2] = values[2] / 5000 |
| return tuple(torch.as_tensor(value, dtype=torch.float32) for value in (*values[:2], coordinates, *values[2:], lat)) |
|
|
|
|
| def validate_data_contract(data, config): |
| expected = config["data"] |
| checks = {"format_version": str(data["format_version"]) == config["data"]["format_version"], |
| "input_shape": data["input_shape"].tolist() == expected["input_shape"], |
| "atmosphere_shape": data["atmosphere_shape"].tolist() == expected["atmosphere_shape"], |
| "output_shape": data["output_shape"].tolist() == expected["output_shape"]} |
| if not all(checks.values()): |
| raise ValueError(f"data contract mismatch: {checks}") |
|
|
|
|
| def main(): |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) |
| torch.set_num_threads(config["runtime"]["num_threads"]) |
| random.seed(config["seed"]); np.random.seed(config["seed"]); torch.manual_seed(config["seed"]) |
| world = int(os.environ.get("WORLD_SIZE", "1")); distributed = world > 1 |
| if distributed: |
| torch.distributed.init_process_group(config["runtime"]["ddp_backend"]) |
| rank = torch.distributed.get_rank() if distributed else 0 |
| local_rank = int(os.environ.get("LOCAL_RANK", "0")) |
| available_devices = torch.cuda.device_count() if torch.cuda.is_available() else 0 |
| requested_auto_gpu = config["runtime"]["device"] != "cpu" and available_devices >= world |
| use_cuda = requested_auto_gpu and local_rank < available_devices |
| device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu") |
| data = np.load(ROOT / config["data"]["path"]) |
| validate_data_contract(data, config) |
| if config["data"]["format_version"] != FORMAT_VERSION: |
| raise ValueError("configuration/model format version mismatch") |
| model_args = {key: value for key, value in config["model"].items() if key != "epsilon"} |
| model = FuXiOcean(**model_args).to(device) |
| if distributed: |
| model = DistributedDataParallel(model, device_ids=[local_rank] if use_cuda else None) |
| optimizer = torch.optim.AdamW(model.parameters(), lr=config["training"]["learning_rate"], |
| weight_decay=config["training"]["weight_decay"]) |
| losses = [] |
| train_count = int(data["train_count"]) |
| dataset = TileDataset(data, range(train_count)) |
| sampler = DistributedSampler(dataset, num_replicas=world, rank=rank, shuffle=True, seed=config["seed"]) if distributed else None |
| loader = DataLoader(dataset, batch_size=config["training"]["batch_size"], sampler=sampler, |
| shuffle=sampler is None, drop_last=False) |
| for epoch in range(config["training"]["epochs"]): |
| if sampler is not None: |
| sampler.set_epoch(epoch) |
| for batch in loader: |
| ocean, atmosphere, coordinates, bathymetry, mask, time_info, target, latitude = [value.to(device) for value in batch] |
| history = ocean |
| step_losses = [] |
| for step in range(config["training"]["multistep_rollout"]): |
| time_info[:, 2] = step |
| prediction = model(history, atmosphere, coordinates, bathymetry, mask, time_info) |
| step_target = target + step * 0.005 |
| step_losses.append(latitude_weighted_charbonnier(prediction, step_target, latitude, |
| channel_mask(mask), config["model"]["epsilon"])) |
| history = torch.cat((history[:, 1:], prediction[:, None]), dim=1) |
| loss = torch.stack(step_losses).mean() |
| optimizer.zero_grad(); loss.backward(); optimizer.step() |
| losses.append(float(loss.detach())) |
| local = torch.tensor([sum(losses), len(losses)], dtype=torch.float64, device=device) |
| if distributed: |
| torch.distributed.all_reduce(local) |
| if rank == 0: |
| raw_model = model.module if distributed else model |
| checkpoint_path = ROOT / config["paths"]["checkpoint"] |
| checkpoint_path.parent.mkdir(parents=True, exist_ok=True) |
| model_config = {"architecture": model_args, "data_format_version": FORMAT_VERSION, |
| "input_shape": data["input_shape"].tolist(), "atmosphere_shape": data["atmosphere_shape"].tolist(), |
| "output_shape": data["output_shape"].tolist()} |
| torch.save({"model": raw_model.state_dict(), "model_config": model_config, "format_version": FORMAT_VERSION, |
| "optimizer": optimizer.state_dict()}, checkpoint_path) |
| metrics_path = ROOT / config["paths"]["training_metrics"] |
| metrics_path.parent.mkdir(parents=True, exist_ok=True) |
| metrics_path.write_text(json.dumps({"mean_loss": local[0].item() / local[1].item(), "world_size": world, |
| "backward_pass": True, "global_loss_all_reduce": distributed, |
| "batch_size": config["training"]["batch_size"], |
| "input_shape": data["input_shape"].tolist(), "synthetic": True}, indent=2) + "\n") |
| print(f"checkpoint={checkpoint_path.relative_to(ROOT)} loss={local[0].item() / local[1].item():.6f}") |
| if distributed: |
| torch.distributed.destroy_process_group() |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|