OneScience-Group/FireCubeNet
025
1"""Train the ConvLSTM classifier with optional distributed data parallelism."""2 3import json4import os5import random6import sys7from pathlib import Path8 9import numpy as np10import torch11import yaml12from torch.nn.parallel import DistributedDataParallel13from torch.utils.data import DataLoader, Dataset, DistributedSampler14 15 16ROOT = Path(__file__).resolve().parents[1]17sys.path.insert(0, str(ROOT))18from model.firecubenet import FireCubeNet19 20 21class WildfireDataset(Dataset):22 def __init__(self, path, config):23 self.data = np.load(path)24 expected = config["data"]25 if str(self.data["format_version"]) != expected["format_version"]:26 raise ValueError("incompatible wildfire data format")27 expected_shape = (int(expected["sequence_days"]), int(expected["channels"]),28 int(expected["patch_height"]), int(expected["patch_width"]))29 if self.data["inputs"].ndim != 5 or self.data["inputs"].shape[1:] != expected_shape:30 raise ValueError(f"inputs must have shape [B,{','.join(map(str, expected_shape))}]")31 count = len(self.data["inputs"])32 if self.data["labels"].shape != (count, 1):33 raise ValueError("labels must have shape [B,1]")34 if self.data["coords"].shape != (count, 2) or self.data["timestamps_unix_s"].shape != (count,):35 raise ValueError("coords/timestamps shape mismatch")36 if not np.isfinite(self.data["inputs"]).all() or not np.isfinite(self.data["labels"]).all():37 raise ValueError("inputs and labels must be finite")38 if not np.isin(self.data["labels"], (0, 1)).all():39 raise ValueError("labels must be binary")40 cover_sum = self.data["inputs"][:, :, 15:25].sum(axis=2)41 if not np.allclose(cover_sum, 1.0, atol=1e-5):42 raise ValueError("land-cover fractions must sum to one")43 44 def __len__(self):45 return len(self.data["labels"])46 47 def __getitem__(self, index):48 return (torch.from_numpy(self.data["inputs"][index]).float(),49 torch.from_numpy(self.data["labels"][index]).float())50 51 52def device_from_config(config, local_rank=0):53 requested = config["runtime"]["device"]54 if requested == "auto":55 return torch.device("cuda", local_rank) if torch.cuda.is_available() else torch.device("cpu")56 return torch.device(requested)57 58 59def main():60 config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())61 seed = int(config["seed"])62 random.seed(seed)63 np.random.seed(seed)64 torch.manual_seed(seed)65 if torch.cuda.is_available():66 torch.cuda.manual_seed_all(seed)67 distributed = int(os.environ.get("WORLD_SIZE", "1")) > 168 local_rank = int(os.environ.get("LOCAL_RANK", "0"))69 if distributed:70 torch.distributed.init_process_group("nccl" if torch.cuda.is_available() else "gloo")71 rank = torch.distributed.get_rank() if distributed else 072 device = device_from_config(config, local_rank)73 if device.type == "cuda":74 torch.cuda.set_device(device)75 dataset = WildfireDataset(ROOT / config["data"]["root"] / "train.npz", config)76 sampler = DistributedSampler(dataset, shuffle=True, seed=seed) if distributed else None77 loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]),78 shuffle=sampler is None, sampler=sampler,79 num_workers=int(config["train"]["num_workers"]))80 channel_mean = dataset.data["inputs"].mean(axis=(0, 1, 3, 4)).astype(np.float32)81 channel_std = dataset.data["inputs"].std(axis=(0, 1, 3, 4)).clip(1e-6).astype(np.float32)82 model = FireCubeNet(**config["model"]).to(device)83 wrapped = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) if distributed else model84 optimizer = torch.optim.Adam(wrapped.parameters(), lr=float(config["train"]["learning_rate"]),85 weight_decay=float(config["train"]["weight_decay"]))86 criterion = torch.nn.BCEWithLogitsLoss()87 mean = torch.from_numpy(channel_mean).to(device).view(1, 1, -1, 1, 1)88 std = torch.from_numpy(channel_std).to(device).view(1, 1, -1, 1, 1)89 history = []90 for epoch in range(int(config["train"]["epochs"])):91 if sampler is not None:92 sampler.set_epoch(epoch)93 total, samples = 0.0, 094 wrapped.train()95 for inputs, labels in loader:96 inputs, labels = inputs.to(device), labels.to(device)97 logits = wrapped((inputs - mean) / std)98 loss = criterion(logits, labels)99 optimizer.zero_grad(set_to_none=True)100 loss.backward()101 torch.nn.utils.clip_grad_norm_(wrapped.parameters(), float(config["train"]["gradient_clip_norm"]))102 optimizer.step()103 total += float(loss.detach()) * len(inputs)104 samples += len(inputs)105 loss_sum = torch.tensor([total, samples], dtype=torch.float64, device=device)106 if distributed:107 torch.distributed.all_reduce(loss_sum)108 if rank == 0:109 history.append({"epoch": epoch + 1, "bce_with_logits": float(loss_sum[0] / loss_sum[1])})110 if rank == 0:111 checkpoint_path = ROOT / config["paths"]["checkpoint"]112 metrics_path = ROOT / config["paths"]["training_metrics"]113 checkpoint_path.parent.mkdir(parents=True, exist_ok=True)114 metrics_path.parent.mkdir(parents=True, exist_ok=True)115 bare_model = wrapped.module if distributed else wrapped116 torch.save({117 "model_state_dict": bare_model.state_dict(),118 "optimizer_state_dict": optimizer.state_dict(),119 "model_config": config["model"], "epoch": int(config["train"]["epochs"]),120 "channel_mean": channel_mean, "channel_std": channel_std,121 "format_version": config["data"]["format_version"], "seed": seed,122 }, checkpoint_path)123 metrics_path.write_text(json.dumps({"history": history}, indent=2) + "\n")124 print(f"checkpoint={checkpoint_path.relative_to(ROOT)} final_loss={history[-1]['bce_with_logits']:.6f}")125 if distributed:126 torch.distributed.destroy_process_group()127 128 129if __name__ == "__main__":130 main()131 