MCPcopy Create free account
hub / github.com/CERT-Lab/lora-sb / train_client

Function train_client

train_eval.py:37–57  ·  view source on GitHub ↗
(model, dataloader, optimizer, scheduler, args)

Source from the content-addressed store, hash-verified

35
36
37def train_client(model, dataloader, optimizer, scheduler, args):
38
39 scaler = GradScaler()
40 model.train()
41
42 for step, data in enumerate(tqdm(dataloader)):
43 data = {k: v.to(args.device) for k, v in data.items()}
44
45 with autocast():
46 outputs = model(**data)
47 loss = outputs.loss
48
49 wandb.log({"client_loss": loss.detach().cpu().numpy()})
50
51 scaler.scale(loss).backward()
52 scaler.step(optimizer)
53 scaler.update()
54 scheduler.step()
55 optimizer.zero_grad()
56
57 return model.state_dict()
58
59
60def calculate_metrics(all_true_labels, all_predictions, task):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected