vastai-utils/tests/test_aggregator.py
Zachery Aaron Shores-Chmielewski 7d11dfa688 feat: Scheduler with basic metrics, docker
Implementation of a mock scheduler for jobs/tasks that will evolve into a job manager for running training jobs over vast.ai instances.
2026-02-05 23:01:05 +07:00

57 lines
1.3 KiB
Python

from collections import OrderedDict
import pytest
import torch
from sched.aggregator import average_weights
def _make_state(val: float) -> OrderedDict:
return OrderedDict(
weight=torch.full((3, 4), val),
bias=torch.full((3,), val),
)
def test_average_identical():
w = _make_state(1.0)
result = average_weights([w, w])
for key in w:
assert torch.allclose(result[key], w[key])
def test_average_different():
a = _make_state(0.0)
b = _make_state(2.0)
result = average_weights([a, b])
expected = _make_state(1.0)
for key in expected:
assert torch.allclose(result[key], expected[key])
def test_average_three():
a = _make_state(0.0)
b = _make_state(3.0)
c = _make_state(6.0)
result = average_weights([a, b, c])
expected = _make_state(3.0)
for key in expected:
assert torch.allclose(result[key], expected[key])
def test_single():
w = _make_state(5.0)
result = average_weights([w])
for key in w:
assert torch.allclose(result[key], w[key])
def test_empty_raises():
with pytest.raises(ValueError):
average_weights([])
def test_preserves_keys():
w = OrderedDict(fc1_weight=torch.ones(2, 2), fc1_bias=torch.zeros(2))
result = average_weights([w, w])
assert list(result.keys()) == ["fc1_weight", "fc1_bias"]