(model, dataloader, optimizer, scheduler, args)
| 35 | |
| 36 | |
| 37 | def 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 | |
| 60 | def calculate_metrics(all_true_labels, all_predictions, task): |
nothing calls this directly
no outgoing calls
no test coverage detected