MCPcopy Create free account
hub / github.com/AdaptiveMotorControlLab/FMPose3D / train

Function train

scripts/FMPose3D_main.py:162–193  ·  view source on GitHub ↗
(opt, train_loader, model, optimizer)

Source from the content-addressed store, hash-verified

160
161
162def 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
196def print_error(data_type, action_error_sum, is_train):

Callers 1

FMPose3D_main.pyFile · 0.70

Calls 3

AccumLossClass · 0.50
get_variableFunction · 0.50
updateMethod · 0.45

Tested by

no test coverage detected