MCPcopy Create free account
hub / github.com/microsoft/TRELLIS / run_step

Method run_step

trellis/trainers/basic.py:341–438  ·  view source on GitHub ↗

Run a training step.

(self, data_list)

Source from the content-addressed store, hash-verified

339 print('Done.')
340
341 def run_step(self, data_list):
342 """
343 Run a training step.
344 """
345 step_log = {'loss': {}, 'status': {}}
346 amp_context = partial(torch.autocast, device_type='cuda') if self.fp16_mode == 'amp' else nullcontext
347 elastic_controller_context = self.elastic_controller.record if self.elastic_controller_config is not None else nullcontext
348
349 # Train
350 losses = []
351 statuses = []
352 elastic_controller_logs = []
353 zero_grad(self.model_params)
354 for i, mb_data in enumerate(data_list):
355 ## sync at the end of each batch split
356 sync_contexts = [self.training_models[name].no_sync for name in self.training_models] if i != len(data_list) - 1 and self.world_size > 1 else [nullcontext]
357 with nested_contexts(*sync_contexts), elastic_controller_context():
358 with amp_context():
359 loss, status = self.training_losses(**mb_data)
360 l = loss['loss'] / len(data_list)
361 ## backward
362 if self.fp16_mode == 'amp':
363 self.scaler.scale(l).backward()
364 elif self.fp16_mode == 'inflat_all':
365 scaled_l = l * (2 ** self.log_scale)
366 scaled_l.backward()
367 else:
368 l.backward()
369 ## log
370 losses.append(dict_foreach(loss, lambda x: x.item() if isinstance(x, torch.Tensor) else x))
371 statuses.append(dict_foreach(status, lambda x: x.item() if isinstance(x, torch.Tensor) else x))
372 if self.elastic_controller_config is not None:
373 elastic_controller_logs.append(self.elastic_controller.log())
374 ## gradient clip
375 if self.grad_clip is not None:
376 if self.fp16_mode == 'amp':
377 self.scaler.unscale_(self.optimizer)
378 elif self.fp16_mode == 'inflat_all':
379 model_grads_to_master_grads(self.model_params, self.master_params)
380 self.master_params[0].grad.mul_(1.0 / (2 ** self.log_scale))
381 if isinstance(self.grad_clip, float):
382 grad_norm = torch.nn.utils.clip_grad_norm_(self.master_params, self.grad_clip)
383 else:
384 grad_norm = self.grad_clip(self.master_params)
385 if torch.isfinite(grad_norm):
386 statuses[-1]['grad_norm'] = grad_norm.item()
387 ## step
388 if self.fp16_mode == 'amp':
389 prev_scale = self.scaler.get_scale()
390 self.scaler.step(self.optimizer)
391 self.scaler.update()
392 elif self.fp16_mode == 'inflat_all':
393 prev_scale = 2 ** self.log_scale
394 if not any(not p.grad.isfinite().all() for p in self.model_params):
395 if self.grad_clip is None:
396 model_grads_to_master_grads(self.model_params, self.master_params)
397 self.master_params[0].grad.mul_(1.0 / (2 ** self.log_scale))
398 self.optimizer.step()

Callers

nothing calls this directly

Calls 9

update_emaMethod · 0.95
zero_gradFunction · 0.85
nested_contextsFunction · 0.85
dict_foreachFunction · 0.85
dict_reduceFunction · 0.85
training_lossesMethod · 0.45
logMethod · 0.45

Tested by

no test coverage detected