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 data.dataloader_nba import NBADataset, seq_collate from models.model_led_initializer import LEDInitializer as InitializationModel from models.model_diffusion import TransformerDenoisingModel as CoreDenoisingModel import pdb NUM_Tau = 5 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) # ------------------------- prepare train/test data loader ------------------------- train_dset = NBADataset( obs_len=self.cfg.past_frames, pred_len=self.cfg.future_frames, training=True) self.train_loader = DataLoader( train_dset, batch_size=self.cfg.train_batch_size, shuffle=True, num_workers=4, collate_fn=seq_collate, pin_memory=True) test_dset = NBADataset( obs_len=self.cfg.past_frames, pred_len=self.cfg.future_frames, training=False) self.test_loader = DataLoader( test_dset, batch_size=self.cfg.test_batch_size, shuffle=False, num_workers=4, collate_fn=seq_collate, pin_memory=True) # data normalization parameters self.traj_mean = torch.FloatTensor(self.cfg.traj_mean).cuda().unsqueeze(0).unsqueeze(0).unsqueeze(0) self.traj_scale = self.cfg.traj_scale # ------------------------- define diffusion parameters ------------------------- self.n_steps = self.cfg.diffusion.steps # define total diffusion steps # make beta schedule and calculate the parameters used in denoising process. 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) # ------------------------- define models ------------------------- self.model = CoreDenoisingModel().cuda() # load pretrained models model_cp = torch.load(self.cfg.pretrained_core_denoising_model, map_location='cpu') self.model.load_state_dict(model_cp['model_dict']) self.model_initializer = InitializationModel(t_h=10, d_h=6, t_f=20, d_f=2, k_pred=20).cuda() self.opt = torch.optim.AdamW(self.model_initializer.parameters(), lr=config.learning_rate) self.scheduler_model = torch.optim.lr_scheduler.StepLR(self.opt, step_size=self.cfg.decay_step, gamma=self.cfg.decay_gamma) # ------------------------- prepare logs ------------------------- self.log = open(os.path.join(self.cfg.log_dir, 'log.txt'), 'a+') self.print_model_param(self.model, name='Core Denoising Model') self.print_model_param(self.model_initializer, name='Initialization Model') # temporal reweight in the loss, it is not necessary. self.temporal_reweight = torch.FloatTensor([21 - i for i in range(1, 21)]).cuda().unsqueeze(0).unsqueeze(0) / 10 def print_model_param(self, model: nn.Module, name: str = 'Model') -> None: ''' Count the trainable/total parameters in `model`. ''' total_num = sum(p.numel() for p in model.parameters()) trainable_num = sum(p.numel() for p in model.parameters() if p.requires_grad) print_log("[{}] Trainable/Total: {}/{}".format(name, trainable_num, total_num), self.log) return None def make_beta_schedule(self, schedule: str = 'linear', n_timesteps: int = 1000, start: float = 1e-5, end: float = 1e-2) -> torch.Tensor: ''' Make beta schedule. Parameters ---- schedule: str, in ['linear', 'quad', 'sigmoid'], n_timesteps: int, diffusion steps, start: float, beta start, `start