sra-trajectory-code / LED /trainer /train_sport_pretrain.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
9.74 kB
"""
Stage-1 pretraining for LED on sport datasets (soccer / football).
This is the training stage that LED's public README leaves on its TODO list:
trains the core 100-step DDPM denoiser (TransformerDenoisingModel) from
scratch on a sport dataset using the standard eps-MSE loss. The resulting
checkpoint is saved to `cfg.pretrained_core_denoising_model` so that
train_sport_led.py (stage 2) can load it as a frozen refiner.
Optimizes: self.model (core denoiser) only.
Loss: noise_estimation_loss — predict eps at a uniformly-sampled diffusion
step t ∈ [0, n_steps) and regress via MSE.
"""
import os
import time
import torch
import random
import numpy as np
import torch.nn as nn
from utils.config import Config
from utils.utils import print_log
from torch.utils.data import DataLoader
from torch.utils.tensorboard import SummaryWriter
from data.dataloader_sport import SportDataset, sport_seq_collate
from models.model_diffusion import TransformerDenoisingModel as CoreDenoisingModel
class Trainer:
def __init__(self, config):
if torch.cuda.is_available():
torch.cuda.set_device(config.gpu)
self.device = torch.device('cuda') if config.cuda else torch.device('cpu')
self.cfg = Config(config.cfg, config.info)
# ------------------------- data -------------------------
self.num_agents = self.cfg.num_agents
train_dset = SportDataset(
data_dir = self.cfg.data_dir,
num_agents = self.num_agents,
obs_len = self.cfg.past_frames,
pred_len = self.cfg.future_frames,
split = 'train',
)
val_dset = SportDataset(
data_dir = self.cfg.data_dir,
num_agents = self.num_agents,
obs_len = self.cfg.past_frames,
pred_len = self.cfg.future_frames,
split = 'val',
)
pre_bs = self.cfg.pretrain['train_batch_size']
self.train_loader = DataLoader(
train_dset, batch_size=pre_bs, shuffle=True,
num_workers=4, collate_fn=sport_seq_collate, pin_memory=True)
self.val_loader = DataLoader(
val_dset, batch_size=self.cfg.test_batch_size, shuffle=False,
num_workers=4, collate_fn=sport_seq_collate, pin_memory=True)
self.traj_mean = torch.FloatTensor(self.cfg.traj_mean).cuda().unsqueeze(0).unsqueeze(0).unsqueeze(0)
self.traj_scale = float(self.cfg.traj_scale)
self.per_scene_norm = bool(self.cfg.get('per_scene_norm', False))
# ------------------------- diffusion parameters -------------------------
self.n_steps = self.cfg.diffusion.steps
self.betas = self.make_beta_schedule(
schedule=self.cfg.diffusion.beta_schedule, n_timesteps=self.n_steps,
start=self.cfg.diffusion.beta_start, end=self.cfg.diffusion.beta_end).cuda()
self.alphas = 1 - self.betas
self.alphas_prod = torch.cumprod(self.alphas, 0)
self.alphas_bar_sqrt = torch.sqrt(self.alphas_prod)
self.one_minus_alphas_bar_sqrt = torch.sqrt(1 - self.alphas_prod)
# ------------------------- model (the only thing being trained) -------------------------
self.model = CoreDenoisingModel().cuda()
pre_lr = float(self.cfg.pretrain['lr'])
pre_decay_step = int(self.cfg.pretrain['decay_step'])
pre_decay_gamma = float(self.cfg.pretrain['decay_gamma'])
self.pre_epochs = int(self.cfg.pretrain['num_epochs'])
self.opt = torch.optim.AdamW(self.model.parameters(), lr=pre_lr)
self.scheduler = torch.optim.lr_scheduler.StepLR(
self.opt, step_size=pre_decay_step, gamma=pre_decay_gamma)
# ------------------------- logs -------------------------
self.log = open(os.path.join(self.cfg.log_dir, 'log.txt'), 'a+')
self.tb = SummaryWriter(log_dir=os.path.join(self.cfg.log_dir, 'tb'))
self.global_step = 0
self.print_model_param(self.model, name='Core Denoising Model')
self.ckpt_path = self.cfg.pretrained_core_denoising_model
def print_model_param(self, model: nn.Module, name: str = 'Model'):
total = sum(p.numel() for p in model.parameters())
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
print_log(f'[{name}] Trainable/Total: {trainable}/{total}', self.log)
def make_beta_schedule(self, schedule='linear', n_timesteps=1000, start=1e-5, end=1e-2):
if schedule == 'linear':
betas = torch.linspace(start, end, n_timesteps)
elif schedule == 'quad':
betas = torch.linspace(start ** 0.5, end ** 0.5, n_timesteps) ** 2
elif schedule == 'sigmoid':
betas = torch.linspace(-6, 6, n_timesteps)
betas = torch.sigmoid(betas) * (end - start) + start
return betas
def extract(self, inp, t, x):
shape = x.shape
out = torch.gather(inp, 0, t.to(inp.device))
reshape = [t.shape[0]] + [1] * (len(shape) - 1)
return out.reshape(*reshape)
# ------------------------------------------------------------------
# Data preprocess — identical math to LED NBA trainer but with
# num_agents parameterized so the block-diagonal mask is A×A, not 11×11.
# ------------------------------------------------------------------
def data_preprocess(self, data):
A = self.num_agents
batch_size = data['pre_motion_3D'].shape[0]
traj_mask = torch.zeros(batch_size * A, batch_size * A).cuda()
for i in range(batch_size):
traj_mask[i * A:(i + 1) * A, i * A:(i + 1) * A] = 1.
pre = data['pre_motion_3D'].cuda()
fut = data['fut_motion_3D'].cuda()
initial_pos = pre[:, :, -1:]
if self.per_scene_norm:
scene_center = pre[:, :, -1, :].mean(dim=1, keepdim=True).unsqueeze(2)
past_traj_abs = ((pre - scene_center) / self.traj_scale).contiguous().view(-1, self.cfg.past_frames, 2)
else:
past_traj_abs = ((pre - self.traj_mean) / self.traj_scale).contiguous().view(-1, self.cfg.past_frames, 2)
past_traj_rel = ((pre - initial_pos) / self.traj_scale).contiguous().view(-1, self.cfg.past_frames, 2)
past_traj_vel = torch.cat(
(past_traj_rel[:, 1:] - past_traj_rel[:, :-1],
torch.zeros_like(past_traj_rel[:, -1:])), dim=1)
past_traj = torch.cat((past_traj_abs, past_traj_rel, past_traj_vel), dim=-1)
fut_traj = ((fut - initial_pos) / self.traj_scale).contiguous().view(-1, self.cfg.future_frames, 2)
return batch_size, traj_mask, past_traj, fut_traj
# ------------------------------------------------------------------
# Pretraining loss (LED's noise_estimation_loss, unchanged)
# ------------------------------------------------------------------
def noise_estimation_loss(self, x, y_0, mask):
batch_size = x.shape[0]
t = torch.randint(0, self.n_steps, size=(batch_size // 2 + 1,)).to(x.device)
t = torch.cat([t, self.n_steps - t - 1], dim=0)[:batch_size]
a = self.extract(self.alphas_bar_sqrt, t, y_0)
beta = self.extract(self.betas, t, y_0)
am1 = self.extract(self.one_minus_alphas_bar_sqrt, t, y_0)
e = torch.randn_like(y_0)
y = y_0 * a + e * am1
out = self.model(y, beta, x, mask)
return (e - out).square().mean()
# ------------------------------------------------------------------
# Training loop
# ------------------------------------------------------------------
def fit(self):
best_val = float('inf')
for epoch in range(self.pre_epochs):
self.model.train()
loss_sum, n_batches = 0.0, 0
for data in self.train_loader:
_, mask, past, fut = self.data_preprocess(data)
loss = self.noise_estimation_loss(past, fut, mask)
self.opt.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)
self.opt.step()
loss_sum += float(loss.item())
n_batches += 1
self.tb.add_scalar('pretrain_step/loss', loss.item(), self.global_step)
self.global_step += 1
train_loss = loss_sum / max(1, n_batches)
print_log(f'[{time.strftime("%Y-%m-%d %H:%M:%S")}] Pretrain Epoch {epoch} train_mse={train_loss:.6f}', self.log)
self.tb.add_scalar('pretrain_epoch/train_loss', train_loss, epoch)
self.tb.add_scalar('pretrain_epoch/lr', self.opt.param_groups[0]['lr'], epoch)
# Quick val MSE to monitor
self.model.eval()
val_sum, val_n = 0.0, 0
with torch.no_grad():
for data in self.val_loader:
_, mask, past, fut = self.data_preprocess(data)
val_sum += float(self.noise_estimation_loss(past, fut, mask).item())
val_n += 1
val_loss = val_sum / max(1, val_n)
print_log(f' val_mse={val_loss:.6f}', self.log)
self.tb.add_scalar('pretrain_epoch/val_loss', val_loss, epoch)
if val_loss < best_val:
best_val = val_loss
os.makedirs(os.path.dirname(self.ckpt_path), exist_ok=True)
torch.save({'model_dict': self.model.state_dict(), 'epoch': epoch}, self.ckpt_path)
print_log(f' -> saved {self.ckpt_path} (best val {best_val:.6f})', self.log)
self.scheduler.step()
self.tb.flush()
self.tb.close()