Spaces:
Running
Running
File size: 5,725 Bytes
8f91935 b799d1d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 | 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"}
|