Download scripts/train.py from OneScience-Group/FuXi-Ocean: direct link, hf CLI and curl.
- Browser
- Download file 6.42 kB
-
https://huggingface.co/OneScience-Group/FuXi-Ocean/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/FuXi-Ocean/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/FuXi-Ocean/resolve/main/scripts/train.py
6.42 kB
| """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() | |