vastai-utils/tests/test_scheduler.py

262 lines
7.6 KiB
Python
Raw Permalink Normal View History

"""Integration tests for the scheduler using MockTransport.
These tests run the actual worker.py as subprocesses in isolated temp dirs,
with real torch training on CPU. No vast.ai instances needed.
"""
from __future__ import annotations
import asyncio
import os
import sys
import pytest
import torch
from sched.aggregator import average_weights
from sched.job import Job
from sched.scheduler import Node, Scheduler
from tests.mock_transport import MockConnection, MockTransport
JOBS_DIR = os.path.join(os.path.dirname(__file__), os.pardir, "jobs")
MNIST_SCRIPT = os.path.abspath(os.path.join(JOBS_DIR, "mnist.py"))
# ------------------------------------------------------------------
# helpers
# ------------------------------------------------------------------
def make_job(**overrides) -> Job:
defaults = dict(
script=MNIST_SCRIPT,
num_nodes=2,
rounds=3,
local_steps=20,
lr=0.01,
batch_size=32,
)
defaults.update(overrides)
return Job(**defaults)
async def setup_scheduler(
job: Job, num_nodes: int = 2
) -> tuple[Scheduler, MockTransport]:
"""Create a scheduler with mock nodes, deploy code. Ready to run rounds."""
transport = MockTransport()
sched = Scheduler(job, transport)
# manually create nodes (skip provision)
for rank in range(num_nodes):
conn = await transport.connect(f"mock{rank}", 22)
sched.nodes.append(
Node(rank=rank, host=f"mock{rank}", port=22, conn=conn, status="ready")
)
# deploy worker + job module
await sched.deploy()
return sched, transport
# ------------------------------------------------------------------
# tests
# ------------------------------------------------------------------
@pytest.mark.asyncio
async def test_single_round_completes():
"""One round with 2 workers. Weights should change from initial."""
job = make_job(rounds=1, local_steps=10)
sched, transport = await setup_scheduler(job)
try:
job_mod = sched._load_job_module()
model = job_mod.make_model()
init_state = {k: v.clone() for k, v in model.state_dict().items()}
new_state = await sched.run_round(model.state_dict(), round_num=0)
# weights should have changed
changed = False
for key in init_state:
if not torch.equal(init_state[key], new_state[key]):
changed = True
break
assert changed, "Weights did not change after training"
finally:
transport.cleanup()
@pytest.mark.asyncio
async def test_multiple_rounds():
"""Multiple rounds. Verify the loop completes and returns valid state dict."""
job = make_job(rounds=3, local_steps=15)
sched, transport = await setup_scheduler(job)
try:
job_mod = sched._load_job_module()
model = job_mod.make_model()
state = model.state_dict()
for r in range(job.rounds):
state = await sched.run_round(state, r)
# final state should be loadable
model.load_state_dict(state)
# quick forward pass shouldn't crash
x = torch.randn(4, 1, 28, 28)
out = model(x)
assert out.shape == (4, 10)
finally:
transport.cleanup()
@pytest.mark.asyncio
async def test_worker_dies_midround():
"""One worker dies before a round. Scheduler continues with survivor."""
job = make_job(rounds=2, local_steps=10)
sched, transport = await setup_scheduler(job)
try:
job_mod = sched._load_job_module()
model = job_mod.make_model()
state = model.state_dict()
# round 0 with both workers
state = await sched.run_round(state, 0)
assert len(sched._active_nodes()) == 2
# kill worker 1
sched.nodes[1].conn.kill()
# round 1 should still work with 1 survivor
state = await sched.run_round(state, 1)
alive = sched._active_nodes()
assert len(alive) == 1
assert alive[0].rank == 0
# state should be valid
model.load_state_dict(state)
finally:
transport.cleanup()
@pytest.mark.asyncio
async def test_all_workers_dead():
"""All workers dead → RuntimeError."""
job = make_job(rounds=1, local_steps=5)
sched, transport = await setup_scheduler(job)
try:
job_mod = sched._load_job_module()
state = job_mod.make_model().state_dict()
for node in sched.nodes:
node.conn.kill()
with pytest.raises(RuntimeError, match="No.*worker"):
await sched.run_round(state, 0)
finally:
transport.cleanup()
@pytest.mark.asyncio
async def test_worker_dies_after_push_before_exec():
"""Worker dies between param push and exec. Round completes with survivor."""
job = make_job(rounds=1, local_steps=10)
sched, transport = await setup_scheduler(job)
try:
job_mod = sched._load_job_module()
state = job_mod.make_model().state_dict()
# wrap node 1's exec to kill it when called
node1_conn = sched.nodes[1].conn
original_exec = node1_conn.exec
async def dying_exec(cmd, timeout=30.0):
node1_conn.kill()
raise ConnectionError("node is dead")
node1_conn.exec = dying_exec
# should still complete with node 0
new_state = await sched.run_round(state, 0)
assert sched.nodes[1].status == "dead"
assert len([k for k, v in new_state.items()]) > 0
finally:
transport.cleanup()
@pytest.mark.asyncio
async def test_three_nodes():
"""FedAvg with 3 workers."""
job = make_job(num_nodes=3, rounds=2, local_steps=10)
sched, transport = await setup_scheduler(job, num_nodes=3)
try:
job_mod = sched._load_job_module()
state = job_mod.make_model().state_dict()
for r in range(job.rounds):
state = await sched.run_round(state, r)
assert len(sched._active_nodes()) == 3
finally:
transport.cleanup()
@pytest.mark.asyncio
async def test_deploy_copies_files():
"""Verify deploy actually puts worker.py and job_module.py on nodes."""
job = make_job()
transport = MockTransport()
sched = Scheduler(job, transport)
try:
for rank in range(2):
conn = await transport.connect(f"mock{rank}", 22)
sched.nodes.append(
Node(rank=rank, host=f"mock{rank}", port=22, conn=conn, status="ready")
)
await sched.deploy()
for node in sched.nodes:
root = node.conn.root_dir
assert os.path.exists(os.path.join(root, "worker.py"))
assert os.path.exists(os.path.join(root, "job_module.py"))
finally:
transport.cleanup()
@pytest.mark.asyncio
async def test_convergence_signal():
"""Run several rounds and check that weight magnitude changes.
Not a rigorous convergence test (synthetic data is random), but verifies
the training loop is actually doing gradient updates.
"""
job = make_job(rounds=5, local_steps=30)
sched, transport = await setup_scheduler(job)
try:
job_mod = sched._load_job_module()
model = job_mod.make_model()
state = model.state_dict()
weight_norms = []
for r in range(job.rounds):
state = await sched.run_round(state, r)
norm = sum(v.float().norm().item() for v in state.values())
weight_norms.append(norm)
# weights should be changing across rounds (not frozen)
assert weight_norms[0] != pytest.approx(weight_norms[-1], abs=1e-3), (
f"Weights did not change across rounds: {weight_norms}"
)
finally:
transport.cleanup()