Implementation of a mock scheduler for jobs/tasks that will evolve into a job manager for running training jobs over vast.ai instances.
62 lines
1.8 KiB
Python
62 lines
1.8 KiB
Python
#!/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()
|