diff --git a/tests/analytical_visualization/app.py b/tests/analytical_visualization/app.py index d381266..be8581a 100644 --- a/tests/analytical_visualization/app.py +++ b/tests/analytical_visualization/app.py @@ -62,9 +62,11 @@ from tests.analytical_visualization.chip_roofline import ( good_batch, good_context, knee_batch, + max_batch_within_slo, memory_budget_curve_vs_batch, memory_budget_curve_vs_skv, per_token_latency_curve, + size_deployment, step_latency_curve, t_com, t_mem_long, @@ -2639,3 +2641,147 @@ with tab_roofline: ] st.dataframe(pd.DataFrame(_formula_rows), width='stretch', hide_index=True) + + st.divider() + + # ── Capacity planning: three-axis GPU count ──────────────────── + st.subheader("How hyperscalers pick GPU count") + st.markdown( + "**N_GPUs = max(capacity_floor, KV_headroom, throughput_SLO) × N_replicas.** " + "Three independent axes gate the minimum deployment size." + ) + + _cp_a, _cp_b, _cp_c = st.columns(3) + with _cp_a: + _cp_users = st.number_input( + "Concurrent users", min_value=1, max_value=100_000, + value=100, step=1, key="_cp_n_users", + help="Target simultaneous active decode sequences.", + ) + with _cp_b: + _cp_ctx = st.number_input( + "Avg context / user (tokens)", min_value=1, + max_value=2_000_000, value=int(_rf_skv), step=128, + key="_cp_avg_ctx", + help="Mean S_kv across the active users.", + ) + with _cp_c: + _cp_slo_ms = st.number_input( + "TPOT SLO (ms/token)", min_value=1, max_value=1000, + value=30, step=1, key="_cp_tpot_slo_ms", + help=("Per-token latency target during decode. " + "Tighter = smaller B per replica, more replicas."), + ) + + _sizing = size_deployment( + _scaled_machine, model, + n_users=int(_cp_users), avg_ctx=int(_cp_ctx), + tpot_slo_s=_cp_slo_ms / 1000.0, + ) + + _s1, _s2, _s3, _s4 = st.columns(4) + _s1.metric("Axis A: Capacity", + f"{_sizing.pes_axis_a_capacity} PEs", + help="Min PEs to hold one replica's weights alone. " + "Below this, weights don't fit at any sharding.") + _s2.metric("Axis B: KV headroom", + f"{_sizing.pes_axis_b_kv} PEs", + help="Min PEs to hold weights + all KV of this replica's users. " + "Grows with users × context.") + _s3.metric("Axis C: Throughput", + (f"{_sizing.n_replicas} replicas" + if _sizing.b_at_slo > 0 + else "SLO INFEASIBLE"), + help=(f"B_at_SLO = {_sizing.b_at_slo} per replica; " + f"need ceil({_cp_users}/{_sizing.users_per_replica}) " + f"= {_sizing.n_replicas} replicas.")) + _s4.metric("Total GPUs (recommended)", + f"{_sizing.total_pes}", + delta=f"binds: {_sizing.binding_axis}", + help="max(A, B) × N_replicas.") + + if _sizing.b_at_slo == 0: + st.error( + f"SLO infeasible on this chip: even B=1 step latency exceeds " + f"{_cp_slo_ms} ms at S_kv={_cp_ctx:,}. Increase the SLO, " + f"shorten the context, or move to a faster chip (raise the " + f"FLOPs/BW knobs above)." + ) + else: + _binding_hint = { + "capacity": ("Weights dominate — the model is big relative to " + "per-PE HBM. Adding replicas or bigger HBM chips " + "buys the most headroom."), + "kv": ("KV cache dominates — users × context is filling HBM " + "after weights. Shorter context / smaller B / more " + "aggressive CP sharding help most."), + "throughput": ("Latency SLO forces you into more replicas — " + "one replica can only serve ~" + f"{_sizing.b_at_slo} users at this SLO. " + "Loosening TPOT or picking a faster chip cuts " + "the replica count."), + }[_sizing.binding_axis] + st.caption( + f"**Binding axis: {_sizing.binding_axis}.** {_binding_hint} " + f"At this config: B_at_SLO = **{_sizing.b_at_slo}** per " + f"replica → **{_sizing.users_per_replica}** users/replica × " + f"**{_sizing.n_replicas}** replicas = **{_sizing.total_pes} " + f"GPUs** total." + ) + + st.divider() + + # ── High-level info: same-setup vs different-setup ──────────── + st.subheader("Same setup or different for short vs long context?") + st.markdown( + "**In practice: hybrid — one elastic pool with length-tier " + "routing on top, plus prefill/decode disaggregation at the " + "frontier.**" + ) + + with st.expander("Layer 1 — One base replica config serving mixed traffic", + expanded=False): + st.markdown( + "- **PagedAttention** (vLLM) manages KV as fixed-size pages, " + "so different users' contexts pack efficiently into the same HBM.\n" + "- The scheduler admits requests based on remaining page pool: " + "*'does this new user's expected max_tokens fit?'*\n" + "- **Continuous batching** mixes different-length contexts in " + "the same forward pass.\n" + "- Works well for ~90% of traffic: short-to-medium context " + "(chat, RAG, code snippets)." + ) + + with st.expander("Layer 2 — Length-tier routing for the tail", + expanded=False): + st.markdown( + "Frontier providers (OpenAI, Anthropic, Google) run **separate " + "GPU pools per context tier**:\n" + "- **Standard tier** (up to 32k or 128k): normal replica " + "config, high B, high utilization.\n" + "- **Long-context tier** (128k–1M+): bigger replicas with " + "more chips per instance, more CP sharding to spread KV, " + "lower B per replica.\n\n" + "A router at the API gateway inspects the request's max " + "context and dispatches accordingly. This is why Gemini's " + "1M-context tier is more expensive per token than standard — " + "those requests can't share efficient replicas with short-" + "context traffic." + ) + + with st.expander("Layer 3 — Disaggregated prefill / decode " + "(DistServe, Splitwise)", expanded=False): + st.markdown( + "- **Prefill** (compute-heavy, high FLOPs/byte) runs on one " + "GPU pool.\n" + "- **Decode** (memory-heavy, HBM BW dominated) runs on " + "another.\n" + "- **KV cache is transferred over NVLink/RDMA** between the " + "pools.\n" + "- Long-context prefill goes to specialized 'prefill GPUs' " + "that then hand off KV to a decode-optimized pool.\n\n" + "Rationale: prefill and decode have opposite roofline " + "signatures. Running both on one homogeneous cluster wastes " + "either FLOPs (when decode dominates) or HBM BW (when " + "prefill dominates)." + ) diff --git a/tests/analytical_visualization/chip_roofline.py b/tests/analytical_visualization/chip_roofline.py index b215891..07bfc75 100644 --- a/tests/analytical_visualization/chip_roofline.py +++ b/tests/analytical_visualization/chip_roofline.py @@ -389,6 +389,106 @@ class AISensitivityPoint: b_star: float # C·b/(2·W) +# ── GPU count sizing (three-axis) ───────────────────────────────── + + +@dataclass +class GpuSizingResult: + # Inputs echoed back for display / debugging + n_users: int + avg_ctx_tokens: int + tpot_slo_s: float + # Per-replica values + pes_axis_a_capacity: int # PEs to hold weights alone (bare floor) + pes_axis_b_kv: int # PEs to hold weights + KV of this replica's users + pes_per_replica: int # max(A, B) + # Throughput axis + b_at_slo: int # largest per-replica B satisfying SLO + users_per_replica: int # min(n_users, b_at_slo), at least 1 + n_replicas: int # ceil(n_users / users_per_replica) + # Grand total + total_pes: int + binding_axis: str # 'capacity' | 'kv' | 'throughput' + + +def max_batch_within_slo(machine: MachineParams, model: ModelConfig, + s_kv: int, tpot_slo_s: float) -> int: + """Largest per-replica B such that decode step latency ≤ SLO. + + step_latency(B) = N·b/W + 2·N·B/C + B·S_kv·kv_bpt/W + = weight_fetch + B · (compute + kv_read) + + Returns 0 if even B=1 exceeds SLO (weight fetch alone too large, + or per-sequence term already blows the budget). + """ + n = total_active_params(model) + b_elem = model.bytes_per_elem + weight_time = n * b_elem / machine.bw_hbm + per_seq_time = (_FLOPS_PER_PARAM_PER_TOKEN * n / machine.peak_flops + + s_kv * kv_bytes_per_token(model) / machine.bw_hbm) + if weight_time + per_seq_time > tpot_slo_s: + return 0 + remaining = tpot_slo_s - weight_time + b_max = int(remaining / per_seq_time) + return max(1, b_max) + + +def size_deployment(machine: MachineParams, model: ModelConfig, + n_users: int, avg_ctx: int, + tpot_slo_s: float) -> GpuSizingResult: + """Three-axis GPU count for a target workload. + + Axis A — capacity floor: PEs to hold one replica's weights. + Axis B — KV headroom: PEs to hold weights + all KV of the + users assigned to this replica. + Axis C — throughput SLO: replicas needed so per-replica B ≤ b_at_slo. + + Result: pes_per_replica = max(A, B); total = pes_per_replica × replicas. + Binding axis is whichever grew the count the most. + """ + n = total_active_params(model) + b_elem = model.bytes_per_elem + weight_bytes = n * b_elem + hbm_pe = int(machine.pe_hbm_gb * 1e9) + kv_bpt = kv_bytes_per_token(model) + + # Axis C: throughput + b_at_slo = max_batch_within_slo(machine, model, avg_ctx, tpot_slo_s) + users_per_replica = max(1, min(n_users, b_at_slo)) if b_at_slo > 0 else 1 + n_replicas = (n_users + users_per_replica - 1) // users_per_replica + + # Axis A: bare weights + pes_a = (weight_bytes + hbm_pe - 1) // hbm_pe + + # Axis B: weights + KV load for this replica's users + kv_per_replica = users_per_replica * avg_ctx * kv_bpt + pes_b = (weight_bytes + kv_per_replica + hbm_pe - 1) // hbm_pe + + pes_per_replica = max(pes_a, pes_b) + total = pes_per_replica * n_replicas + + if n_replicas > 1: + binding = "throughput" + elif pes_b > pes_a: + binding = "kv" + else: + binding = "capacity" + + return GpuSizingResult( + n_users=n_users, + avg_ctx_tokens=avg_ctx, + tpot_slo_s=tpot_slo_s, + pes_axis_a_capacity=int(pes_a), + pes_axis_b_kv=int(pes_b), + pes_per_replica=int(pes_per_replica), + b_at_slo=b_at_slo, + users_per_replica=int(users_per_replica), + n_replicas=int(n_replicas), + total_pes=int(total), + binding_axis=binding, + ) + + def ai_sensitivity_curve(machine: MachineParams, model: ModelConfig, multipliers: list[float], axis: str = "flops") -> list[AISensitivityPoint]: diff --git a/tests/analytical_visualization/test_chip_roofline.py b/tests/analytical_visualization/test_chip_roofline.py index 443a0ef..768f1d4 100644 --- a/tests/analytical_visualization/test_chip_roofline.py +++ b/tests/analytical_visualization/test_chip_roofline.py @@ -15,9 +15,11 @@ from tests.analytical_visualization.chip_roofline import ( good_context, knee_batch, kv_bytes_per_token, + max_batch_within_slo, memory_budget_curve_vs_batch, memory_budget_curve_vs_skv, per_token_latency_curve, + size_deployment, step_latency_curve, t_com, t_mem_long, @@ -430,6 +432,79 @@ def test_ai_sensitivity_bad_axis_raises(): [1.0], axis="bogus") +# ── GPU sizing (three-axis) ──────────────────────────────────────── + + +def test_capacity_floor_scales_with_model_size(): + """A model with 2× params needs ~2× PEs for weights alone (axis A).""" + m = MachineParams() + small = PRESETS["Llama 3 8B"].model + big = PRESETS["Llama 3 70B"].model + r_small = size_deployment(m, small, n_users=1, avg_ctx=1024, + tpot_slo_s=1.0) + r_big = size_deployment(m, big, n_users=1, avg_ctx=1024, + tpot_slo_s=1.0) + ratio = r_big.pes_axis_a_capacity / r_small.pes_axis_a_capacity + # 70B / 8B ≈ 8.75× params — allow generous tolerance + assert 5 < ratio < 12, (ratio, r_small, r_big) + + +def test_kv_headroom_scales_with_users_and_context(): + """Doubling n_users OR avg_ctx grows KV load roughly linearly.""" + m = MachineParams() + model = PRESETS["Llama 3 8B"].model + base = size_deployment(m, model, n_users=64, avg_ctx=4096, + tpot_slo_s=1.0) + twice_users = size_deployment(m, model, n_users=128, avg_ctx=4096, + tpot_slo_s=1.0) + twice_ctx = size_deployment(m, model, n_users=64, avg_ctx=8192, + tpot_slo_s=1.0) + # KV bytes doubled either way → axis-B PE count must strictly grow. + assert twice_users.pes_axis_b_kv > base.pes_axis_b_kv + assert twice_ctx.pes_axis_b_kv > base.pes_axis_b_kv + + +def test_binding_axis_flips_with_workload(): + """Huge N_users + short ctx → throughput binds. + Huge ctx + few users → kv (or capacity) binds.""" + m = MachineParams() + model = PRESETS["Llama 3 8B"].model + big_users = size_deployment(m, model, n_users=10_000, avg_ctx=1024, + tpot_slo_s=0.05) + long_ctx = size_deployment(m, model, n_users=1, avg_ctx=500_000, + tpot_slo_s=1.0) + assert big_users.binding_axis == "throughput", big_users + assert long_ctx.binding_axis in ("kv", "capacity"), long_ctx + + +def test_max_batch_shrinks_with_shorter_slo(): + """Tighter SLO → smaller B fits.""" + m = MachineParams() + model = PRESETS["Llama 3 8B"].model + loose = max_batch_within_slo(m, model, s_kv=1024, tpot_slo_s=1.0) + tight = max_batch_within_slo(m, model, s_kv=1024, tpot_slo_s=0.05) + assert tight < loose, (tight, loose) + + +def test_max_batch_returns_zero_when_weight_time_exceeds_slo(): + """If weight-fetch alone > SLO, no batch fits.""" + m = MachineParams() + model = PRESETS["Llama 3 70B"].model + # 70B in BF16 = 140 GB. Reading at 256 GB/s takes ~547 ms. + # SLO of 100 ms cannot be met. + assert max_batch_within_slo(m, model, s_kv=100, + tpot_slo_s=0.1) == 0 + + +def test_total_pes_equals_per_replica_times_replicas(): + """Sanity: total_pes = pes_per_replica × n_replicas.""" + m = MachineParams() + model = PRESETS["Llama 3 8B"].model + r = size_deployment(m, model, n_users=500, avg_ctx=4096, + tpot_slo_s=0.1) + assert r.total_pes == r.pes_per_replica * r.n_replicas + + # ── App wiring: tab exists on the Streamlit app ────────────────────