Files
ModelForge/backend/tests/test_gpu_scheduler.py
T

204 lines
6.8 KiB
Python

from __future__ import annotations
from datetime import UTC, datetime, timedelta
import pytest
from modelforge_api.services.gpu_scheduler import (
AttributionConfidence,
LeaseAllocation,
PlacementInput,
PlacementVerdict,
PressureHysteresis,
PressureState,
ResidentCandidate,
SchedulerPolicy,
calculate_accounting,
dynamic_reserve,
plan_placement,
required_envelope,
)
GIB = 1024**3
POLICY = SchedulerPolicy()
def accounting(
observed: int | None, resident: int = 0, leases: list[LeaseAllocation] | None = None
):
return calculate_accounting(
total_bytes=16 * GIB,
observed_used_bytes=observed,
resident_bytes=resident,
leases=leases or [],
policy=POLICY,
)
@pytest.mark.parametrize(
("observed", "resident", "external"),
[(0, 0, 0), (4 * GIB, 0, 4 * GIB), (2 * GIB, 2 * GIB, 0), (6 * GIB, 2 * GIB, 4 * GIB)],
)
def test_accounting_attributes_observed_usage_without_double_count(
observed: int, resident: int, external: int
) -> None:
result = accounting(observed, resident)
assert result.external_bytes == external
assert result.schedulable_bytes == 16 * GIB - observed - result.reserve_bytes
assert result.invariant_delta_bytes == 0
assert result.attribution is AttributionConfidence.KNOWN
def test_accounting_only_subtracts_unmaterialized_lease_bytes() -> None:
result = accounting(
5 * GIB,
2 * GIB,
[LeaseAllocation(3 * GIB, 2 * GIB), LeaseAllocation(GIB, GIB)],
)
assert result.future_lease_bytes == GIB
assert result.schedulable_bytes == 16 * GIB - 5 * GIB - GIB - result.reserve_bytes
def test_unknown_or_stale_telemetry_fails_closed() -> None:
unknown = accounting(None)
stale = calculate_accounting(
total_bytes=16 * GIB,
observed_used_bytes=0,
resident_bytes=0,
leases=[],
policy=POLICY,
telemetry_fresh=False,
)
assert unknown.attribution is AttributionConfidence.UNKNOWN
assert unknown.schedulable_bytes == stale.schedulable_bytes == 0
assert unknown.pressure is stale.pressure is PressureState.CRITICAL
def test_accounting_is_non_negative_under_noisy_attribution() -> None:
result = accounting(GIB, 2 * GIB)
assert result.external_bytes == 0
assert result.schedulable_bytes >= 0
assert result.attribution is AttributionConfidence.ESTIMATED
def test_reserve_and_envelope_margin_are_central_and_conservative() -> None:
assert dynamic_reserve(16 * GIB, POLICY) >= GIB
assert required_envelope(GIB, POLICY) == GIB + POLICY.deployment_margin_minimum_bytes
assert required_envelope(10 * GIB, POLICY) == 11 * GIB
def request(**updates: object) -> PlacementInput:
values: dict[str, object] = {
"deployment_id": "requested",
"capability": "speech.transcription",
"node_id": "gpu_node",
"accelerator_id": "gpu",
"priority": "interactive",
"required_bytes": 2 * GIB,
"cold_load_ms": 3000.0,
"is_resident": False,
"envelope_stale": False,
"runtime_healthy": True,
"node_eligible": True,
}
values.update(updates)
return PlacementInput(**values) # type: ignore[arg-type]
def candidate(**updates: object) -> ResidentCandidate:
values: dict[str, object] = {
"deployment_id": "lab",
"capability": "vision.embedding",
"resident_bytes": 2 * GIB,
"active_requests": 0,
"priority": "lab",
"policy": "lab_only",
"idle_since": datetime.now(UTC) - timedelta(minutes=5),
"cold_load_ms": 1000.0,
"last_used_at": datetime.now(UTC) - timedelta(minutes=5),
}
values.update(updates)
return ResidentCandidate(**values) # type: ignore[arg-type]
def test_planner_admits_with_measured_headroom() -> None:
decision = plan_placement(request(), accounting(4 * GIB), [], POLICY)
assert decision.verdict is PlacementVerdict.ADMIT
assert decision.headroom_after_bytes >= 0
def test_planner_evicts_only_idle_lower_priority_managed_residency() -> None:
snapshot = accounting(13 * GIB, 2 * GIB)
decision = plan_placement(request(required_bytes=3 * GIB), snapshot, [candidate()], POLICY)
assert decision.verdict is PlacementVerdict.ADMIT_AFTER_EVICTION
assert decision.evictions[0].reason == "IDLE_LAB_EVICTED_FOR_PRODUCTION"
@pytest.mark.parametrize(
"protected", [{"active_requests": 1}, {"pinned": True}, {"policy": "always_warm"}]
)
def test_planner_never_evicts_active_or_pinned_residency(protected: dict[str, object]) -> None:
decision = plan_placement(
request(required_bytes=3 * GIB),
accounting(13 * GIB, 2 * GIB),
[candidate(**protected)],
POLICY,
)
assert not decision.evictions
assert decision.verdict in {PlacementVerdict.QUEUE, PlacementVerdict.REJECT_CAPACITY}
@pytest.mark.parametrize(
("updates", "verdict", "reason"),
[
({"runtime_healthy": False}, PlacementVerdict.REJECT_HEALTH, "RUNTIME_UNHEALTHY"),
({"node_eligible": False}, PlacementVerdict.REJECT_HEALTH, "NODE_UNAVAILABLE"),
({"envelope_stale": True}, PlacementVerdict.REJECT_HEALTH, "SCHEDULER_STATE_STALE"),
(
{"deadline_remaining_ms": 100.0},
PlacementVerdict.REJECT_POLICY,
"DEADLINE_CANNOT_BE_MET",
),
],
)
def test_planner_typed_rejections(
updates: dict[str, object], verdict: PlacementVerdict, reason: str
) -> None:
decision = plan_placement(request(**updates), accounting(4 * GIB), [], POLICY)
assert decision.verdict is verdict
assert reason in decision.reason_codes
def test_lab_pause_is_policy_block_and_fingerprint_is_deterministic() -> None:
policy = SchedulerPolicy(lab_paused=True)
first = plan_placement(request(priority="lab"), accounting(4 * GIB), [], policy)
second = plan_placement(request(priority="lab"), accounting(4 * GIB), [], policy)
assert first.verdict is PlacementVerdict.REJECT_POLICY
assert first.fingerprint == second.fingerprint
def test_hysteresis_requires_sustained_candidate_and_prevents_flapping() -> None:
start = datetime.now(UTC)
hysteresis = PressureHysteresis()
assert hysteresis.observe(PressureState.HIGH, start, 30) is PressureState.NORMAL
assert (
hysteresis.observe(PressureState.NORMAL, start + timedelta(seconds=10), 30)
is PressureState.NORMAL
)
assert (
hysteresis.observe(PressureState.HIGH, start + timedelta(seconds=20), 30)
is PressureState.NORMAL
)
assert (
hysteresis.observe(PressureState.HIGH, start + timedelta(seconds=51), 30)
is PressureState.HIGH
)
assert (
hysteresis.observe(PressureState.NORMAL, start + timedelta(seconds=60), 30)
is PressureState.HIGH
)
assert (
hysteresis.observe(PressureState.NORMAL, start + timedelta(seconds=91), 30)
is PressureState.NORMAL
)