InferScale-Sim / tests /test_prediction.py
ArchitSharma's picture
Add online predictive KV tiering experiments
d2258e5
Raw History Blame Contribute Delete
2.8 kB
from inferscale.agentic import adaptive_alpha_sweep, adaptive_tiering_study, run_agent_session_simulation
from inferscale.prediction import OnlineToolGapPredictor
def _agent_cfg():
return {
"model": "Qwen2.5-3B",
"accelerator": "L4",
"quantization": "int8",
"duration_s": 24,
"session_rate_rps": 0.3,
"replicas": 2,
"seed": 7,
"retention_policy": "adaptive",
"routing_policy": "bounded_affinity",
"host_memory_gb": 4,
"gap_aware_threshold_s": 1.5,
}
def test_online_predictor_uses_completed_history_only():
predictor = OnlineToolGapPredictor(initial_mean_s=1.5, alpha=0.5, min_observations=2, scope="per_tool_ema")
first, source, count = predictor.predict("search")
assert first == 1.5
assert source == "global"
assert count == 0
predictor.observe("search", 0.4)
second, source, count = predictor.predict("search")
assert source == "global" # tool estimate is still warming up
assert count == 1
predictor.observe("search", 0.6)
third, source, count = predictor.predict("search")
assert source == "tool"
assert count == 2
assert 0.4 <= third <= 0.6
def test_adaptive_agent_run_reports_prediction_provenance():
result = run_agent_session_simulation(_agent_cfg())
assert result["provenance"]["adaptive_policy"] == "online-tool-gap-ewma-no-lookahead"
assert result["resource"]["adaptive_prediction_count"] > 0
assert result["resource"]["adaptive_prediction_mae_s"] >= 0
assert 0 <= result["resource"]["adaptive_oracle_action_agreement"] <= 1
assert result["prediction"]["rows"]
def test_predictive_tiering_study_has_common_shifted_trace():
result = adaptive_tiering_study(
_agent_cfg(), horizon_s=60, shift_fraction=0.5, shift_multiplier=2.0, alpha=0.3
)
assert result["protocol"] == "common-nonstationary-agent-program-trace"
assert len(result["rows"]) == 5
assert result["shift_observation"] > 0
labels = {row["label"] for row in result["rows"]}
assert "Adaptive per-tool EWMA" in labels
assert "Oracle gap-aware" in labels
adaptive = next(row for row in result["rows"] if row["label"] == "Adaptive per-tool EWMA")
assert adaptive["prediction_count"] > 0
assert result["learning_curves"]["per_tool"]
def test_adaptation_rate_sweep_uses_same_trace():
result = adaptive_alpha_sweep(
_agent_cfg(), [0.1, 0.3, 0.8], horizon_s=60, shift_fraction=0.5, shift_multiplier=2.0
)
assert len(result["rows"]) == 3
assert result["best_post_shift_alpha"] in {0.1, 0.3, 0.8}
for row in result["rows"]:
assert row["pre_shift_mae_s"] >= 0
assert row["post_shift_mae_s"] >= 0
assert 0 <= row["oracle_action_agreement"] <= 1