File size: 7,507 Bytes
387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d 7a2d30b 387a20d | 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 | 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()
|