(opt, train_loader, model, optimizer)
| 160 | |
| 161 | |
| 162 | def train(opt, train_loader, model, optimizer): |
| 163 | loss_all = {"loss": AccumLoss()} |
| 164 | model_3d = model["CFM"] |
| 165 | model_3d.train() |
| 166 | split = "train" |
| 167 | |
| 168 | for i, data in enumerate(tqdm(train_loader, 0)): |
| 169 | batch_cam, gt_3D, input_2D, action, subject, scale, bb_box, cam_ind = data |
| 170 | [input_2D, gt_3D, batch_cam, scale, bb_box] = get_variable( |
| 171 | split, [input_2D, gt_3D, batch_cam, scale, bb_box] |
| 172 | ) |
| 173 | |
| 174 | if split == "train": |
| 175 | B, F, J, C = input_2D.shape |
| 176 | |
| 177 | x0_noise = torch.randn(B, F, J, 3, device=gt_3D.device, dtype=gt_3D.dtype) |
| 178 | x0 = x0_noise |
| 179 | |
| 180 | # t on correct device/dtype and broadcastable: (B,1,1,1) |
| 181 | t = torch.rand(B, 1, 1, 1, device=gt_3D.device, dtype=gt_3D.dtype) |
| 182 | y_t = (1.0 - t) * x0 + t * gt_3D |
| 183 | v_target = gt_3D - x0 |
| 184 | v_pred = model_3d(input_2D, y_t, t) |
| 185 | |
| 186 | loss = ((v_pred - v_target) ** 2).mean() |
| 187 | N = input_2D.size(0) |
| 188 | loss_all["loss"].update(loss.detach().cpu().numpy() * N, N) |
| 189 | optimizer.zero_grad() |
| 190 | loss.backward() |
| 191 | optimizer.step() |
| 192 | |
| 193 | return loss_all["loss"].avg |
| 194 | |
| 195 | |
| 196 | def print_error(data_type, action_error_sum, is_train): |
no test coverage detected