b8e1a3322f
Both auto tabs now accept a scope choice at run time:
- "Run sweep — Attention" → include_ffn=False
- "Run sweep — Attn + FFN/MoE" → include_ffn=True
Each button runs an independent sweep and caches its result under its
own session_state key. The most recently clicked button determines the
displayed view; both caches persist so users can flip between the two
scopes without re-running.
Tabs:
- Renamed "Auto Explore" → "Auto Suggest Parallelism"
(accurately reflects that it only varies parallelism knobs; HW is
held at the sidebar values).
- "Auto Hardware" tab unchanged.
- Still 6 top-level tabs; no additional tabs added.
Core changes:
auto_explore.py:
- New include_ffn: bool = True parameter on _sum_visible_latency,
_efficiency, score_config, run_auto_explore, compute_parallelism_
sensitivity. False drops all FFN stages from the summed latency.
auto_hardware.py:
- New include_ffn: bool = True parameter on joint_explore,
compute_sensitivity, _best_parallelism_for_hw, _best_parallelism_
two_stage. Forwards to score_config.
Both defaults keep existing tests byte-identical.
Verified:
- 23 pytest tests pass (added 3 new: attn-only latency lower, attn-only
Pareto non-empty, joint HW attn-only faster than full).
- Smoke: Llama 70B decode 128K:
* Attn+FFN best latency: 12.85 ms (unchanged)
* Attention-only best: 7.35 ms (~57% of full)
* Both sensitivities top-rank bw_hbm_gbs (physics preserved).
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
133 lines
5.3 KiB
Python
133 lines
5.3 KiB
Python
"""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}"
|
||
)
|