File size: 12,286 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
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
"""
Stage-2 LED trainer for SDD. Variable-A scenes, batch_size=1, grad_accum.
Optional graph module (--use_graph --use_v6_graph).
Eval: only target agent (index 0) counts toward ADE/FDE.
"""

import os, time, torch, random, 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_led_initializer import LEDInitializer as InitializationModel
from models.model_diffusion import TransformerDenoisingModel as CoreDenoisingModel

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)
        self.use_graph   = bool(getattr(config, 'use_graph', False))
        self.use_v6_graph = bool(getattr(config, 'use_v6_graph', False))
        self.residual_on = getattr(config, 'residual_on', 'y0')
        self.grad_accum  = getattr(config, 'grad_accum', 16)

        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.test_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()
        ckpt_path = self.cfg.pretrained_core_denoising_model
        if not os.path.isfile(ckpt_path):
            raise FileNotFoundError(f'Missing pretrained denoiser: {ckpt_path}')
        self.model.load_state_dict(torch.load(ckpt_path, map_location='cpu')['model_dict'])

        self.model_initializer = InitializationModel(
            t_h=self.cfg.past_frames, d_h=6,
            t_f=self.cfg.future_frames, d_f=2, k_pred=20).cuda()

        params = list(self.model_initializer.parameters())
        self.interaction_graph = None
        if self.use_graph:
            from models.future_interaction_graph_v6 import FutureInteractionGraphV6Wrapper
            self.interaction_graph = FutureInteractionGraphV6Wrapper(
                num_agents=64, future_steps=self.cfg.future_frames,
                past_steps=self.cfg.past_frames, past_channels=6,
                node_dim=128, top_n=5, num_denoise_steps=NUM_Tau).cuda()
            params += list(self.interaction_graph.parameters())

        self.opt = torch.optim.AdamW(params, lr=config.learning_rate)
        self.scheduler = torch.optim.lr_scheduler.StepLR(
            self.opt, step_size=self.cfg.decay_step, gamma=self.cfg.decay_gamma)

        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_param(self.model, 'Core Denoiser')
        self._print_param(self.model_initializer, 'Initializer')
        if self.interaction_graph:
            self._print_param(self.interaction_graph, 'Graph')

        T = self.cfg.future_frames
        self.temporal_reweight = torch.FloatTensor(
            [(T + 1) - i for i in range(1, T + 1)]).cuda().unsqueeze(0).unsqueeze(0) / (T / 2)

    def _print_param(self, m, name):
        t = sum(p.numel() for p in m.parameters())
        tr = sum(p.numel() for p in m.parameters() if p.requires_grad)
        print_log(f'[{name}] {tr}/{t}', self.log)

    def _make_beta_schedule(self, schedule, n, start, end):
        if schedule == 'linear': return torch.linspace(start, end, n)
        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 p_sample_accelerate(self, x, mask, cur_y, t, sigma=None):
        t_tensor = torch.tensor([int(t)]).cuda()
        eps_factor = ((1 - self._extract(self.alphas, t_tensor, cur_y))
                      / self._extract(self.one_minus_alphas_bar_sqrt, t_tensor, cur_y))
        beta = self._extract(self.betas, t_tensor.repeat(x.shape[0]), cur_y)
        eps_theta = self.model.generate_accelerate(cur_y, beta, x, mask)

        if self.interaction_graph is not None:
            abs_t = self._extract(self.alphas_bar_sqrt, t_tensor, cur_y)
            am1_t = self._extract(self.one_minus_alphas_bar_sqrt, t_tensor, cur_y)
            y0_hat = (cur_y - am1_t * eps_theta) / abs_t
            delta = self.interaction_graph(
                y0_hat, x, int(t), sigma=sigma, A_override=x.size(0))
            eps_theta = eps_theta - (abs_t / am1_t) * delta

        mean = (1 / self._extract(self.alphas, t_tensor, cur_y).sqrt()) \
               * (cur_y - eps_factor * eps_theta)
        z = torch.randn_like(cur_y)
        sigma_t = self._extract(self.betas, t_tensor, cur_y).sqrt()
        return mean + sigma_t * z * 0.00001

    def p_sample_loop_accelerate(self, x, mask, loc, sigma=None):
        cur_y = loc[:, :10]
        for i in reversed(range(NUM_Tau)):
            cur_y = self.p_sample_accelerate(x, mask, cur_y, i, sigma=sigma)
        cur_y_ = loc[:, 10:]
        for i in reversed(range(NUM_Tau)):
            cur_y_ = self.p_sample_accelerate(x, mask, cur_y_, i, sigma=sigma)
        return torch.cat((cur_y_, cur_y), dim=1)

    def data_preprocess(self, data):
        pre = data['pre_motion_3D'].cuda()
        fut = data['fut_motion_3D'].cuda()
        A = pre.size(1)
        initial_pos = pre[:, :, -1:]
        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)
        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 fit(self):
        for epoch in range(self.cfg.num_epochs):
            lt, ld, lu = self._train_epoch(epoch)
            print_log(f'[{time.strftime("%Y-%m-%d %H:%M:%S")}] Epoch: {epoch}\t'
                      f'Loss: {lt:.6f}\tDist: {ld:.6f}\tUnc: {lu:.6f}', self.log)
            self.tb.add_scalar('train/loss', lt, epoch)
            self.tb.add_scalar('train/loss_dist', ld, epoch)

            if (epoch + 1) % self.cfg.test_interval == 0:
                perf, n = self._test_epoch()
                # MID protocol: scale normalized ADE/FDE by 50 to report in pixels
                ade_px = perf['ADE'] / n * 50.0
                fde_px = perf['FDE'] / n * 50.0
                print_log(f'Epoch {epoch} Best Of 20: ADE: {ade_px:.4f} FDE: {fde_px:.4f}', self.log)
                self.tb.add_scalar('val/ADE_px', ade_px, epoch)
                self.tb.add_scalar('val/FDE_px', fde_px, epoch)

                cp = {'model_initializer_dict': self.model_initializer.state_dict()}
                if self.interaction_graph:
                    cp['interaction_graph_dict'] = self.interaction_graph.state_dict()
                torch.save(cp, self.cfg.model_path % (epoch + 1))
            self.scheduler.step()
        self.tb.flush(); self.tb.close()

    def _train_epoch(self, epoch):
        self.model.train(); self.model_initializer.train()
        if self.interaction_graph: self.interaction_graph.train()
        lt, ld, lu, cnt = 0, 0, 0, 0
        self.opt.zero_grad()
        for i, data in enumerate(self.train_loader):
            A, mask, past, fut = self.data_preprocess(data)
            sp, me, ve = self.model_initializer(past, mask)
            ve = ve.clamp(min=-5, max=5)
            sp = torch.exp(ve / 2)[..., None, None] * sp \
                 / (sp.std(dim=1).mean(dim=(1, 2))[:, None, None, None] + 1e-6)
            loc = sp + me[:, None]
            sigma_in = ve if self.use_v6_graph else None
            gen = self.p_sample_loop_accelerate(past, mask, loc, sigma=sigma_in)
            loss_d = ((gen - fut.unsqueeze(1)).norm(p=2, dim=-1)
                      * self.temporal_reweight).mean(dim=-1).min(dim=1)[0].mean()
            loss_u = (torch.exp(-ve)
                      * (gen - fut.unsqueeze(1)).norm(p=2, dim=-1).mean(dim=(1, 2))
                      + ve).mean()
            loss = loss_d * 50 + loss_u
            (loss / self.grad_accum).backward()
            if (i + 1) % self.grad_accum == 0:
                params = list(self.model_initializer.parameters())
                if self.interaction_graph: params += list(self.interaction_graph.parameters())
                nn.utils.clip_grad_norm_(params, 1.0)
                self.opt.step(); self.opt.zero_grad()
            lt += loss.item(); ld += loss_d.item() * 50; lu += loss_u.item(); cnt += 1
            self.global_step += 1
        self.opt.step(); self.opt.zero_grad()
        return lt / cnt, ld / cnt, lu / cnt

    def _test_epoch(self):
        """MID-style SDD protocol:
        per-pedestrian full-horizon ADE (mean L2 over 12 future frames) and
        FDE (L2 at final frame), best_of_20 per pedestrian, then average
        across all evaluated pedestrians. Coordinates are already in
        MID's ÷50 mean-centered space, so the final ADE/FDE is multiplied
        by 50 to report in pixels.
        Each scene in the SDD dataloader corresponds to one target
        pedestrian (index 0) + its neighbors; we evaluate only the target
        per scene so each pedestrian is counted exactly once (matches
        MID's get_timesteps_data qualification intent).
        """
        T = self.cfg.future_frames
        perf = {'ADE': 0.0, 'FDE': 0.0}
        n = 0
        np.random.seed(0); random.seed(0)
        torch.manual_seed(0); torch.cuda.manual_seed_all(0)
        self.model_initializer.eval()
        if self.interaction_graph: self.interaction_graph.eval()
        with torch.no_grad():
            for data in self.test_loader:
                A, mask, past, fut = self.data_preprocess(data)
                sp, me, ve = self.model_initializer(past, mask)
                ve = ve.clamp(min=-5, max=5)
                sp = torch.exp(ve / 2)[..., None, None] * sp \
                     / (sp.std(dim=1).mean(dim=(1, 2))[:, None, None, None] + 1e-6)
                loc = sp + me[:, None]
                sigma_in = ve if self.use_v6_graph else None
                pred = self.p_sample_loop_accelerate(past, mask, loc, sigma=sigma_in)
                # MID protocol: only target agent (index 0) per scene
                pred_0 = pred[0:1]              # [1, 20, T, 2]
                fut_0  = fut[0:1]               # [1, T, 2]
                dist = torch.norm(fut_0.unsqueeze(1) - pred_0, dim=-1) * self.traj_scale  # [1, 20, T]
                # best_of_20 per pedestrian, then full-horizon ADE / final FDE
                ade_per_ped = dist.mean(dim=-1).min(dim=-1)[0]       # [1]
                fde_per_ped = dist[:, :, -1].min(dim=-1)[0]           # [1]
                perf['ADE'] += ade_per_ped.sum().item()
                perf['FDE'] += fde_per_ped.sum().item()
                n += 1
        return perf, n