File size: 9,142 Bytes
1558db5 2d7eead 1558db5 2d7eead 1558db5 2d7eead 1558db5 2d7eead 1558db5 2d7eead 1558db5 2d7eead 1558db5 2d7eead 1558db5 2d7eead 1558db5 2d7eead 1558db5 2d7eead 1558db5 2d7eead 1558db5 2d7eead 1558db5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 | """Train compact SkySense on NPZ multi-modal temporal samples."""
import importlib.util
import json
import os
import random
from pathlib import Path
import numpy as np
import torch
import yaml
from torch import distributed as dist
from torch.nn.parallel import DistributedDataParallel
from torch.utils.data import DataLoader, Dataset, DistributedSampler
ROOT = Path(__file__).resolve().parents[1]
def load_model_class():
spec = importlib.util.spec_from_file_location("skysense_model", ROOT / "model" / "skysense.py")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.SkySense
def load_config():
with (ROOT / "conf" / "config.yaml").open(encoding="utf-8") as handle:
return yaml.safe_load(handle)
class NPZDataset(Dataset):
def __init__(self, path):
archive = np.load(path)
sample_keys = {"hr", "s2", "s1", "dates_hr", "dates_s2", "dates_s1", "region", "labels"}
missing = sample_keys.difference(archive.files)
if missing:
raise ValueError(f"Dataset {path} is missing arrays: {sorted(missing)}")
self.arrays = {key: archive[key] for key in sample_keys}
self.data_source = str(archive["data_source"]) if "data_source" in archive.files else "unknown"
self.protocol = str(archive["protocol"]) if "protocol" in archive.files else "unknown"
def __len__(self):
return len(self.arrays["hr"])
def __getitem__(self, index):
return {key: torch.as_tensor(value[index]) for key, value in self.arrays.items()}
def setup_device(config):
world_size = int(os.environ.get("WORLD_SIZE", "1"))
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
use_accelerator = torch.cuda.is_available() and config["runtime"].get("device", "auto") != "cpu"
if world_size > 1:
backend = "nccl" if use_accelerator else "gloo"
dist.init_process_group(backend=backend)
if use_accelerator:
torch.cuda.set_device(local_rank)
return torch.device("cuda", local_rank), world_size, local_rank
return torch.device("cpu"), world_size, local_rank
def main():
config = load_config()
seed = config["seed"]
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
device, world_size, local_rank = setup_device(config)
amp_enabled = bool(config["train"].get("amp", True) and device.type == "cuda")
dataset_path = ROOT / config["data"]["root"] / "train.npz"
if not dataset_path.exists():
raise FileNotFoundError(
f"Missing training data: {dataset_path.relative_to(ROOT)}. "
"Run `python scripts/fake_data.py` for a synthetic connectivity test."
)
dataset = NPZDataset(dataset_path)
expected = {
"hr": (config["data"]["hr_timesteps"], config["data"]["hr_channels"], config["data"]["hr_size"], config["data"]["hr_size"]),
"s2": (config["data"]["s2_timesteps"], config["data"]["s2_channels"], config["data"]["s2_size"], config["data"]["s2_size"]),
"s1": (config["data"]["s1_timesteps"], config["data"]["s1_channels"], config["data"]["s1_size"], config["data"]["s1_size"]),
"labels": (config["data"]["hr_size"], config["data"]["hr_size"]),
"dates_hr": (config["data"]["hr_timesteps"],),
"dates_s2": (config["data"]["s2_timesteps"],),
"dates_s1": (config["data"]["s1_timesteps"],),
"region": (),
}
for key, shape in expected.items():
if tuple(dataset.arrays[key].shape[1:]) != shape:
raise ValueError(f"Expected {key} shaped [N,{','.join(map(str, shape))}], got {dataset.arrays[key].shape}")
if len(dataset.arrays[key]) != len(dataset):
raise ValueError(f"Array {key} has {len(dataset.arrays[key])} samples, expected {len(dataset)}")
for key in ("hr", "s2", "s1"):
if not np.issubdtype(dataset.arrays[key].dtype, np.floating):
raise TypeError(f"{key} must use a floating dtype, got {dataset.arrays[key].dtype}")
for key in ("dates_hr", "dates_s2", "dates_s1", "region", "labels"):
if dataset.arrays[key].dtype != np.int64:
raise TypeError(f"{key} must use int64, got {dataset.arrays[key].dtype}")
for key in ("dates_hr", "dates_s2", "dates_s1"):
if np.any((dataset.arrays[key] < 0) | (dataset.arrays[key] > 364)):
raise ValueError(f"{key} must contain day-of-year values in [0, 364]")
if np.any((dataset.arrays["region"] < 0) | (dataset.arrays["region"] >= config["model"]["num_regions"])):
raise ValueError(f"region IDs must be in [0, {config['model']['num_regions'] - 1}]")
if np.any((dataset.arrays["labels"] < 0) | (dataset.arrays["labels"] >= config["data"]["num_classes"])):
raise ValueError(f"labels must be in [0, {config['data']['num_classes'] - 1}]")
if local_rank == 0:
print(
f"data_source={dataset.data_source} protocol={dataset.protocol} "
f"samples={len(dataset)} hr={config['data']['hr_size']} s2={config['data']['s2_size']} s1={config['data']['s1_size']}"
)
sampler = DistributedSampler(dataset, shuffle=True) if world_size > 1 else None
loader = DataLoader(
dataset,
batch_size=config["train"]["batch_size"],
shuffle=sampler is None,
sampler=sampler,
num_workers=config["train"]["num_workers"],
)
SkySense = load_model_class()
model = SkySense(
**config["model"],
hr_channels=config["data"]["hr_channels"],
s2_channels=config["data"]["s2_channels"],
s1_channels=config["data"]["s1_channels"],
num_classes=config["data"]["num_classes"],
).to(device)
if world_size > 1:
model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None)
optimizer = torch.optim.AdamW(model.parameters(), lr=config["train"]["learning_rate"], weight_decay=config["train"]["weight_decay"])
scaler = torch.amp.GradScaler("cuda", enabled=amp_enabled)
final_loss = float("nan")
final_segmentation = float("nan")
final_alignment = float("nan")
for epoch in range(config["train"]["epochs"]):
if sampler is not None:
sampler.set_epoch(epoch)
model.train()
totals = torch.zeros(4, dtype=torch.float64, device=device)
for batch in loader:
batch = {key: value.to(device) for key, value in batch.items()}
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=amp_enabled):
output = model(batch["hr"], batch["s2"], batch["s1"], batch["dates_hr"], batch["dates_s2"], batch["dates_s1"], batch["region"])
base_model = model.module if hasattr(model, "module") else model
segmentation = torch.nn.functional.cross_entropy(output["logits"], batch["labels"])
alignment = base_model.cross_modal_alignment_loss(output["features"])
loss = segmentation + config["train"]["alignment_weight"] * alignment
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
batch_size = batch["hr"].shape[0]
totals += torch.tensor([loss.item() * batch_size, segmentation.item() * batch_size,
alignment.item() * batch_size, batch_size],
dtype=torch.float64, device=device)
if world_size > 1:
dist.all_reduce(totals, op=dist.ReduceOp.SUM)
final_loss = float((totals[0] / totals[3]).item())
final_segmentation = float((totals[1] / totals[3]).item())
final_alignment = float((totals[2] / totals[3]).item())
if local_rank == 0:
print(f"epoch={epoch + 1} loss={final_loss:.6f}")
if local_rank == 0:
checkpoint = ROOT / config["paths"]["checkpoint"]
checkpoint.parent.mkdir(parents=True, exist_ok=True)
base_model = model.module if hasattr(model, "module") else model
torch.save({"model": base_model.state_dict(), "config": config, "final_loss": final_loss}, checkpoint)
metrics = ROOT / config["paths"]["training_metrics"]
metrics.parent.mkdir(parents=True, exist_ok=True)
metrics.write_text(
json.dumps(
{
"final_loss": final_loss,
"segmentation_loss": final_segmentation,
"alignment_loss": final_alignment,
"backward_completed": True,
"amp_enabled": amp_enabled,
"epochs": config["train"]["epochs"],
"samples": len(dataset),
"data_source": dataset.data_source,
"protocol": dataset.protocol,
},
indent=2,
)
+ "\n",
encoding="utf-8",
)
print(f"checkpoint={checkpoint.relative_to(ROOT)} final_loss={final_loss:.6f}")
if world_size > 1:
dist.destroy_process_group()
if __name__ == "__main__":
main()
|