InferScale-Sim / tests /test_agentic.py
ArchitSharma's picture
Deepen InferScale simulation research workflow
94910ac
Raw History Blame Contribute Delete
4.42 kB
from inferscale.agentic import compare_agent_policies, run_agent_session_simulation, ttl_retention_sweep
BASE = {
"model": "Qwen2.5-3B",
"accelerator": "L4",
"quantization": "int8",
"duration_s": 30,
"session_rate_rps": 0.20,
"replicas": 2,
"turns_mean": 4,
"tool_gap_mean_s": 1.5,
"seed": 7,
}
def test_agent_session_run_is_deterministic():
first = run_agent_session_simulation(BASE | {"retention_policy": "ttl", "routing_policy": "session_affinity", "kv_ttl_s": 3})
second = run_agent_session_simulation(BASE | {"retention_policy": "ttl", "routing_policy": "session_affinity", "kv_ttl_s": 3})
assert first["summary"] == second["summary"]
assert first["latency"] == second["latency"]
assert first["resource"] == second["resource"]
assert first["provenance"]["mode"] == "stateful-agent-session-simulation"
def test_evict_policy_has_no_cross_turn_hits():
result = run_agent_session_simulation(BASE | {"retention_policy": "evict", "routing_policy": "least_load"})
assert result["resource"]["cross_turn_cache_hit_rate"] == 0
assert result["resource"]["recomputed_history_tokens"] > 0
def test_retention_with_affinity_reuses_state():
baseline = run_agent_session_simulation(BASE | {"retention_policy": "evict", "routing_policy": "least_load"})
retained = run_agent_session_simulation(BASE | {"retention_policy": "retain", "routing_policy": "session_affinity"})
assert retained["resource"]["cross_turn_cache_hit_rate"] > baseline["resource"]["cross_turn_cache_hit_rate"]
assert retained["resource"]["recomputed_history_tokens"] < baseline["resource"]["recomputed_history_tokens"]
assert retained["resource"]["mean_kv_gb"] >= baseline["resource"]["mean_kv_gb"]
def test_policy_compare_uses_four_common_trace_candidates():
result = compare_agent_policies(BASE | {"kv_ttl_s": 3})
assert result["protocol"] == "common-agent-program-trace"
assert result["candidate_count"] == 4
assert {row["label"] for row in result["rows"]} == {
"Stateless / least-load",
"Retain / least-load",
"TTL / affinity",
"Retain / affinity",
}
def test_ttl_sweep_reports_latency_memory_frontier():
result = ttl_retention_sweep(BASE, [0, 1, 3, 8])
assert result["protocol"] == "common-agent-program-trace"
assert result["objective"] == "minimize-p95-turn-ttft-and-mean-kv-residency"
assert result["pareto_count"] >= 1
assert len(result["rows"]) == 4
assert any(row["pareto"] for row in result["rows"])
def test_host_offload_reuses_state_without_hbm_retention():
from inferscale.agentic import compare_agent_memory_policies
result = compare_agent_memory_policies(BASE | {"host_memory_gb": 4, "host_bandwidth_gbps": 32})
rows = {row["label"]: row for row in result["rows"]}
offload = rows["Host offload / bounded affinity"]
stateless = rows["Stateless / least-load"]
assert offload["host_hit_rate"] > 0
assert offload["offloaded_gb"] > 0
assert offload["recomputed_history_tokens"] < stateless["recomputed_history_tokens"]
def test_gap_aware_policy_is_marked_as_oracle_upper_bound():
result = run_agent_session_simulation(
BASE
| {
"retention_policy": "gap_aware",
"routing_policy": "bounded_affinity",
"gap_aware_threshold_s": 1.0,
"host_memory_gb": 4,
}
)
assert result["provenance"]["gap_aware_policy"] == "oracle-upper-bound"
assert result["resource"]["cross_turn_cache_hit_rate"] > 0
def test_memory_budget_sweep_uses_common_trace_and_finite_budgets():
from inferscale.agentic import agent_memory_budget_sweep
result = agent_memory_budget_sweep(BASE, [0.5, 1.0])
assert result["protocol"] == "common-agent-program-trace"
assert result["budget_count"] == 2
assert result["policy_count"] == 3
assert len(result["rows"]) == 6
assert all(row["budget_gb_per_replica"] > 0 for row in result["rows"])
def test_bounded_affinity_sweep_reports_locality_frontier_inputs():
from inferscale.agentic import agent_affinity_sweep
result = agent_affinity_sweep(BASE, [0, 100, 500])
assert result["study"] == "bounded-affinity-routing-frontier"
assert [row["affinity_slack_ms"] for row in result["rows"]] == [0.0, 100.0, 500.0]
assert all(0 <= row["routing_locality_rate"] <= 1 for row in result["rows"])