DOFA / scripts /train.py
zhangrenchao's picture
Upload DOFA model package
1d4cac8 verified
Raw
History Blame Contribute Delete
7.68 kB
"""Train compact DOFA across strictly validated, variable-channel sensors."""
import importlib.util
import json
import os
import random
from contextlib import nullcontext
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("dofa_model", ROOT / "model" / "dofa.py")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.DOFA
def expected_wavelengths(modality):
if modality.get("wavelength_mode") == "synthetic_uniform":
return modality["wavelength_start"] + np.arange(modality["channels"], dtype=np.float32) * modality["wavelength_step"]
return np.asarray(modality["wavelengths"], dtype=np.float32)
def scalar(archive, key, default=None):
if key not in archive:
if default is not None:
return default
raise ValueError(f"NPZ is missing required metadata: {key}")
if archive[key].ndim != 0:
raise ValueError(f"NPZ metadata {key} must be a scalar")
return archive[key].item()
class SensorDataset(Dataset):
def __init__(self, path, name, data_config):
archive = np.load(path)
if "images" not in archive or "wavelengths" not in archive:
raise ValueError(f"{path} must contain images and wavelengths")
self.images = archive["images"]
wavelengths = archive["wavelengths"]
modality = data_config["modalities"][name]
expected = expected_wavelengths(modality)
if self.images.ndim != 4:
raise ValueError(f"{path}: images must be NCHW")
if self.images.dtype != np.float32:
raise ValueError(f"{path}: images must use float32")
if self.images.shape[1:] != (modality["channels"], data_config["image_size"], data_config["image_size"]):
raise ValueError(f"{path}: image shape does not match configured channels/224x224")
if wavelengths.shape != (modality["channels"],) or not np.issubdtype(wavelengths.dtype, np.floating):
raise ValueError(f"{path}: wavelengths must be a floating [C] array")
if not np.isfinite(wavelengths).all() or not np.allclose(wavelengths, expected, rtol=1e-5, atol=1e-6):
raise ValueError(f"{path}: wavelengths do not match configured sensor wavelengths")
self.modality = str(scalar(archive, "modality"))
self.protocol = str(scalar(archive, "protocol"))
self.data_source = str(scalar(archive, "data_source", "unknown"))
if self.modality != name:
raise ValueError(f"{path}: modality {self.modality} does not match {name}")
if self.protocol != data_config["protocol"]:
raise ValueError(f"{path}: protocol {self.protocol} does not match config")
self.wavelengths = torch.from_numpy(wavelengths.astype("float32", copy=False))
def __len__(self):
return len(self.images)
def __getitem__(self, index):
return torch.from_numpy(self.images[index])
def reduced_average(total, count, device, distributed):
values = torch.tensor([total, count], dtype=torch.float64, device=device)
if distributed:
dist.all_reduce(values, op=dist.ReduceOp.SUM)
if values[1].item() == 0:
raise RuntimeError("Training processed no batches")
return (values[0] / values[1]).item()
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text(encoding="utf-8"))
world_size = int(os.environ.get("WORLD_SIZE", "1"))
global_rank = int(os.environ.get("RANK", "0"))
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
distributed = world_size > 1
if distributed:
dist.init_process_group("nccl" if torch.cuda.is_available() else "gloo")
requested = config["runtime"]["device"]
use_accelerator = torch.cuda.is_available() and requested != "cpu"
device = torch.device(f"cuda:{local_rank}" if use_accelerator else "cpu")
if use_accelerator:
torch.cuda.set_device(local_rank)
seed = config["seed"] + global_rank
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
data_root = ROOT / config["data"]["root"]
loaders = []
for modality in config["data"]["modalities"]:
path = data_root / f"train_{modality}.npz"
if not path.exists():
raise FileNotFoundError(f"Missing training data: {path.relative_to(ROOT)}. Run `python scripts/fake_data.py` first.")
dataset = SensorDataset(path, modality, config["data"])
sampler = DistributedSampler(dataset, shuffle=True) if distributed else None
loader = DataLoader(dataset, batch_size=config["training"]["batch_size"],
shuffle=sampler is None, sampler=sampler,
num_workers=config["training"]["num_workers"])
loaders.append((dataset, loader, sampler))
model = load_model_class()(**config["model"]).to(device)
if distributed:
model = DistributedDataParallel(model, device_ids=[local_rank] if use_accelerator else None)
optimizer = torch.optim.AdamW(model.parameters(), lr=config["training"]["learning_rate"],
weight_decay=config["training"]["weight_decay"])
amp = bool(config["training"]["amp"] and use_accelerator)
scaler = torch.amp.GradScaler("cuda", enabled=amp)
autocast = (lambda: torch.amp.autocast("cuda", enabled=True)) if amp else nullcontext
history = []
for epoch in range(config["training"]["epochs"]):
model.train()
records = {}
for dataset, loader, sampler in loaders:
if sampler is not None:
sampler.set_epoch(epoch)
total = count = 0
for images in loader:
optimizer.zero_grad(set_to_none=True)
with autocast():
output = model(images.to(device), dataset.wavelengths.to(device))
scaler.scale(output["loss"]).backward()
scaler.step(optimizer)
scaler.update()
total += output["loss"].item()
count += 1
records[dataset.modality] = reduced_average(total, count, device, distributed)
history.append({"epoch": epoch + 1, "reconstruction_loss": records})
if global_rank == 0:
print(f"epoch={epoch + 1} " + " ".join(f"{name}={value:.6f}" for name, value in records.items()))
if global_rank == 0:
checkpoint = ROOT / config["paths"]["checkpoint"]
metrics = ROOT / config["paths"]["training_metrics"]
checkpoint.parent.mkdir(parents=True, exist_ok=True)
metrics.parent.mkdir(parents=True, exist_ok=True)
base_model = model.module if hasattr(model, "module") else model
sources = sorted({dataset.data_source for dataset, _, _ in loaders})
torch.save({"model": base_model.state_dict(), "config": config,
"protocol": config["data"]["protocol"], "data_sources": sources}, checkpoint)
metrics.write_text(json.dumps({"history": history,
"modalities": list(config["data"]["modalities"]),
"protocol": config["data"]["protocol"],
"data_sources": sources}, indent=2) + "\n", encoding="utf-8")
print(f"checkpoint={checkpoint.relative_to(ROOT)}")
if distributed:
dist.destroy_process_group()
if __name__ == "__main__":
main()