MCPcopy Create free account
hub / github.com/Alpha-VLLM/LLaMA2-Accessory / train_one_epoch

Function train_one_epoch

accessory/engine_pretrain.py:14–116  ·  view source on GitHub ↗
(model: torch.nn.Module,
                    data_loader, val_loader, optimizer: torch.optim.Optimizer,
                    epoch: int, start_iter, loss_scaler,
                    log_writer=None,
                    args=None)

Source from the content-addressed store, hash-verified

12
13
14def train_one_epoch(model: torch.nn.Module,
15 data_loader, val_loader, optimizer: torch.optim.Optimizer,
16 epoch: int, start_iter, loss_scaler,
17 log_writer=None,
18 args=None):
19 model.train(True)
20 metric_logger = misc.MetricLogger(delimiter=" ")
21 metric_logger.add_meter('lr', misc.SmoothedValue(window_size=1, fmt='{value:.6f}'))
22 header = 'Epoch: [{}]'.format(epoch)
23 print_freq = 10
24
25 accum_iter = args.accum_iter
26
27 model.zero_grad(set_to_none=True)
28
29 dataset_state = {}
30
31 if log_writer is not None:
32 print('log_dir: {}'.format(log_writer.log_dir))
33 for data_iter_step, (examples, labels, item_states) in enumerate(
34 metric_logger.log_every(data_loader, print_freq, header, start_iter), start=start_iter
35 ):
36
37 if data_iter_step % accum_iter == 0:
38 lr_sched.adjust_learning_rate(optimizer, data_iter_step, args)
39
40 autocast_ctx = {
41 "bf16": torch.cuda.amp.autocast(dtype=torch.bfloat16),
42 "fp16": torch.cuda.amp.autocast(dtype=torch.float16),
43 "tf32": contextlib.nullcontext(),
44 }[args.precision]
45 with autocast_ctx:
46 c_loss, additional_loss_dict = model(examples, labels)
47 loss = c_loss
48 for (add_loss, weight) in additional_loss_dict.values():
49 loss = loss + add_loss * weight
50 loss_value = loss.item()
51 c_loss_value = c_loss.item()
52 if not math.isfinite(loss_value):
53 print("Loss is {}, stopping training".format(loss_value))
54 sys.exit(1)
55
56 loss /= accum_iter
57
58 update_grad = (data_iter_step + 1) % accum_iter == 0
59 grad_norm = loss_scaler(
60 loss, optimizer, model,
61 parameters=model.parameters(),
62 update_grad=update_grad,
63 clip_grad=None if args.clip_grad <= 0 else args.clip_grad,
64 )
65
66 if update_grad:
67 assert grad_norm is not None
68 if torch.any(torch.isinf(grad_norm)):
69 print("grad norm is inf")
70 else:
71 metric_logger.update(grad_norm=grad_norm)

Callers 1

mainFunction · 0.90

Calls 6

add_meterMethod · 0.95
log_everyMethod · 0.95
updateMethod · 0.95
printFunction · 0.85
val_one_epochFunction · 0.85

Tested by

no test coverage detected