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"}