Files
mukesh b8e1a3322f analytical-viz: attention-only vs attn+FFN via two sweep buttons
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>
2026-07-28 13:58:29 -07:00

133 lines
5.3 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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}"
)