#!/usr/bin/env python3 """Generic FedAvg worker. Deployed to /workspace/ by the scheduler. Loads model params, trains locally for K steps, saves updated weights. The job-specific model and data come from job_module.py (deployed alongside). """ from __future__ import annotations import argparse import os import sys def main(): parser = argparse.ArgumentParser() parser.add_argument("--params", required=True) parser.add_argument("--output", required=True) parser.add_argument("--rank", type=int, required=True) parser.add_argument("--world-size", type=int, required=True) parser.add_argument("--local-steps", type=int, required=True) parser.add_argument("--lr", type=float, default=0.01) parser.add_argument("--batch-size", type=int, default=64) args = parser.parse_args() import torch import torch.nn.functional as F # import job module from same directory sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from job_module import make_dataloader, make_model device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = make_model().to(device) model.load_state_dict( torch.load(args.params, map_location=device, weights_only=True) ) model.train() loader = make_dataloader(args.rank, args.world_size, args.batch_size) opt = torch.optim.SGD(model.parameters(), lr=args.lr) step = 0 while step < args.local_steps: for x, y in loader: if step >= args.local_steps: break x, y = x.to(device), y.to(device) opt.zero_grad() loss = F.cross_entropy(model(x), y) loss.backward() opt.step() step += 1 torch.save(model.state_dict(), args.output) print(f"rank={args.rank} done steps={args.local_steps}") if __name__ == "__main__": main()