File size: 6,790 Bytes
d4cbafd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Stage-1 pretraining for LED on SDD.
Variable-A scenes with batch_size=1 + gradient accumulation.
past_frames=8, future_frames=12. Pixel coordinates.
"""

import os, time, torch, 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_sdd import SDDDataset, sdd_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)
        self.grad_accum = getattr(config, 'grad_accum', 32)

        train_dset = SDDDataset(obs_len=self.cfg.past_frames,
                                pred_len=self.cfg.future_frames, split='train')
        test_dset  = SDDDataset(obs_len=self.cfg.past_frames,
                                pred_len=self.cfg.future_frames, split='test')
        self.train_loader = DataLoader(train_dset, batch_size=1, shuffle=True,
            num_workers=2, collate_fn=sdd_seq_collate)
        self.val_loader = DataLoader(test_dset, batch_size=1, shuffle=False,
            num_workers=2, collate_fn=sdd_seq_collate)

        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.n_steps = self.cfg.diffusion.steps
        self.betas = self._make_beta_schedule(
            self.cfg.diffusion.beta_schedule, self.n_steps,
            self.cfg.diffusion.beta_start, 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)

        self.model = CoreDenoisingModel(past_len=self.cfg.past_frames).cuda()

        pre_lr = float(self.cfg.pretrain['lr'])
        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=int(self.cfg.pretrain.get('decay_step', 30)),
            gamma=float(self.cfg.pretrain.get('decay_gamma', 0.5)))

        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.ckpt_path = self.cfg.pretrained_core_denoising_model
        total = sum(p.numel() for p in self.model.parameters())
        print_log(f'Core Denoiser params: {total:,}', self.log)

    def _make_beta_schedule(self, schedule, n, start, end):
        if schedule == 'linear': return torch.linspace(start, end, n)
        elif schedule == 'quad': return torch.linspace(start**0.5, end**0.5, n)**2
        return torch.linspace(start, end, n)

    def _extract(self, a, t, x):
        out = torch.gather(a, 0, t.to(a.device))
        return out.reshape(t.shape[0], *([1] * (len(x.shape) - 1)))

    def data_preprocess(self, data):
        pre = data['pre_motion_3D'].cuda()       # [1, A, 8, 2]
        fut = data['fut_motion_3D'].cuda()       # [1, A, 12, 2]
        A = pre.size(1)
        initial_pos = pre[:, :, -1:]              # [1, A, 1, 2]
        past_abs = ((pre - self.traj_mean) / self.traj_scale).contiguous().view(-1, self.cfg.past_frames, 2)
        past_rel = ((pre - initial_pos) / self.traj_scale).contiguous().view(-1, self.cfg.past_frames, 2)
        past_vel = torch.cat([past_rel[:, 1:] - past_rel[:, :-1],
                              torch.zeros_like(past_rel[:, -1:])], dim=1)
        past_traj = torch.cat([past_abs, past_rel, past_vel], dim=-1)  # [A, 8, 6]
        fut_traj  = ((fut - initial_pos) / self.traj_scale).contiguous().view(-1, self.cfg.future_frames, 2)
        mask = torch.ones(A, A).cuda()
        return A, mask, past_traj, fut_traj

    def noise_estimation_loss(self, x, y_0, mask):
        B = x.shape[0]
        t = torch.randint(0, self.n_steps, size=(B // 2 + 1,)).to(x.device)
        t = torch.cat([t, self.n_steps - t - 1], dim=0)[:B]
        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()

    def fit(self):
        best_val = float('inf')
        for epoch in range(self.pre_epochs):
            self.model.train()
            loss_sum, n = 0.0, 0
            self.opt.zero_grad()
            for i, data in enumerate(self.train_loader):
                _, mask, past, fut = self.data_preprocess(data)
                loss = self.noise_estimation_loss(past, fut, mask)
                (loss / self.grad_accum).backward()
                if (i + 1) % self.grad_accum == 0:
                    nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)
                    self.opt.step(); self.opt.zero_grad()
                loss_sum += loss.item(); n += 1
                self.tb.add_scalar('pretrain_step/loss', loss.item(), self.global_step)
                self.global_step += 1
            self.opt.step(); self.opt.zero_grad()

            train_loss = loss_sum / max(1, n)
            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)

            if (epoch + 1) % 5 == 0:
                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 += 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}', self.log)

            self.scheduler.step()
        self.tb.flush(); self.tb.close()