SpectralGPT / scripts /train.py
zhangrenchao's picture
Update SpectralGPT model package
7a2d30b verified
Raw
History Blame Contribute Delete
7.51 kB
import argparse
import json
import math
import os
from pathlib import Path
import sys
import numpy as np
import torch
import torch.nn.functional as F
import yaml
from torch.nn.parallel import DistributedDataParallel
from torch.utils.data import DataLoader, Dataset, DistributedSampler
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from model.spectralgpt import SpectralGPT
class SpectralDataset(Dataset):
def __init__(self, path, image_size, stage):
with np.load(path) as data:
if "images" not in data.files:
raise ValueError(f"Dataset {path} is missing images")
self.images = data["images"].copy()
self.data_source = str(data["data_source"]) if "data_source" in data.files else "unknown"
self.protocol = str(data["protocol"]) if "protocol" in data.files else "unknown"
self.normalization = str(data["normalization"]) if "normalization" in data.files else "unknown"
stored_stage = str(data["stage"]) if "stage" in data.files else "unknown"
expected = (12, image_size, image_size)
if self.images.dtype != np.float32 or self.images.ndim != 4 or tuple(self.images.shape[1:]) != expected:
raise ValueError(f"Expected float32 [N,{','.join(map(str, expected))}], got {self.images.dtype} {self.images.shape}")
if stored_stage != stage:
raise ValueError(f"Expected stage metadata {stage}, got {stored_stage}")
def __len__(self):
return len(self.images)
def __getitem__(self, index):
return torch.from_numpy(self.images[index])
def resize_spatial_position(state, old_size, new_size, patch_size):
if old_size == new_size:
return state
key = "spatial_pos"
position = state[key]
old_grid, new_grid = old_size // patch_size, new_size // patch_size
if position.shape[1] != old_grid * old_grid:
raise ValueError("Checkpoint spatial position shape does not match previous stage")
position = position.reshape(1, old_grid, old_grid, -1).permute(0, 3, 1, 2)
state[key] = F.interpolate(position, size=(new_grid, new_grid), mode="bicubic", align_corners=False).permute(0, 2, 3, 1).reshape(1, new_grid * new_grid, -1)
return state
def main():
parser = argparse.ArgumentParser(description="Progressive two-stage SpectralGPT training")
parser.add_argument("--config", default="conf/config.yaml")
args = parser.parse_args()
with open(args.config, encoding="utf-8") as handle:
config = yaml.safe_load(handle)
distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1
rank = int(os.environ.get("RANK", "0"))
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
requested = config["runtime"]["device"]
device = torch.device(f"cuda:{local_rank}" if torch.cuda.is_available() and requested != "cpu" else "cpu")
if device.type == "cuda":
torch.cuda.set_device(device)
if distributed:
torch.distributed.init_process_group("nccl" if device.type == "cuda" else "gloo")
torch.manual_seed(config["runtime"]["seed"] + rank)
amp_enabled = bool(config["training"].get("amp", True) and device.type == "cuda")
save_dir = Path(config["training"]["save_dir"])
history = []
previous_state = None
previous_size = None
for stage in config["stages"]:
path = Path(stage["train_path"])
if not path.exists():
raise FileNotFoundError(f"Missing {stage['name']} data: {path}. Run scripts/fake_data.py")
dataset = SpectralDataset(path, stage["image_size"], stage["name"])
sampler = DistributedSampler(dataset, shuffle=True) if distributed else None
loader = DataLoader(dataset, batch_size=config["training"]["batch_size"], sampler=sampler,
shuffle=sampler is None)
model = SpectralGPT(image_size=stage["image_size"], **config["model"])
if previous_state is not None:
model.load_state_dict(resize_spatial_position(previous_state, previous_size,
stage["image_size"], config["model"]["patch_size"]))
model = model.to(device)
if distributed:
model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None)
optimizer = torch.optim.AdamW(model.parameters(), lr=config["training"]["learning_rate"],
weight_decay=config["training"]["weight_decay"], betas=(0.9, 0.95))
scaler = torch.amp.GradScaler("cuda", enabled=amp_enabled)
for epoch in range(stage["epochs"]):
if sampler is not None:
sampler.set_epoch(epoch)
model.train()
totals = torch.zeros(5, dtype=torch.float64, device=device)
for images in loader:
images = images.to(device)
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=amp_enabled):
output = model(images)
scaler.scale(output["loss"]).backward()
scaler.step(optimizer)
scaler.update()
count = images.shape[0]
totals += torch.tensor([output[name].item() * count for name in
("loss", "masked_mse", "spectral_angle", "spectral_gradient")] + [count],
dtype=torch.float64, device=device)
if distributed:
torch.distributed.all_reduce(totals)
values = (totals[:4] / totals[4]).tolist()
record = {"stage": stage["name"], "dataset": stage["dataset"], "image_size": stage["image_size"],
"patch_size": config["model"]["patch_size"], "epoch": epoch + 1,
**dict(zip(("loss", "masked_mse", "spectral_angle", "spectral_gradient"), values))}
history.append(record)
if rank == 0:
print(f"stage={stage['name']} epoch={epoch + 1} size={stage['image_size']} loss={values[0]:.6f}")
base_model = model.module if distributed else model
previous_state = {key: value.detach().cpu() for key, value in base_model.state_dict().items()}
previous_size = stage["image_size"]
if rank == 0:
save_dir.mkdir(parents=True, exist_ok=True)
checkpoint = {"model": previous_state, "config": config, "stage": stage["name"],
"image_size": stage["image_size"], "stage_history": history,
"data_source": dataset.data_source, "protocol": dataset.protocol,
"normalization": dataset.normalization, "backward_completed": True,
"format": "spectralgpt-progressive-v2"}
torch.save(checkpoint, save_dir / f"{stage['name']}.pth")
if stage is config["stages"][-1]:
torch.save(checkpoint, Path(config["training"]["checkpoint"]))
if rank == 0:
metrics = Path(config["training"]["metrics"])
metrics.parent.mkdir(parents=True, exist_ok=True)
metrics.write_text(json.dumps({"stage_history": history, "backward_completed": True,
"amp_enabled": amp_enabled}, indent=2) + "\n", encoding="utf-8")
print(f"saved: {config['training']['checkpoint']}")
if distributed:
torch.distributed.destroy_process_group()
if __name__ == "__main__":
main()