InferScale-Sim / tests /test_execution.py
ArchitSharma's picture
Add online predictive KV tiering experiments
b799d1d
Raw
History Blame Contribute Delete
5.73 kB
from inferscale.execution import (
ExecutionLearningConfig,
OnlineTransitionPredictor,
execution_decay_sweep,
execution_prefetch_study,
execution_threshold_sweep,
generate_workflows,
run_execution_learning,
)
def _cfg() -> dict:
return {
"model": "Qwen2.5-3B",
"accelerator": "L4",
"quantization": "int8",
"seed": 7,
"duration_s": 50,
"workflow_rate_rps": 0.2,
"shift_fraction": 0.5,
"prefix_cache_budget_fraction": 0.62,
"confidence_threshold": 0.5,
}
def test_transition_predictor_updates_without_lookahead() -> None:
predictor = OnlineTransitionPredictor(decay=1.0, prior_strength=0.2)
before, _, _ = predictor.predict("planner")
assert before == "planner"
for _ in range(4):
predictor.observe("planner", "retriever")
after, confidence, probs = predictor.predict("planner")
assert after == "retriever"
assert confidence > probs["reasoner"]
assert predictor.observations == 4
def test_execution_run_reports_prefetch_and_transition_provenance() -> None:
result = run_execution_learning(_cfg() | {"prefetch_policy": "decayed"})
assert result["provenance"]["mode"] == "online-agent-execution-learning"
assert result["provenance"]["lookahead"] == "no-future-transition-lookahead"
assert result["summary"]["steps_completed"] > 0
assert result["prediction"]["count"] > 0
assert 0 <= result["prediction"]["top1_accuracy"] <= 1
assert result["resource"]["prefix_cache_capacity_gb"] > 0
def test_oracle_prefetch_has_perfect_transition_accuracy() -> None:
result = run_execution_learning(_cfg() | {"prefetch_policy": "oracle", "confidence_threshold": 0.9})
assert result["prediction"]["top1_accuracy"] == 1.0
assert result["prediction"]["pre_shift_accuracy"] == 1.0
assert result["prediction"]["post_shift_accuracy"] == 1.0
assert result["provenance"]["lookahead"] == "oracle-upper-bound"
def test_common_trace_policy_study_returns_four_candidates() -> None:
result = execution_prefetch_study(_cfg())
assert result["protocol"] == "common-shifted-agent-workflow-trace"
assert len(result["rows"]) == 4
assert {row["label"] for row in result["rows"]} == {
"Learn only / no prefetch",
"Cumulative transitions",
"Decayed transitions",
"Oracle next-role",
}
def test_threshold_and_decay_sweeps_reuse_one_workflow_generator() -> None:
thresholds = execution_threshold_sweep(_cfg(), [0.0, 0.5, 0.9])
assert [row["threshold"] for row in thresholds["rows"]] == [0.0, 0.5, 0.9]
assert "coverage_vs_ttft_r" in thresholds["association"]
decay = execution_decay_sweep(_cfg(), [0.5, 0.85, 1.0])
assert [row["decay"] for row in decay["rows"]] == [0.5, 0.85, 1.0]
assert decay["best_post_shift_accuracy_decay"] in {0.5, 0.85, 1.0}
assert decay["best_ttft_decay"] in {0.5, 0.85, 1.0}
def test_workflow_generation_has_structured_roles_and_shift() -> None:
cfg = ExecutionLearningConfig.from_dict(_cfg() | {"duration_s": 120, "workflow_rate_rps": 0.5})
workflows = generate_workflows(cfg)
assert workflows
roles = {step.role for workflow in workflows for step in workflow.steps}
assert "planner" in roles
assert len(roles) >= 3
assert any(step.shifted_regime for workflow in workflows for step in workflow.steps)
def test_multistep_forecast_is_normalized_and_no_lookahead() -> None:
predictor = OnlineTransitionPredictor(decay=0.85, prior_strength=0.2)
for _ in range(5):
predictor.observe("planner", "retriever")
forecast = predictor.forecast("planner", horizon=3, discount=0.75)
assert forecast["horizon"] == 3
assert len(forecast["distributions"]) == 3
assert abs(sum(forecast["normalized_scores"].values()) - 1.0) < 1e-9
assert forecast["ranked_roles"][0] == "retriever"
def test_multistep_and_utility_runs_report_forecast_and_calibration_metrics() -> None:
for policy in ("multistep", "utility"):
result = run_execution_learning(
_cfg()
| {
"prefetch_policy": policy,
"forecast_horizon": 3,
"prefetch_top_k": 2,
"forecast_min_score": 0.05,
}
)
assert result["provenance"]["lookahead"] == "no-future-transition-lookahead"
assert 0 <= result["resource"]["forecast_recall"] <= 1
assert 0 <= result["resource"]["prefetch_utilization"] <= 1
assert result["prediction"]["calibration"]["brier"] >= 0
assert 0 <= result["prediction"]["calibration"]["ece"] <= 1
def test_planning_and_horizon_studies_return_controlled_candidates() -> None:
from inferscale.execution import execution_horizon_sweep, execution_planning_study
planning = execution_planning_study(_cfg())
assert len(planning["rows"]) == 4
assert planning["best_ttft_policy"] in {row["label"] for row in planning["rows"]}
assert {row["label"] for row in planning["rows"]} == {
"Top-1 decayed",
"Multi-step top-k",
"Utility-aware multi-step",
"Oracle future-set",
}
horizons = execution_horizon_sweep(_cfg(), [1, 2, 3])
assert [row["horizon"] for row in horizons["rows"]] == [1, 2, 3]
assert horizons["best_ttft_horizon"] in {1, 2, 3}
def test_cache_budget_sweep_compares_three_policies_per_budget() -> None:
from inferscale.execution import execution_budget_sweep
result = execution_budget_sweep(_cfg(), [0.3, 0.6])
assert len(result["rows"]) == 6
assert len(result["winners"]) == 2
assert {row["policy_label"] for row in result["rows"]} == {"Top-1", "Multi-step", "Utility-aware"}