Download scripts/train.py from OneScience-Group/ConvLSTM: direct link, hf CLI and curl.
- Browser
- Download file 4.59 kB
-
https://huggingface.co/OneScience-Group/ConvLSTM/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/ConvLSTM/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/ConvLSTM/resolve/main/scripts/train.py
4.59 kB
| """Train the ConvLSTM radar encoder-forecaster with full-sequence BPTT.""" | |
| 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.convlstm import ConvLSTM | |
| class RadarDataset(Dataset): | |
| def __init__(self, path, config): | |
| self.data = np.load(path) | |
| data = config["data"] | |
| if str(self.data["format_version"]) != data["format_version"]: | |
| raise ValueError("incompatible radar data format") | |
| expected_input = (int(data["input_frames"]), int(data["channels"]), int(data["height"]), int(data["width"])) | |
| expected_target = (int(data["output_frames"]), int(data["channels"]), int(data["height"]), int(data["width"])) | |
| if self.data["inputs"].shape[1:] != expected_input or self.data["targets"].shape[1:] != expected_target: | |
| raise ValueError("radar tensors do not preserve the paper dimensions") | |
| def __len__(self): | |
| return len(self.data["inputs"]) | |
| def __getitem__(self, index): | |
| return torch.from_numpy(self.data["inputs"][index]).float(), torch.from_numpy(self.data["targets"][index]).float() | |
| def device_from_config(config, rank=0): | |
| if config["runtime"]["device"] == "auto": | |
| return torch.device("cuda", rank) if torch.cuda.is_available() else torch.device("cpu") | |
| return torch.device(config["runtime"]["device"]) | |
| 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) | |
| dataset = RadarDataset(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 = ConvLSTM(config["model"]).to(device) | |
| if distributed: | |
| model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) | |
| optimizer = torch.optim.RMSprop(model.parameters(), lr=float(config["train"]["learning_rate"]), | |
| alpha=float(config["train"]["rmsprop_alpha"]), | |
| weight_decay=float(config["train"]["weight_decay"])) | |
| history = [] | |
| for epoch in range(int(config["train"]["epochs"])): | |
| model.train() | |
| total, steps = 0.0, 0 | |
| for inputs, targets in loader: | |
| _, logits = model(inputs.to(device)) | |
| patched_target = torch.nn.functional.pixel_unshuffle(targets.to(device).flatten(0, 1), | |
| int(config["model"]["patch_size"])).unflatten(0, targets.shape[:2]) | |
| loss = torch.nn.functional.binary_cross_entropy_with_logits(logits, patched_target) | |
| optimizer.zero_grad(set_to_none=True) | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), float(config["train"]["gradient_clip_norm"])) | |
| optimizer.step() | |
| total += float(loss.detach()) | |
| steps += 1 | |
| metrics = {"epoch": epoch + 1, "binary_cross_entropy": total / max(steps, 1)} | |
| history.append(metrics) | |
| if rank == 0: | |
| print(f"epoch={epoch + 1} binary_cross_entropy={metrics['binary_cross_entropy']:.6f}") | |
| if rank == 0: | |
| checkpoint, metrics_path = ROOT / config["paths"]["checkpoint"], 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"], | |
| "format_version": config["data"]["format_version"]}, checkpoint) | |
| metrics_path.write_text(json.dumps({"history": history}, indent=2) + "\n") | |
| if distributed: | |
| torch.distributed.destroy_process_group() | |
| if __name__ == "__main__": | |
| main() | |