Spaces:
Running
Running
Download tests/test_prediction.py from ArchitSharma/InferScale-Sim: direct link, hf CLI and curl.
- Browser
- Download file 2.8 kB
-
https://huggingface.co/spaces/ArchitSharma/InferScale-Sim/resolve/main/tests/test_prediction.py
- Command line
-
hf download hf://spaces/ArchitSharma/InferScale-Sim/tests/test_prediction.py
-
curl -L -o test_prediction.py https://huggingface.co/spaces/ArchitSharma/InferScale-Sim/resolve/main/tests/test_prediction.py
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 | |