vastai-utils/worker/worker.py

63 lines
1.8 KiB
Python
Raw Normal View History

#!/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()