"""Interface + invariant tests for auto_hardware.""" from __future__ import annotations from tests.analytical_visualization.auto_hardware import ( HardwareCandidate, JointScore, _HW_KNOB_DEFAULTS, enumerate_hardware, joint_explore, ) from tests.analytical_visualization.model_presets import PRESETS # ── Enumeration ───────────────────────────────────────────────────── def test_enumerate_two_stage_yields_one(): """Two-stage depth yields exactly 1 HW candidate (defaults).""" hws = list(enumerate_hardware("two_stage")) assert len(hws) == 1 for knob, default_val in _HW_KNOB_DEFAULTS.items(): assert getattr(hws[0], knob) == default_val def test_enumerate_balanced_yields_64(): """Balanced = 2 values × 6 knobs = 2^6 = 64 candidates.""" hws = list(enumerate_hardware("balanced")) assert len(hws) == 64 def test_enumerate_coarse_yields_729(): """Coarse = 3 values × 6 knobs = 3^6 = 729 candidates.""" hws = list(enumerate_hardware("coarse")) assert len(hws) == 729 def test_default_cost_score_is_6(): """At every knob = its default, cost_score = 6 (one per knob).""" hw = HardwareCandidate( pe_hbm_gb=_HW_KNOB_DEFAULTS["pe_hbm_gb"], bw_hbm_gbs=_HW_KNOB_DEFAULTS["bw_hbm_gbs"], peak_tflops_f16=_HW_KNOB_DEFAULTS["peak_tflops_f16"], bw_intra_gbs=_HW_KNOB_DEFAULTS["bw_intra_gbs"], bw_inter_gbs=_HW_KNOB_DEFAULTS["bw_inter_gbs"], bw_intersip_gbs=_HW_KNOB_DEFAULTS["bw_intersip_gbs"], ) assert hw.cost_score == 6.0 # ── Joint explore + Pareto ────────────────────────────────────────── def test_joint_explore_returns_pareto_non_empty(): """Given a reasonable model+workload, at least one HW+parallelism fits.""" model = PRESETS["Llama 3.1 70B"].model res = joint_explore(model, s_kv=131072, mode="decode", depth="two_stage") assert res.total_joint >= 1 assert len(res.pareto_scores) >= 1 def test_pareto_subset_of_all_scores(): """Every Pareto entry is present in all_scores.""" model = PRESETS["Llama 3.1 70B"].model res = joint_explore(model, s_kv=131072, mode="decode", depth="two_stage") pareto_ids = {id(s) for s in res.pareto_scores} all_ids = {id(s) for s in res.all_scores} assert pareto_ids.issubset(all_ids) def test_pareto_non_dominated(): """No Pareto entry is dominated by another Pareto entry on (lat, cost).""" model = PRESETS["Llama 3.1 70B"].model res = joint_explore(model, s_kv=131072, mode="decode", depth="balanced") for i, a in enumerate(res.pareto_scores): for j, b in enumerate(res.pareto_scores): if i == j: continue no_worse = ( b.total_latency_ns <= a.total_latency_ns and b.cost_score <= a.cost_score ) strictly_better = ( b.total_latency_ns < a.total_latency_ns or b.cost_score < a.cost_score ) assert not (no_worse and strictly_better), ( f"Pareto entry {i} is dominated by entry {j}" ) # ── Sensitivity ───────────────────────────────────────────────────── def test_sensitivity_all_knobs_monotone_non_worsening(): """Doubling any HW knob should not slow the sim down (rel_speedup ≥ 0 within floating-point tolerance).""" model = PRESETS["Llama 3.1 70B"].model res = joint_explore(model, s_kv=131072, mode="decode", depth="two_stage") assert len(res.sensitivity) == 6 for row in res.sensitivity: # Allow a tiny slop for floating-point but disallow real regressions. assert row.rel_speedup >= -1e-9, ( f"knob {row.knob}: doubling slowed latency from " f"{row.baseline_latency_ns} to {row.doubled_latency_ns}" ) def test_joint_explore_attention_only_faster_than_full(): """include_ffn=False produces smaller best-latency than include_ffn=True for the same HW+model.""" model = PRESETS["Llama 3.1 70B"].model r_full = joint_explore(model, s_kv=131072, mode="decode", depth="two_stage", include_ffn=True) r_attn = joint_explore(model, s_kv=131072, mode="decode", depth="two_stage", include_ffn=False) b_full = min(r_full.pareto_scores, key=lambda s: s.total_latency_ns) b_attn = min(r_attn.pareto_scores, key=lambda s: s.total_latency_ns) assert b_attn.total_latency_ns < b_full.total_latency_ns, ( f"attn-only ({b_attn.latency_ms:.2f}ms) should be less than " f"full ({b_full.latency_ms:.2f}ms)" ) def test_sensitivity_hbm_bw_dominant_for_llama_decode(): """For Llama 70B decode (memory-bound), HBM BW should be the top sensitivity knob — doubling it gives more speedup than any other knob.""" model = PRESETS["Llama 3.1 70B"].model res = joint_explore(model, s_kv=131072, mode="decode", depth="balanced") assert res.sensitivity[0].knob == "bw_hbm_gbs", ( f"expected bw_hbm_gbs to top the sensitivity ranking for Llama 70B " f"decode; got {res.sensitivity[0].knob}" )