MCPcopy Create free account
hub / github.com/OpenSparseLLMs/Linear-MoE / train_step

Function train_step

linear_moe/training.py:307–379  ·  view source on GitHub ↗

Single training step.

(forward_step_func, data_iterator,
               model, optimizer, opt_param_scheduler, config)

Source from the content-addressed store, hash-verified

305 return model
306
307def train_step(forward_step_func, data_iterator,
308 model, optimizer, opt_param_scheduler, config):
309 """Single training step."""
310 args = get_args()
311 timers = get_timers()
312
313 # Set grad to zero.
314 for partition in model:
315 try:
316 partition.zero_grad_buffer()
317 except:
318 partition.zero_grad_buffer(zero_buffer=(not args.use_distributed_optimizer))
319 optimizer.zero_grad()
320
321 # Forward pass.
322 forward_backward_func = get_forward_backward_func()
323 losses_reduced = forward_backward_func(
324 forward_step_func=forward_step_func,
325 data_iterator=data_iterator,
326 model=model,
327 num_microbatches=get_num_microbatches(),
328 seq_length=args.seq_length,
329 micro_batch_size=args.micro_batch_size,
330 decoder_seq_length=args.decoder_seq_length,
331 forward_only=False)
332
333 # Empty unused memory.
334 if args.empty_unused_memory_level >= 1:
335 torch.cuda.empty_cache()
336
337 # Vision gradients.
338 if args.vision_pretraining and args.vision_pretraining_type == "dino":
339 unwrapped_model = unwrap_model(model[0])
340 unwrapped_model.cancel_gradients_last_layer(args.curr_iteration)
341
342 # Update parameters.
343 timers('optimizer', log_level=1).start(barrier=args.barrier_with_L1_time)
344 update_successful, grad_norm, num_zeros_in_grad = optimizer.step(args, timers)
345 timers('optimizer').stop()
346
347 try:
348 if update_successful:
349 optimizer.gather_model_params(args, timers)
350 except:
351 pass
352
353 # Vision momentum.
354 if args.vision_pretraining and args.vision_pretraining_type == "dino":
355 unwrapped_model = unwrap_model(model[0])
356 unwrapped_model.update_momentum(args.curr_iteration)
357
358 # Update learning rate.
359 if update_successful:
360 increment = get_num_microbatches() * \
361 args.micro_batch_size * \
362 args.data_parallel_size
363 opt_param_scheduler.step(increment=increment)
364 skipped_iter = 0

Callers 2

_trainFunction · 0.90
trainFunction · 0.85

Calls 2

get_argsFunction · 0.90
stepMethod · 0.80

Tested by

no test coverage detected