Files
kernbench2/tests/analytical_visualization/test_auto_explore.py
T
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

208 lines
8.2 KiB
Python

"""Interface + smoke tests for auto_explore.
Covers:
- enumerate_configs yields valid TopologyConfigs with expected pruning
- score_config returns a ConfigScore with all fields populated
- pareto_frontier is a subset of feasible and is non-empty for a
reasonable model+workload
- run_auto_explore returns a coherent AutoExploreResult (all_scores
sorted by latency, pareto ⊆ feasible ⊆ all_scores)
"""
from __future__ import annotations
import pytest
from tests.analytical_visualization.auto_explore import (
ConfigScore,
enumerate_configs,
pareto_frontier,
run_auto_explore,
score_config,
)
from tests.analytical_visualization.model_config import (
FullConfig, MachineParams,
)
from tests.analytical_visualization.model_presets import PRESETS
# ── Enumeration ─────────────────────────────────────────────────────
def test_enumerate_yields_valid_configs():
"""Enumerator produces TopologyConfigs with all 9 knobs set."""
model = PRESETS["Llama 3.1 70B"].model
it = enumerate_configs(model, s_kv=8192, mode="decode")
first = next(it)
for f in ("cp", "tp", "pp", "dp", "kv_shard_mode",
"ffn_shard_scope", "tp_placement", "cp_placement",
"cp_ring_variant"):
assert getattr(first, f) is not None, f"knob {f} not set"
assert first.s_kv == 8192
assert first.mode == "decode"
def test_enumerate_prunes_pp_beyond_layers():
"""PP > model.layers is skipped by the enumerator."""
model = PRESETS["Llama 3.1 70B"].model # layers=80
for cfg in enumerate_configs(model, s_kv=8192, mode="decode"):
assert cfg.pp <= model.layers
def test_enumerate_prunes_ffn_dp_when_dp_is_1():
"""ffn_shard_scope containing 'DP' is skipped when dp=1 (redundant)."""
model = PRESETS["Llama 3.1 70B"].model
for cfg in enumerate_configs(model, s_kv=8192, mode="decode"):
if cfg.dp == 1:
assert "DP" not in cfg.ffn_shard_scope
# ── Scoring ─────────────────────────────────────────────────────────
def test_score_returns_populated_config_score():
"""score_config returns a fully-populated ConfigScore."""
model = PRESETS["Llama 3.1 70B"].model
topo = next(enumerate_configs(model, s_kv=8192, mode="decode"))
machine = MachineParams()
cfg = FullConfig(model=model, topo=topo, machine=machine)
score = score_config(cfg)
assert isinstance(score, ConfigScore)
assert score.total_latency_ns > 0
assert score.pes_used == topo.total_pes
assert 0.0 <= score.efficiency_score <= 1.0
assert score.hbm_utilization >= 0.0
# ── Pareto ──────────────────────────────────────────────────────────
def test_pareto_subset_of_feasible():
"""Pareto scores are always feasible (memory + placement)."""
model = PRESETS["Llama 3.1 70B"].model
machine = MachineParams()
res = run_auto_explore(model, machine, s_kv=8192, mode="decode")
for p in res.pareto_scores:
assert p.fits_memory
assert p.placement_valid
def test_pareto_non_empty_when_feasible_configs_exist():
"""When at least one config fits, Pareto must have ≥ 1 entry."""
model = PRESETS["Llama 3.1 70B"].model
machine = MachineParams()
res = run_auto_explore(model, machine, s_kv=8192, mode="decode")
assert res.total_feasible > 0
assert len(res.pareto_scores) > 0
def test_pareto_frontier_non_dominated():
"""No Pareto entry is dominated by another."""
from tests.analytical_visualization.auto_explore import _dominates
model = PRESETS["Llama 3.1 70B"].model
machine = MachineParams()
res = run_auto_explore(model, machine, s_kv=8192, mode="decode")
for i, a in enumerate(res.pareto_scores):
for j, b in enumerate(res.pareto_scores):
if i == j:
continue
assert not _dominates(b, a), (
f"Pareto entry {i} is dominated by entry {j}"
)
# ── End-to-end sanity ───────────────────────────────────────────────
def test_run_auto_explore_shape():
"""all_scores sorted asc by latency; pareto ⊆ feasible ⊆ all_scores."""
model = PRESETS["Llama 3.1 70B"].model
machine = MachineParams()
res = run_auto_explore(model, machine, s_kv=8192, mode="decode")
assert res.total_enumerated == len(res.all_scores)
assert res.total_feasible == sum(
1 for s in res.all_scores if s.fits_memory and s.placement_valid
)
assert len(res.pareto_scores) <= res.total_feasible
latencies = [s.total_latency_ns for s in res.all_scores]
assert latencies == sorted(latencies)
def test_pp_does_not_reduce_single_request_latency():
"""A single request traverses all layers regardless of PP.
Latency-optimal Pareto configs should NOT prefer PP>1 for decode."""
model = PRESETS["Llama 3.1 70B"].model
machine = MachineParams()
res = run_auto_explore(model, machine, s_kv=8192, mode="decode")
# Sort Pareto by latency; the fastest should not need PP>1 to win.
fastest = res.pareto_scores[0]
assert fastest.pp == 1, (
f"expected PP=1 for the fastest decode config; got PP={fastest.pp}"
)
# ── Parallelism sensitivity ─────────────────────────────────────────
def test_parallelism_sensitivity_has_all_five_knobs():
"""compute_parallelism_sensitivity returns one row per knob."""
from tests.analytical_visualization.auto_explore import (
compute_parallelism_sensitivity,
)
model = PRESETS["Llama 3.1 70B"].model
machine = MachineParams()
res = run_auto_explore(model, machine, s_kv=8192, mode="decode")
baseline = res.pareto_scores[0]
rows = compute_parallelism_sensitivity(
baseline, model, machine, s_kv=8192, mode="decode",
)
knobs = {r.knob for r in rows}
assert knobs == {"cp", "tp", "pp", "dp", "ep"}
def test_attention_only_latency_is_lower_than_full():
"""include_ffn=False must produce lower latency than include_ffn=True
for the same model+workload (FFN cost is dropped)."""
model = PRESETS["Llama 3.1 70B"].model
machine = MachineParams()
r_full = run_auto_explore(model, machine, s_kv=8192, mode="decode",
include_ffn=True)
r_attn = run_auto_explore(model, machine, s_kv=8192, mode="decode",
include_ffn=False)
best_full = min(r_full.pareto_scores, key=lambda s: s.total_latency_ns)
best_attn = min(r_attn.pareto_scores, key=lambda s: s.total_latency_ns)
assert best_attn.total_latency_ns < best_full.total_latency_ns, (
f"attn-only ({best_attn.latency_us:.2f} us) should be less than "
f"full ({best_full.latency_us:.2f} us)"
)
def test_attention_only_pareto_non_empty():
"""The attention-only sweep must still produce a non-empty Pareto set."""
model = PRESETS["Llama 3.1 70B"].model
machine = MachineParams()
res = run_auto_explore(model, machine, s_kv=8192, mode="decode",
include_ffn=False)
assert res.total_feasible > 0
assert len(res.pareto_scores) > 0
def test_parallelism_sensitivity_includes_baseline_value():
"""Each row's values contain the baseline_value."""
from tests.analytical_visualization.auto_explore import (
compute_parallelism_sensitivity,
)
model = PRESETS["Llama 3.1 70B"].model
machine = MachineParams()
res = run_auto_explore(model, machine, s_kv=8192, mode="decode")
baseline = res.pareto_scores[0]
rows = compute_parallelism_sensitivity(
baseline, model, machine, s_kv=8192, mode="decode",
)
for r in rows:
# Baseline value must be sweepable OR mapped to something in values.
assert r.baseline_value in r.values or r.baseline_value == 1, (
f"knob {r.knob}: baseline_value={r.baseline_value} not in {r.values}"
)