diff --git a/docs/docs.json b/docs/docs.json index ae1b38dd1584..c5fd8759ad44 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -149,10 +149,6 @@ "source": "/advanced_features/pd_disaggregation.html", "destination": "/docs/advanced_features/pd_disaggregation" }, - { - "source": "/advanced_features/piecewise_cuda_graph.html", - "destination": "/docs/advanced_features/piecewise_cuda_graph" - }, { "source": "/advanced_features/pipeline_parallelism.html", "destination": "/docs/advanced_features/pipeline_parallelism" @@ -964,7 +960,6 @@ "docs/advanced_features/dp_for_multi_modal_encoder", "docs/advanced_features/cuda_graph_for_multi_modal_encoder", "docs/advanced_features/breakable_cuda_graph", - "docs/advanced_features/piecewise_cuda_graph", "docs/advanced_features/sgl_model_gateway", "docs/advanced_features/llm-d", "docs/advanced_features/deterministic_inference", diff --git a/docs/docs/advanced_features/breakable_cuda_graph.mdx b/docs/docs/advanced_features/breakable_cuda_graph.mdx index 546af64844ab..5bc0589e8e9c 100644 --- a/docs/docs/advanced_features/breakable_cuda_graph.mdx +++ b/docs/docs/advanced_features/breakable_cuda_graph.mdx @@ -35,7 +35,7 @@ For production use, you can mark specific functions as "non-graphable" using the ```python from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import eager_on_graph -@eager_on_graph(enable=True) +@eager_on_graph def my_dynamic_op(x): # This op is incompatible with CUDA graph capture return some_dynamic_operation(x) diff --git a/docs/docs/advanced_features/cuda_graph_for_multi_modal_encoder.mdx b/docs/docs/advanced_features/cuda_graph_for_multi_modal_encoder.mdx index b1c1aed28998..3476d83e9972 100644 --- a/docs/docs/advanced_features/cuda_graph_for_multi_modal_encoder.mdx +++ b/docs/docs/advanced_features/cuda_graph_for_multi_modal_encoder.mdx @@ -69,7 +69,7 @@ it; naming a backend locks the choice and skips that rule: SGLANG_VIT_ENABLE_CUDA_GRAPH=1 \ python3 -m sglang.launch_server \ --model Qwen/Qwen3-VL-8B-Instruct \ - --cuda-graph-backend-prefill tc_piecewise \ + --cuda-graph-backend-prefill breakable \ --cuda-graph-max-bs-prefill 4096 \ --cuda-graph-tc-compiler eager ``` diff --git a/docs/docs/advanced_features/piecewise_cuda_graph.mdx b/docs/docs/advanced_features/piecewise_cuda_graph.mdx deleted file mode 100644 index 4f03665e824c..000000000000 --- a/docs/docs/advanced_features/piecewise_cuda_graph.mdx +++ /dev/null @@ -1,289 +0,0 @@ ---- -title: "Piecewise CUDA Graph" -metatags: - description: "Use Piecewise CUDA Graph to reduce prefill and extend kernel launch overhead while supporting dynamic token shapes." ---- - -## Motivation - -Standard CUDA graphs capture the entire model forward pass as a single graph. This works well for decode (fixed batch size), but not for extend/prefill where the number of tokens varies across iterations. - -Piecewise CUDA Graph (PCG) solves this by splitting the model's computation graph into pieces (roughly one per layer) at "split points" (e.g., MoE dispatch ops). Each piece is captured as a separate CUDA graph for a set of pre-defined token lengths. At runtime, the input is padded to the nearest captured size, and each piece is replayed. This eliminates kernel launch overhead for prefill/extend while still supporting dynamic shapes. - -PCG is **enabled by default**. Pass `--cuda-graph-backend-prefill=disabled` to turn it off. - -## Usage - -PCG is enabled by default for supported configurations. No extra flags needed: - -```bash -python3 -m sglang.launch_server \ - --model-path meta-llama/Llama-3.1-8B-Instruct -``` - -### Disable PCG - -```bash -python3 -m sglang.launch_server \ - --model-path meta-llama/Llama-3.1-8B-Instruct \ - --cuda-graph-backend-prefill=disabled -``` - -### Custom capture sizes - -```bash -python3 -m sglang.launch_server \ - --model-path meta-llama/Llama-3.1-8B-Instruct \ - --cuda-graph-max-bs-prefill 2048 -``` - -### Server Args - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
ArgumentDefaultDescription
--cuda-graph-backend-prefillNone (auto)Backend for the prefill phase. Choices: full, breakable, tc_piecewise, disabled. Pass disabled to turn PCG off for extend/prefill, or tc_piecewise to force it on, skipping all auto-disable conditions (testing only).
--cuda-graph-max-bs-prefillNone (auto)Maximum token count to capture. Defaults to chunked_prefill_size (non-MLA) or 2048 (MLA).
--cuda-graph-bs-prefillNone (auto)Explicit list of token lengths to capture. Auto-generated if not set.
--cuda-graph-tc-compiler"eager"Compiler backend for the captured subgraphs. Choices: eager, inductor.
- -## Bug Report - -PCG is enabled by default but is still in an experimental stage. Since PCG relies on `torch.compile` to trace the model's forward pass, most bugs are introduced by torch compile tracing failures (e.g., untraceable ops, dynamic control flow, or graph breaks). If you encounter any issues related to PCG, please disable it by adding `--cuda-graph-backend-prefill=disabled` to your launch command and report the bug at [GitHub Issues](https://github.com/sgl-project/sglang/issues/new/choose). We greatly appreciate your help in improving this feature. - -### For Users - -If you see an error message like the following during server startup, it is a PCG bug: - -``` -Piecewise CUDA Graph is enabled by default as an experimental feature. -To work around this error, add --cuda-graph-backend-prefill=disabled to your launch command. -Please report this issue at https://github.com/sgl-project/sglang/issues/new/choose -``` - -To work around it, add `--cuda-graph-backend-prefill=disabled` to your launch command. When filing a bug report, please include: -1. The full error traceback -2. Model name and quantization method -3. Launch command with all arguments -4. GPU type and driver version - -### For Developers - -Since PCG relies on `torch.compile` to trace the model's forward pass, newly developed CUDA kernels (both JIT kernels and sgl-kernels) are typically not compatible with `torch.compile` out of the box. The tracing will fail on untraceable operations such as JIT compilation, file I/O, or dynamic module loading inside the kernel. - -To make a kernel compatible with PCG, you need to register it as a custom op using `register_custom_op` from `sglang.srt.utils.custom_op`. This wraps the kernel as an opaque node in the compiled graph so that `torch.compile` will not trace inside it. - -**Example usage (JIT kernel):** - -```python -from sglang.srt.utils.custom_op import register_custom_op - -# Inplace operator (no return value) -@register_custom_op(mutates_args=["output_q", "output_s"]) -def per_token_group_quant_8bit( - input: torch.Tensor, - output_q: torch.Tensor, - output_s: torch.Tensor, -) -> None: - # kernel implementation ... -``` - -**Example usage (operator with output):** - -```python -# out_shape indicates which argument has the same shape as the output -@register_custom_op(mutates_args=["x"], out_shape=0) -def add(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: - return x.add_(y) -``` - -For wrapping external library functions (e.g., FlashInfer kernels), use `register_custom_op_from_extern` instead. See `python/sglang/srt/utils/custom_op.py` for full API documentation. - -## How it works - -### Torch compile backend - -PCG uses `torch.compile` with a custom backend (`SGLangBackend`) to split and compile the model's forward pass. The flow is: - -``` -model.forward wrapper -→ torch.compile(..., backend=SGLangBackend) -→ FX graph -→ split_graph() at registered split ops -→ split_gm (top-level graph that chains the pieces) -→ replace capturable submodules with CUDAPiecewiseBackend -→ runtime dispatch: eager split ops + per-piece capture/replay -``` - -- **Install**: `install_torch_compiled()` replaces `model.forward` with a wrapper function. When `is_in_piecewise_cuda_graph()` returns True, the wrapper dispatches to the compiled callable; otherwise it falls back to the original forward. The first invocation through this path triggers Dynamo tracing and graph compilation — CUDA graph replay only happens after the capture phase completes. - -- **Split**: When `torch.compile` traces the model, `SGLangBackend` receives the FX graph and calls `split_graph()`. Ops listed in `CompilationConfig.split_ops` are treated as split points, so the graph is cut at each one. These split-op submodules are left to run eagerly at runtime, while the surrounding submodules are compiled and wrapped by `CUDAPiecewiseBackend`. The result is a top-level "stitching graph" (`split_gm`) with children such as `submod_0`, `submod_1`, … interleaving capturable subgraphs and eager split-op submodules. - -- **Replace**: `PiecewiseCompileInterpreter` iterates over each capturable submodule in `split_gm`, compiles it for general (dynamic) shapes, and replaces it in-place with a `CUDAPiecewiseBackend` instance. Split-op submodules (e.g., attention, all-reduce) are left as-is and run eagerly at runtime. - -- **Dispatch**: At runtime, calling `split_gm` executes the stitching graph, which calls each submodule in order. Split-op submodules run eagerly. Each `CUDAPiecewiseBackend` submodule goes through three phases: - - **Compile warmup** — runs the general-shape compiled path. - - **Capture** — for each capture size, runs one warmup pass then records a CUDA graph. - - **Steady-state replay** — replays the captured CUDA graph for each forward pass. - -### Piecewise cuda graph runner - -`PiecewiseCudaGraphRunner` orchestrates the full lifecycle through three phases: - -- **Compile** — Warms up JIT kernels with a dummy forward pass, then wraps the model with `torch.compile`, triggering Dynamo tracing to split the FX graph and create `CUDAPiecewiseBackend` instances for each subgraph piece. - -- **Capture** — Iterates over capture sizes in reverse order (largest first). For each size, runs the forward pass twice (one warmup, one CUDA graph capture). - -- **Replay** — At runtime, finds the smallest captured size >= actual token count via binary search, copies inputs into static buffers with zero-padding, replays the captured CUDA graphs, and slices outputs back to the actual token count. - -### Memory optimization - -The memory cost of PCG comes from two parts: **torch memory allocator** and **non-torch memory**. - -The torch memory allocator overhead is trivial thanks to several optimizations: a global shared memory pool is reused across all CUDA graph runners and capture sizes, capture is done in reverse order (large to small) so smaller graphs reuse memory allocated by larger ones, and output tensors of the last subgraph are stored as weak references to maximize memory reuse. - -The main memory overhead comes from non-torch memory — the CUDA graph objects themselves require GPU memory to store the recorded kernel launch parameters and internal state. This overhead scales with the number of captured sizes, which is why `piecewise_cuda_graph_max_tokens` is capped conservatively by default. - -### Shape configuration - -Piecewise CUDA graph pre-captures graphs for a set of token counts. At runtime, the actual token count is rounded up to the nearest captured size (via binary search), and the corresponding graph is replayed. If the token count exceeds the largest captured size, the runtime falls back to the normal (non-graph) forward path. - -The default capture schedule is auto-generated with increasing granularity: - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
Token rangeStep size
4 – 324
48 – 25616
288 – 51232
576 – 102464
1280 – 4096256
4096+512
- -For the auto-generated schedule, sizes are capped at `--cuda-graph-max-bs-prefill`. The default cap is `chunked_prefill_size` for non-MLA models and `2048` for MLA backend models. If `--max-total-tokens` is set, the cap is further limited to not exceed it. Additionally, Llama-2 models are auto-capped at 4096 tokens as a temporary workaround. - -## Compatibility - -PCG is auto-disabled in the following scenarios. We are actively working on expanding compatibility — support for many of these will be coming soon. - -- Disabled model architectures (e.g., `DeepseekV32ForCausalLM`) -- Speculative decoding -- DP attention -- Pipeline parallelism (`pp_size > 1`) -- Non-CUDA hardware (AMD ROCm, Ascend NPU) -- MoE A2A backend -- LoRA -- Multimodal / VLM models -- DLLM (diffusion LLM) -- Deterministic inference -- PD disaggregation -- Expert distribution recorder / EPLB - -Use `--cuda-graph-backend-prefill=tc_piecewise` to skip all auto-disable checks (for testing/debugging only). - -## Code Reference - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
FileDescription
python/sglang/srt/model_executor/runner_backend/tc_piecewise_cuda_graph_backend.pyBackend implementation: compile, capture, replay
python/sglang/srt/compilation/compile.pyinstall_torch_compiled trampoline
python/sglang/srt/compilation/backend.pySGLangBackend, graph splitting, piecewise compilation
python/sglang/srt/compilation/cuda_piecewise_backend.pyPer-subgraph CUDA graph capture/replay
python/sglang/srt/compilation/piecewise_context_manager.pyGlobal context flags and ForwardContext
python/sglang/srt/compilation/compilation_config.pyCapture sizes, split ops, compiler config
python/sglang/srt/utils/custom_op.pyregister_custom_op for torch.compile compatibility
python/sglang/srt/server_args.pyServer arguments and auto-disable logic
diff --git a/docs/docs/advanced_features/server_arguments.mdx b/docs/docs/advanced_features/server_arguments.mdx index 15f0543c4b7c..dc83be97f94f 100644 --- a/docs/docs/advanced_features/server_arguments.mdx +++ b/docs/docs/advanced_features/server_arguments.mdx @@ -2408,7 +2408,7 @@ Combining `--enable-response-store` with `--disaggregation-mode=prefill` or `dec `--cuda-graph-config` - Canonical per-phase CUDA graph settings as JSON, e.g. {`{"decode":{"backend":"full","max_bs":256},"prefill":{"backend":"tc_piecewise","tc_compiler":"eager"}}`}. JSON wins over the per-phase --cuda-graph-* convenience flags and over the legacy flags. Allowed backends: full, breakable, tc_piecewise, disabled (full is decode-only). + Canonical per-phase CUDA graph settings as JSON, e.g. {`{"decode":{"backend":"full","max_bs":256},"prefill":{"backend":"breakable"}}`}. JSON wins over the per-phase --cuda-graph-* convenience flags and over the legacy flags. Allowed backends: full, breakable, disabled. `None` Type: JSON (dict-of-dicts) @@ -2416,13 +2416,13 @@ Combining `--enable-response-store` with `--disaggregation-mode=prefill` or `dec `--cuda-graph-backend-decode` Backend for the decode phase. Folds into cuda_graph_config[decode].backend. `None` - full, breakable, tc_piecewise, disabled + full, breakable, disabled `--cuda-graph-backend-prefill` Backend for the prefill phase. Folds into cuda_graph_config[prefill].backend. `None` - breakable, tc_piecewise, disabled + breakable, disabled `--cuda-graph-max-bs-decode` @@ -2454,12 +2454,6 @@ Combining `--enable-response-store` with `--disaggregation-mode=prefill` or `dec `None` List[int] - - `--cuda-graph-tc-compiler` - Compiler used by the tc_piecewise backend (only the prefill phase consumes it today). - `None` - eager, inductor - `--disable-cuda-graph-padding` Disable cuda graph when padding is needed. Still uses cuda graph when padding is not needed. diff --git a/docs/docs/hardware-platforms/ascend-npus/model-deployment/best-practices/mimo_v2_5_pro.mdx b/docs/docs/hardware-platforms/ascend-npus/model-deployment/best-practices/mimo_v2_5_pro.mdx index c18f549457b3..a7fce758e3fc 100644 --- a/docs/docs/hardware-platforms/ascend-npus/model-deployment/best-practices/mimo_v2_5_pro.mdx +++ b/docs/docs/hardware-platforms/ascend-npus/model-deployment/best-practices/mimo_v2_5_pro.mdx @@ -81,7 +81,7 @@ python3 -m sglang.launch_server \ --swa-full-tokens-ratio 0.3 \ --disaggregation-mode prefill --disaggregation-transfer-backend ascend \ --disaggregation-bootstrap-port 8996 \ - --disable-piecewise-cuda-graph \ + --cuda-graph-backend-prefill disabled \ --dp-size 2 --enable-dp-attention --enable-dp-lm-head \ --moe-a2a-backend deepep --deepep-mode normal ``` diff --git a/docs/docs/hardware-platforms/ascend-npus/reference/support_features.mdx b/docs/docs/hardware-platforms/ascend-npus/reference/support_features.mdx index d66366ba923c..47fd3193d2dd 100644 --- a/docs/docs/hardware-platforms/ascend-npus/reference/support_features.mdx +++ b/docs/docs/hardware-platforms/ascend-npus/reference/support_features.mdx @@ -2061,7 +2061,7 @@ If the value is int8, you must also set the environment variable:DEEP_NORMAL_MOD `--cuda-graph-backend-prefill` `None` - `disabled`, `tc_piecewise`
(`tc_piecewise` currently supports Llama-3.1-8B-Instruct and Qwen2.5-7B-Instruct) + `disabled`
(the former `tc_piecewise` prefill backend has been removed) A2/A3 Series diff --git a/docs/docs/hardware-platforms/plugin.mdx b/docs/docs/hardware-platforms/plugin.mdx index 2e4f78c411a3..d822bbc5b1f4 100644 --- a/docs/docs/hardware-platforms/plugin.mdx +++ b/docs/docs/hardware-platforms/plugin.mdx @@ -556,11 +556,6 @@ python -c "from sglang.srt.platforms import current_platform; print(current_plat False Whether device graph capture is supported (plain CUDA graph) - - support_piecewise_cuda_graph() - False - Whether piecewise CUDA graph (torch.compile backend) is supported - supports_fp8() False @@ -620,11 +615,6 @@ python -c "from sglang.srt.platforms import current_platform; print(current_plat raise NotImplementedError Return hardware-specific quantization config for the specific quantization scheme, raise an error if not supported or return None to use the default config. - - get_piecewise_backend_cls() - raise NotImplementedError - Piecewise compilation backend class - get_compile_backend(mode) "inductor" diff --git a/docs/docs/hardware-platforms/xpu.mdx b/docs/docs/hardware-platforms/xpu.mdx index 336e22a5ace8..f34d3ecb4ecb 100644 --- a/docs/docs/hardware-platforms/xpu.mdx +++ b/docs/docs/hardware-platforms/xpu.mdx @@ -177,7 +177,6 @@ SGLang enables XPU graph capture to reduce per-step kernel-launch overhead. | Phase | Backend | Mechanism | Default | |---|---|---|---| | Decode | `full` | One `torch.xpu.XPUGraph` per batch size, captured on startup | **Off** (opt-in) | -| Prefill | `tc_piecewise` | `torch.compile` + XPU graph, one graph segment per token-length bucket | **Off** (opt-in) | | Prefill | `breakable` | Segmented `torch.xpu.XPUGraph` capture/replay (no `torch.compile`); eager break points at attention / MoE boundaries | **Off** (opt-in) | ### Enable Decode Graph @@ -192,26 +191,7 @@ python -m sglang.launch_server --model-path --device xpu \ ### Enable Prefill Graph Prefill graph capture is **opt-in** on XPU and must be enabled explicitly. -Two backends are available: `tc_piecewise` and `breakable`. - -#### tc_piecewise - -Uses `torch.compile` plus an XPU graph, one graph segment per token-length -bucket: - -```bash -python -m sglang.launch_server --model-path --device xpu \ - --cuda-graph-backend-prefill tc_piecewise -``` - -By default the prefill subgraphs are compiled with `eager` mode. Switch to -`inductor` for higher-quality generated code at the cost of longer startup: - -```bash -python -m sglang.launch_server --model-path --device xpu \ - --cuda-graph-backend-prefill tc_piecewise \ - --cuda-graph-tc-compiler inductor -``` +Use the `breakable` backend. The former `tc_piecewise` backend has been removed. #### breakable @@ -227,7 +207,7 @@ You can also configure both phases together with a single `--cuda-graph-config` ```bash python -m sglang.launch_server --model-path --device xpu \ - --cuda-graph-config '{"decode":{"backend":"full"},"prefill":{"backend":"tc_piecewise","tc_compiler":"eager"}}' + --cuda-graph-config '{"decode":{"backend":"full"},"prefill":{"backend":"breakable"}}' ``` ### Enable torch.compile for Decode @@ -242,11 +222,6 @@ python -m sglang.launch_server --model-path --device xpu \ --enable-torch-compile ``` -> **Note:** `--enable-torch-compile` is mutually exclusive with the prefill -> `tc_piecewise` graph (the compatibility rules auto-disable it). Use them -> separately or lock the prefill backend explicitly via `--cuda-graph-config` -> if you need both. - ### Disable XPU Graph Both phases are disabled by default. To explicitly disable them anyway: @@ -274,7 +249,7 @@ To specify explicit token-length buckets: ```bash python -m sglang.launch_server \ --model-path --device xpu \ - --cuda-graph-backend-prefill tc_piecewise \ + --cuda-graph-backend-prefill breakable \ --cuda-graph-bs-prefill 64 128 256 512 ``` @@ -291,11 +266,10 @@ python -m sglang.launch_server \ | Argument | XPU allowed values | Default | Description | |---|---|---|---| | `--cuda-graph-backend-decode` | `full`, `disabled` | `disabled` | Backend for the decode phase. Only `full` is supported on XPU. Set to `full` to enable. | -| `--cuda-graph-backend-prefill` | `tc_piecewise`, `breakable`, `disabled` | `disabled`* | Backend for the prefill phase. Set to `tc_piecewise` or `breakable` explicitly to enable. | -| `--cuda-graph-tc-compiler` | `eager`, `inductor` | `eager` | Compiler for `tc_piecewise` prefill subgraphs. `inductor` produces more optimized code but has longer startup. | +| `--cuda-graph-backend-prefill` | `breakable`, `disabled` | `disabled`* | Backend for the prefill phase. Set to `breakable` explicitly to enable. | | `--cuda-graph-bs-prefill` | list of ints | auto | Explicit token-length buckets to capture for prefill. | | `--cuda-graph-bs-decode` | list of ints | auto | Explicit batch sizes to capture for decode. | -| `--cuda-graph-config` | JSON string | — | One-shot JSON config for both phases, e.g. `'{"decode":{"backend":"full"},"prefill":{"backend":"tc_piecewise","tc_compiler":"eager"}}'`. Overrides all per-phase flags. | +| `--cuda-graph-config` | JSON string | — | One-shot JSON config for both phases, e.g. `'{"decode":{"backend":"full"},"prefill":{"backend":"breakable"}}'`. Overrides all per-phase flags. | | `--disable-decode-cuda-graph` | — | `False` | Shorthand for `--cuda-graph-backend-decode=disabled`. | | `--disable-prefill-cuda-graph` | — | `False` | Shorthand for `--cuda-graph-backend-prefill=disabled`. | | `--enable-torch-compile` | — | `False` | Apply `torch.compile` on top of the decode XPU graph for further kernel optimization. | diff --git a/python/sglang/kernels/ops/attention/fla/layernorm_gated.py b/python/sglang/kernels/ops/attention/fla/layernorm_gated.py index c0de08c4585a..e0e8ea31ddd7 100644 --- a/python/sglang/kernels/ops/attention/fla/layernorm_gated.py +++ b/python/sglang/kernels/ops/attention/fla/layernorm_gated.py @@ -17,11 +17,6 @@ from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.batch_invariant_ops import is_batch_invariant_mode_enabled -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.utils import ( cdiv, cpu_has_amx_support, @@ -211,9 +206,7 @@ def _get_sm_count(device: torch.device) -> int: def calc_rows_per_block(M: int, device: torch.device) -> int: # Use a constant value when the row count must not affect kernel numerics. - if is_batch_invariant_mode_enabled() or check_cuda_graph_backend( - Phase.PREFILL, Backend.TC_PIECEWISE - ): + if is_batch_invariant_mode_enabled(): return MAX_ROWS_PER_BLOCK sm_count = _get_sm_count(device) rows_per_block = next_power_of_2(cdiv(M, 2 * sm_count)) diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index c1374fec86b9..992b9aed54cf 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -2101,16 +2101,15 @@ def _make_breakable_attention_forward(forward_method): disabled this is a transparent pass-through to the original method. """ - def _forward_boxing_tuples(*args, **kwargs): + @eager_on_graph + def _eager_attention(*args, **kwargs): out = forward_method(*args, **kwargs) return _BCGBoxedTupleOutput(out) if isinstance(out, tuple) else out - bcg_forward = eager_on_graph(True)(_forward_boxing_tuples) - @functools.wraps(forward_method) def forward(self, *args, **kwargs): if is_in_breakable_cuda_graph(): - out = bcg_forward(self, *args, **kwargs) + out = _eager_attention(self, *args, **kwargs) return out.astuple() if isinstance(out, _BCGBoxedTupleOutput) else out return forward_method(self, *args, **kwargs) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py index 1c83cecab743..6eec8c0d9430 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py @@ -775,7 +775,9 @@ def _minimax_h3_attention_core_impl( return out -_minimax_h3_attention_core_bcg = eager_on_graph(True)(_minimax_h3_attention_core_impl) +@eager_on_graph +def _eager_attention_core(*args, **kwargs): + return _minimax_h3_attention_core_impl(*args, **kwargs) class MiniMaxH3Attention(nn.Module): @@ -1173,7 +1175,7 @@ def forward( gate_compress = gate_flat.view(total, self.num_heads, self.head_dim) attention_core = ( - _minimax_h3_attention_core_bcg + _eager_attention_core if self.bcg_breakpoint else _minimax_h3_attention_core_impl ) @@ -2395,8 +2397,8 @@ def build_rope_cache( self.release_mps_non_layer_weights("rope") return result - @eager_on_graph(True) - def _embed( + @eager_on_graph + def _eager_embed( self, *, x: torch.Tensor, @@ -2428,7 +2430,7 @@ def _embed( elif torch.is_tensor(refined_prompt_embeds_length): # BCG turns this request-varying host constant into a scalar input # so different live lengths can replay one padded-text signature. - # _embed is an eager graph break, so this value is read outside + # _eager_embed is an eager graph break, so this value is read outside # captured CUDA graphs. text_len = int(refined_prompt_embeds_length.item()) else: @@ -2692,7 +2694,7 @@ def forward(self, **kwargs: Any) -> tuple[torch.Tensor, torch.Tensor]: audio_pos = audio_pos.to(device) text_pos = text_pos.to(device) - decoder_input, t_emb = self._embed( + decoder_input, t_emb = self._eager_embed( x=x, audio_x=audio_x, text_embeddings_selected=text_selected, diff --git a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3_vdn_attention.py b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3_vdn_attention.py index 3413084fff66..96771a58df8d 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3_vdn_attention.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3_vdn_attention.py @@ -140,7 +140,7 @@ def forward( beta = self.linear_attention.beta(x) gate_hidden, _ = self.linear_attention.output_gate.down(x) attention_core = ( - _hybrid_attention_core_bcg + _eager_hybrid_attention_core if attention.bcg_breakpoint else _minimax_h3_hybrid_attention_core_impl ) @@ -563,9 +563,9 @@ def _vdn_return_to_rows( return merged[0], linear_rows -_hybrid_attention_core_bcg = eager_on_graph(True)( - _minimax_h3_hybrid_attention_core_impl -) +@eager_on_graph +def _eager_hybrid_attention_core(*args, **kwargs): + return _minimax_h3_hybrid_attention_core_impl(*args, **kwargs) def prepare_hybrid_attention_metadata( diff --git a/python/sglang/srt/arg_groups/cuda_graph_hook.py b/python/sglang/srt/arg_groups/cuda_graph_hook.py index 1aeee2839b60..81e40fadca19 100644 --- a/python/sglang/srt/arg_groups/cuda_graph_hook.py +++ b/python/sglang/srt/arg_groups/cuda_graph_hook.py @@ -22,14 +22,10 @@ default_cuda_graph_config, with_phase, ) -from sglang.srt.platforms import current_platform from sglang.srt.runtime_context import get_platform from sglang.srt.utils.common import ( - is_cpu, - is_mps, parse_connector_type, ) -from sglang.srt.utils.hf_transformers_utils import check_gguf_file logger = logging.getLogger(__name__) @@ -84,11 +80,6 @@ def _set(phase: str, key: str, value: Any) -> None: _set(Phase.DECODE, "bs", cfg.cuda_graph_bs_decode) if cfg.cuda_graph_bs_prefill is not None: _set(Phase.PREFILL, "bs", cfg.cuda_graph_bs_prefill) - if cfg.cuda_graph_tc_compiler is not None: - # Written to both phases so the value is in place when TC_PIECEWISE - # decode is implemented; today decode ignores it. - _set(Phase.DECODE, "tc_compiler", cfg.cuda_graph_tc_compiler) - _set(Phase.PREFILL, "tc_compiler", cfg.cuda_graph_tc_compiler) if cfg.cuda_graph_prefill_max_context is not None: _set( Phase.PREFILL, @@ -113,10 +104,8 @@ def _set(phase: str, key: str, value: Any) -> None: def apply_cuda_graph_compatibility(server_args: Any): """Auto-disable prefill cuda graph for incompatible configs. - Rules are split per backend — TcPiecewise and Breakable have - different constraints. Skipped when the user explicitly set the - prefill backend, whichever value they chose (the contract the removed - --enforce-piecewise-cuda-graph used to spell). + Breakable and Full have different constraints. Skip automatic selection + when the user explicitly sets the prefill backend. """ cfg = resolving_view(server_args) @@ -145,133 +134,13 @@ def apply_cuda_graph_compatibility(server_args: Any): # piecewise-allowlisted archs run their validated decoder prefill # there instead. Archs also on the breakable allowlist keep it -- # this runs first, so piecewise would otherwise silently win. - if ( - cfg.cuda_graph_config.prefill.backend == Backend.BREAKABLE - and model_config_of(server_args).is_multimodal_piecewise_cuda_graph_supported - and not model_config_of( - server_args - ).is_multimodal_breakable_cuda_graph_supported - # Keep trtllm_mla on the preferred breakable path, which now serves - # MLA by falling back to the flashinfer MLA impl for extend. - and attention_backends_of(resolved_view(server_args))[0] != "trtllm_mla" - ): - logger.info( - "Using tc_piecewise CUDA graph for validated multimodal decoder prefill." - ) - declare_resolution( - server_args, - "_apply_cuda_graph_compatibility", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.TC_PIECEWISE - ), - ) - if cfg.cuda_graph_config.prefill.backend == Backend.TC_PIECEWISE: - disable_tc_piecewise_cudagraph_if_incompatible(server_args) - elif cfg.cuda_graph_config.prefill.backend == Backend.BREAKABLE: + if cfg.cuda_graph_config.prefill.backend == Backend.BREAKABLE: disable_breakable_cudagraph_if_incompatible(server_args) elif cfg.cuda_graph_config.prefill.backend == Backend.FULL: disable_full_prefill_cudagraph_if_incompatible(server_args) -def disable_tc_piecewise_cudagraph_if_incompatible(server_args: Any): - """TcPiecewise (torch.compile + piecewise) is incompatible with - these configurations. Most are torch.compile / dynamo limitations. - """ - - cfg = resolving_view(server_args) - - rules = [ - ( - "model-arch blacklist", - lambda: model_config_of(server_args).is_piecewise_cuda_graph_disabled_model, - ), - ("DP attention", lambda: resolved_view(server_args).enable_dp_attention), - ("full torch.compile mode", lambda: cfg.enable_torch_compile), - ("pipeline parallelism (pp_size > 1)", lambda: cfg.pp_size > 1), - ( - "non-CUDA hardware (HIP/NPU/CPU/MPS/XPU)", - lambda: ( - get_platform().is_hip - or get_platform().is_npu - or is_cpu() - or is_mps() - or get_platform().is_xpu - ), - ), - ( - "OOT platform without piecewise support", - lambda: ( - current_platform.is_out_of_tree() - and not current_platform.support_piecewise_cuda_graph() - ), - ), - ( - "MoE A2A backend", - lambda: resolved_view(server_args).moe_a2a_backend != "none", - ), - # Dynamo blocks LoRA under tc_piecewise (per-batch LoRABatchInfo - # rebinds break guards); breakable/full support LoRA. - ("LoRA", lambda: bool(cfg.lora_paths) or cfg.enable_lora), - ( - "multimodal model", - lambda: ( - model_config_of(server_args).is_multimodal - and not model_config_of( - server_args - ).is_multimodal_piecewise_cuda_graph_supported - ), - ), - ( - "GGUF quantization", - lambda: ( - cfg.load_format == "gguf" - or resolved_view(server_args).quantization == "gguf" - or check_gguf_file(cfg.model_path) - ), - ), - ("DLLM (diffusion LLM)", lambda: cfg.dllm_algorithm is not None), - ( - "CPU offload / hierarchical cache", - lambda: cfg.cpu_offload_gb > 0 or cfg.enable_hierarchical_cache, - ), - ( - "deterministic inference", - lambda: cfg.enable_deterministic_inference, - ), - ("PD disaggregation", lambda: cfg.disaggregation_mode != "null"), - ("symmetric memory", lambda: cfg.enable_symm_mem), - ( - "expert distribution recorder", - lambda: ( - cfg.enable_eplb or cfg.expert_distribution_recorder_mode is not None - ), - ), - ( - "context parallel (attn_cp_size > 1)", - lambda: resolved_view(server_args).attn_cp_size > 1, - ), - ("CUDA graph debug mode", lambda: cfg.debug_cuda_graph), - # Capture builds a dummy extend forward with attn_dcp_metadata=None. - ( - "decode context parallel (dcp_size > 1)", - lambda: cfg.dcp_size > 1, - ), - ] - for _name, predicate in rules: - if predicate(): - declare_resolution( - server_args, - "_disable_tc_piecewise_cudagraph_if_incompatible", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED - ), - ) - # One decision, one declaration: every rule declares the same - # value, so a later match would only append a duplicate entry. - break - - def disable_breakable_cudagraph_if_incompatible(server_args: Any): """Breakable (segmented capture, no torch.compile). Breakable enforces memory-saver rejection in its own __init__; config-time rules can be @@ -374,8 +243,7 @@ def disable_full_prefill_cudagraph_if_incompatible(server_args: Any): def disable_prefill_cuda_graph_for_deepseek_trtllm_mla(server_args: Any): """Disable prefill CUDA graph for dsr1 by default when using the trtllm_mla - attention backend. Under any captured prefill CUDA graph (tc_piecewise or - breakable) trtllm_mla falls back to FlashAttention for prefill and regresses + attention backend. Under any captured prefill CUDA graph trtllm_mla falls back to FlashAttention for prefill and regresses performance, so disable whichever prefill graph backend is in effect. """ @@ -397,7 +265,7 @@ def disable_prefill_cuda_graph_for_deepseek_trtllm_mla(server_args: Any): "Disabling prefill CUDA graph (%s) by default for the DeepSeek-V3 arch on " "the trtllm_mla attention backend (a captured prefill graph forces a " "FlashAttention fallback that regresses prefill). Set the prefill cuda graph " - "backend explicitly (e.g. --cuda-graph-backend-prefill tc_piecewise) to override.", + "backend explicitly (e.g. --cuda-graph-backend-prefill breakable) to override.", cfg.cuda_graph_config.prefill.backend, ) declare_resolution( @@ -545,7 +413,7 @@ def handle_cuda_graph_config(server_args: Any): if cfg.cuda_graph_config.prefill.backend == Backend.FULL: logger.warning( "cuda_graph_config[prefill].backend='full' is experimental. " - "Use breakable or tc_piecewise for production workloads." + "Use breakable for production workloads." ) @@ -622,7 +490,7 @@ def finalize_cuda_graph_prefill_max_context(server_args: Any) -> None: def generate_prefill_cuda_graph_batch_sizes(max_bs: int): """ Generate the list of batch sizes for prefill CUDA graph capture - based on max_bs. For tc_piecewise prefill, bs carries the + based on max_bs. For prefill, bs carries the captured token count (one shape knob per phase). """ capture_sizes = ( diff --git a/python/sglang/srt/arg_groups/field_order.py b/python/sglang/srt/arg_groups/field_order.py index 20c35da68093..13de911e4cf6 100644 --- a/python/sglang/srt/arg_groups/field_order.py +++ b/python/sglang/srt/arg_groups/field_order.py @@ -226,7 +226,6 @@ "cuda_graph_bs_decode", "cuda_graph_bs_prefill", "cuda_graph_prefill_max_context", - "cuda_graph_tc_compiler", "disable_prefill_cuda_graph", "disable_decode_cuda_graph", "disable_cuda_graph", diff --git a/python/sglang/srt/arg_groups/fields/exec_.py b/python/sglang/srt/arg_groups/fields/exec_.py index 21afe7f02b3c..261a0e43058f 100644 --- a/python/sglang/srt/arg_groups/fields/exec_.py +++ b/python/sglang/srt/arg_groups/fields/exec_.py @@ -479,19 +479,19 @@ class ExecGraph(msgspec.Struct): cuda_graph_config: A[ Optional[CudaGraphConfig], Arg( - help='Per-phase CUDA graph settings as JSON, e.g. \'{"decode":{"backend":"full","max_bs":256},"prefill":{"backend":"tc_piecewise","tc_compiler":"eager"}}\'. Allowed backends per phase: full, breakable, tc_piecewise, disabled (full is decode-only). JSON wins over the per-phase --cuda-graph-* convenience flags and over legacy flags.', + help='Per-phase CUDA graph settings as JSON, e.g. \'{"decode":{"backend":"full","max_bs":256},"prefill":{"backend":"breakable"}}\'. Allowed backends per phase: full, breakable, disabled. JSON wins over the per-phase --cuda-graph-* convenience flags and over legacy flags.', type_parser=parse_cuda_graph_config_arg, ), ] = None cuda_graph_backend_decode: A[ - Optional[Literal["full", "breakable", "tc_piecewise", "disabled"]], + Optional[Literal["full", "breakable", "disabled"]], Arg( help="Backend for the decode phase. Folds into cuda_graph_config[decode].backend.", choices=Backend.ALL, ), ] = None cuda_graph_backend_prefill: A[ - Optional[Literal["full", "breakable", "tc_piecewise", "disabled"]], + Optional[Literal["full", "breakable", "disabled"]], Arg( help="Backend for the prefill phase. Folds into cuda_graph_config[prefill].backend.", choices=Backend.ALL, @@ -530,10 +530,6 @@ class ExecGraph(msgspec.Struct): aliases=["--context-bucket"], ), ] = None - cuda_graph_tc_compiler: A[ - Optional[Literal["eager", "inductor"]], - "Compiler used by the tc_piecewise backend (currently only the prefill phase consumes it).", - ] = None disable_prefill_cuda_graph: A[ bool, "Disable the prefill-phase CUDA graph. Convenience for --cuda-graph-backend-prefill=disabled.", diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index a2de0bc2937b..d640b7487b24 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -1765,8 +1765,6 @@ def cutedsl_moe_max_num_tokens(server_args: Any) -> int: num_tokens_per_req = 1 prefill_tokens = cfg.max_prefill_tokens cg_config = cfg.cuda_graph_config - if cg_config is not None and cg_config.prefill.backend == Backend.TC_PIECEWISE: - prefill_tokens = max(prefill_tokens, cg_config.prefill.max_bs or 0) decode_max_bs = (cg_config.decode.max_bs if cg_config is not None else 0) or 0 decode_tokens = decode_max_bs * num_tokens_per_req return max(prefill_tokens, decode_tokens) diff --git a/python/sglang/srt/arg_groups/platform_hook.py b/python/sglang/srt/arg_groups/platform_hook.py index 7fdffb6c2aef..3b51e9268b99 100644 --- a/python/sglang/srt/arg_groups/platform_hook.py +++ b/python/sglang/srt/arg_groups/platform_hook.py @@ -37,20 +37,6 @@ def handle_npu_backends(server_args: Any): set_default_server_args(server_args) - current = cfg.cuda_graph_config.prefill.tc_compiler - if current is not None and current != "eager": - logger.warning( - "At this moment Ascend platform only support prefill graph compilation with " - "cuda_graph_config[prefill].tc_compiler='eager'." - ) - declare_resolution( - server_args, - "_handle_npu_backends", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.PREFILL, tc_compiler="eager" - ), - ) - def handle_mps_backends(server_args: Any): cfg = resolving_view(server_args) diff --git a/python/sglang/srt/compilation/torch_compile_decoration.py b/python/sglang/srt/compilation/torch_compile_decoration.py index c7c305ce5a6e..644185cb2ec4 100644 --- a/python/sglang/srt/compilation/torch_compile_decoration.py +++ b/python/sglang/srt/compilation/torch_compile_decoration.py @@ -6,11 +6,6 @@ otherwise. ``set_torch_compile_config`` flips the inductor/dynamo config flags expected by that path. -Note: the prefill-tc_piecewise path (``TcPiecewiseCudaGraphBackend``) does NOT -use ``patch_model`` — it goes through ``compilation/compile.py``'s -``install_torch_compiled``. ``_to_torch`` here is duplicated by -tc_piecewise's local ``_toggle_fused_ops``; the duplication is kept -because the two paths have different lifecycle requirements. """ from __future__ import annotations diff --git a/python/sglang/srt/distributed/device_communicators/custom_all_reduce.py b/python/sglang/srt/distributed/device_communicators/custom_all_reduce.py index 55c292cde30c..cf94b033f59a 100644 --- a/python/sglang/srt/distributed/device_communicators/custom_all_reduce.py +++ b/python/sglang/srt/distributed/device_communicators/custom_all_reduce.py @@ -19,9 +19,6 @@ is_weak_contiguous, ) from sglang.srt.environ import envs -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.utils import ( get_bool_env_var, is_cuda, @@ -318,13 +315,7 @@ def custom_all_reduce(self, input: torch.Tensor) -> Optional[torch.Tensor]: # Could be warmup OR piecewise cuda graph split op execution. # In piecewise cuda graph, split ops run eagerly outside the graph # but _IS_CAPTURING is still True. We need to do real all-reduce. - if is_in_tc_piecewise_cuda_graph(): - # Split op execution - do real all-reduce - return self._all_reduce_impl(input, registered=False) - else: - # True warmup - mimic the allocation pattern since custom - # allreduce is out-of-place. - return torch.zeros_like(input) + return torch.zeros_like(input) else: return self._all_reduce_impl(input, registered=False) diff --git a/python/sglang/srt/distributed/device_communicators/custom_all_reduce_v2.py b/python/sglang/srt/distributed/device_communicators/custom_all_reduce_v2.py index 15b0bbadaedd..8fb3f560c145 100644 --- a/python/sglang/srt/distributed/device_communicators/custom_all_reduce_v2.py +++ b/python/sglang/srt/distributed/device_communicators/custom_all_reduce_v2.py @@ -42,9 +42,6 @@ ) from sglang.srt.distributed.parallel_state import in_the_same_node_as from sglang.srt.environ import envs -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.utils.cuda_vmm_utils import ( VmmGraphInputManager, compute_graph_capture_bases, @@ -326,11 +323,7 @@ def _can_use_graph(self) -> bool: # hot path never reaches the cudart capture query. During capture, # warm-up runs execute immediately and must not consume a # graph_params row (it would be dereferenced before registration). - return ( - self._graph_mode_allowed - and not is_in_tc_piecewise_cuda_graph() - and torch.cuda.is_current_stream_capturing() - ) + return (self._graph_mode_allowed) and (torch.cuda.is_current_stream_capturing()) def _pick_config(self, nbytes: int, can_use_graph: bool) -> AllReduceConfig | None: # TODO: refactor this along with the config file diff --git a/python/sglang/srt/distributed/device_communicators/pymscclpp.py b/python/sglang/srt/distributed/device_communicators/pymscclpp.py index 2442c9fc5b15..cd82d2c734e8 100644 --- a/python/sglang/srt/distributed/device_communicators/pymscclpp.py +++ b/python/sglang/srt/distributed/device_communicators/pymscclpp.py @@ -10,13 +10,6 @@ import torch.distributed as dist from torch.distributed import ProcessGroup, ReduceOp -from sglang.srt.compilation.compile_phase import ( - get_pcg_capture_stream, - is_in_torch_compile_warmup, -) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.runtime_context import get_exec logger = logging.getLogger(__name__) @@ -1151,12 +1144,6 @@ def should_mscclpp_allreduce( # mscclpp must not be used during any piecewise CUDA graph phase # (compile, capture, or replay) as it changes the allreduce dispatch # path and triggers recompilation. - if ( - is_in_tc_piecewise_cuda_graph() - or is_in_torch_compile_warmup() - or get_pcg_capture_stream() is not None - ): - return False return True def should_mscclpp_allgather( @@ -1181,12 +1168,6 @@ def should_mscclpp_allgather( config = self._get_allgather_tuned_config(output_nbytes) if config is None or not config.supports_dtype(input_tensor.dtype): return False - if ( - is_in_tc_piecewise_cuda_graph() - or is_in_torch_compile_warmup() - or get_pcg_capture_stream() is not None - ): - return False return True def dtype_to_mscclpp_dtype(self, dtype: torch.dtype): diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 7676ef3676c9..a13d0f502aab 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -45,16 +45,12 @@ from torch.distributed import Backend, ProcessGroup from sglang.srt import platforms -from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.distributed.utils import ( all_gather_single, reduce_scatter_single, set_global_tcp_store, ) from sglang.srt.environ import envs -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.platforms.device_mixin import _DEVICE_TO_DISTRIBUTED_BACKEND from sglang.srt.runtime_context import ( derive_parallel_widths, @@ -176,7 +172,6 @@ def _register_group(group: "GroupCoordinator") -> None: @register_custom_op(mutates_args=["tensor"]) -@register_split_op() def inplace_all_reduce(tensor: torch.Tensor, group_name: str) -> None: assert group_name in _groups, f"Group {group_name} is not found." group = _groups[group_name]() @@ -826,21 +821,6 @@ def fused_allreduce_rmsnorm( total_bytes = input_.numel() * input_.element_size() use_1stage_ar = total_bytes <= 128 * 1024 - if ( - getattr(ca_comm, "_IS_CAPTURING", False) - and not torch.cuda.is_current_stream_capturing() - and is_in_tc_piecewise_cuda_graph() - ): - if not hasattr(ca_comm, "fused_ar_rms"): - return None - return ca_comm.fused_ar_rms( - input_, - residual_inp_, - w=weight_, - eps=eps, - registered=False, - use_1stage=use_1stage_ar, - ) fused_outputs = ca_comm.custom_fused_ar_rms( input_, residual_inp_, @@ -958,9 +938,6 @@ def _resolve_outplace_all_reduce_method( and self.torch_symm_mem_comm.should_torch_symm_mem_allreduce(input_) ): return "torch_symm_mem" - if is_in_tc_piecewise_cuda_graph() and self.pynccl_comm is not None: - # For piecewise cuda graph, we use pynccl outplace allreduce - return "pynccl" return None def _can_use_flashinfer_allreduce(self, input_: torch.Tensor) -> bool: @@ -1160,11 +1137,7 @@ def _maybe_aiter_reduce_scatter( ca_comm.reduce_scatter(input, output, registered=False) else: ca_comm.reduce_scatter(input, output, registered=True) - elif is_in_tc_piecewise_cuda_graph(): - ca_comm.reduce_scatter(input, output, registered=False) - else: - # True CUDA graph warmup: avoid a different host collective. - output.zero_() + output.zero_() return True ca_comm.reduce_scatter(input, output, registered=False) return True @@ -1279,11 +1252,7 @@ def _all_gather_into_tensor(self, output: torch.Tensor, input: torch.Tensor): ca_comm.all_gather_unreg(input, out=output, dim=0) else: ca_comm.all_gather_reg(input, out=output, dim=0) - elif is_in_tc_piecewise_cuda_graph(): - ca_comm.all_gather_unreg(input, out=output, dim=0) - else: - # True CUDA graph warmup: avoid a different host collective. - output.zero_() + output.zero_() return else: ca_comm.all_gather_unreg(input, out=output, dim=0) diff --git a/python/sglang/srt/kv_canary/api.py b/python/sglang/srt/kv_canary/api.py index da9032f7b19d..e27c7cac6321 100644 --- a/python/sglang/srt/kv_canary/api.py +++ b/python/sglang/srt/kv_canary/api.py @@ -59,12 +59,6 @@ def install_canary( if config.mode is CanaryMode.NONE: return None - assert not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE), ( - "kv-canary: piecewise cuda graph is not supported by the current " - "SingleForwardManager design; set --cuda-graph-backend-prefill=disabled " - "(or =breakable) when canary is enabled" - ) - perturb_config = PerturbConfig.from_env() device = torch.device(model_runner.device) if torch_reference_conflicts_with_decode_graph(device): diff --git a/python/sglang/srt/layers/activation.py b/python/sglang/srt/layers/activation.py index 7c1bf82fed1f..aea0dd9155ef 100644 --- a/python/sglang/srt/layers/activation.py +++ b/python/sglang/srt/layers/activation.py @@ -28,11 +28,6 @@ ) from sglang.srt.environ import envs from sglang.srt.layers.quantization.base_config import QuantizationConfig -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.runtime_context import get_exec, get_parallel, publish_role from sglang.srt.utils import ( cpu_has_amx_support, @@ -175,8 +170,6 @@ def forward_xpu(self, x: torch.Tensor) -> torch.Tensor: return out def forward_musa(self, x: torch.Tensor) -> torch.Tensor: - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE): - return self.forward_native(x) if not hasattr(self, "_musa_swish_glu"): # XXX (MUSA): nn.SwishGLU seems to have better performance than silu_and_mul on MUSA, we can switch to it for now. We can consider implementing a silu_and_mul kernel for MUSA in the future if needed. diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index ce8d926134ec..4de3cfac33b4 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -107,6 +107,9 @@ from sglang.srt.mem_cache.deepseek_v4_compress_state import KVAndScore from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode +from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( + eager_on_graph, +) from sglang.srt.runtime_context import ( get_exec, get_parallel, @@ -248,58 +251,6 @@ def _maybe_precompute_flashmla_sched_meta( flashmla_metadata.num_splits = num_splits -def _low_ratio_source_projections(layer, x, q_lora, positions, bufs): - from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, - ) - - real = ( - get_tc_piecewise_forward_context().forward_batch.global_num_token_non_padded_cpu - ) - if real is None: - real = x.shape[0] - - # These GEMMs pick their algorithm by M, so at the bucket size the live rows - # differ from eager; everything downstream is row-independent. - def put(name, value): - buf = bufs[name] - buf[:real].copy_(value) - buf[real:].zero_() - - if real == 0: - # An idle DP-attention rank replays on fabricated rows with no live - # token; a zero-row GEMM is a launch error, so only zero the buffers. - for buf in bufs.values(): - buf.zero_() - return - - if layer.compressor is not None: - kv, score = layer.compressor.project(x[:real]) - put("kv", kv) - if score is not None: - put("score", score) - if layer.indexer is not None: - indexer = layer.indexer - put("q", indexer.queries(q_lora[:real], layer.freqs_cis[positions[:real]])) - put("w", indexer.head_weights(x[:real])) - - -def _bcg_low_ratio_source_projections(*args): - from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.breakable_cuda_graph import ( - eager_on_graph, - ) - - global _bcg_low_ratio_source_projections_fn - if _bcg_low_ratio_source_projections_fn is None: - _bcg_low_ratio_source_projections_fn = eager_on_graph(True)( - _low_ratio_source_projections - ) - return _bcg_low_ratio_source_projections_fn(*args) - - -_bcg_low_ratio_source_projections_fn = None - - def _as_int_list(values) -> Optional[List[int]]: if values is None: return None @@ -2735,7 +2686,9 @@ def forward_low_ratio_sources( and self._low_ratio_in_prefill_graph() ): bufs = self._source_projection_buffers(x.shape[0], layer.compress_ratio) - _bcg_low_ratio_source_projections(layer, x, q_lora, pos, bufs) + self._eager_low_ratio_source_projections( + layer, x, q_lora, pos, bufs, forward_batch + ) if run_compressor and layer.compressor is not None: self._low_ratio_compress_torch( layer, x, req, pos, projected=(bufs["kv"], bufs.get("score")) @@ -3999,6 +3952,39 @@ def get_dspark_swa_page_indices( ) return swa_page_indices, swa_topk_lengths + @eager_on_graph + def _eager_low_ratio_source_projections( + self, layer, x, q_lora, positions, bufs, forward_batch + ): + + real = forward_batch.global_num_token_non_padded_cpu + if real is None: + real = x.shape[0] + + # These GEMMs pick their algorithm by M, so at the bucket size the live rows + # differ from eager; everything downstream is row-independent. + def put(name, value): + buf = bufs[name] + buf[:real].copy_(value) + buf[real:].zero_() + + if real == 0: + # An idle DP-attention rank replays on fabricated rows with no live + # token; a zero-row GEMM is a launch error, so only zero the buffers. + for buf in bufs.values(): + buf.zero_() + return + + if layer.compressor is not None: + kv, score = layer.compressor.project(x[:real]) + put("kv", kv) + if score is not None: + put("score", score) + if layer.indexer is not None: + indexer = layer.indexer + put("q", indexer.queries(q_lora[:real], layer.freqs_cis[positions[:real]])) + put("w", indexer.head_weights(x[:real])) + class DeepseekV4MultiStepBackend(DeepseekV4AttnBackend): def __init__( diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index f4e23f74b6a7..8189bb5b1c5b 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -19,23 +19,16 @@ fused_store_index_k_cache, ) from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz -from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.environ import envs from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import BaseIndexerMetadata from sglang.srt.layers.attention.dsa.dsa_npu_indexer import DSANPUIndexerMixin -from sglang.srt.layers.attention.dsa.dsa_prefill_cuda_graph import ( - GRAPH_WEIGHTS_PROJ_LORA_ERROR, - _is_in_piecewise_or_breakable_cuda_graph, - bcg_dsa_indexer_prefill_split, - pcg_dsa_indexer_prefill_split, -) from sglang.srt.layers.attention.dsa.paged_mqa_logits_backend import ( DSAPagedMQALogitsBackend, ) from sglang.srt.layers.attention.dsa.utils import ( aiter_can_use_preshuffle_paged_mqa, + is_dsa_bcg_prefill, is_dsa_enable_prefill_cp, - is_graph_dsa_split_op_surface, ) from sglang.srt.layers.attention.graph_variants import DSA_DENSE from sglang.srt.layers.attention.mqa_logits_utils import ( @@ -49,12 +42,10 @@ mqa_logits_static_budget_bytes, ) from sglang.srt.layers.layernorm import LayerNorm, RMSNorm -from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( +from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( + eager_on_graph, is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.runtime_context import ( get_device, get_exec, @@ -144,12 +135,20 @@ from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool +GRAPH_WEIGHTS_PROJ_LORA_ERROR = ( + "DSA indexer weights_proj LoRA is incompatible with " + "breakable CUDA graph; remove the explicit " + "prefill cuda-graph backend override or drop " + "indexer.weights_proj from the LoRA target modules." +) + + DUAL_STREAM_TOKEN_THRESHOLD = 1024 if _is_cuda else 0 if _is_cuda or _is_hip: - # Plain-torch graph helpers: usable wherever the split-op surface is. - from sglang.srt.layers.attention.dsa.dsa_prefill_cuda_graph import ( + # Head-gate custom ops support torch.compile on CUDA and HIP. + from sglang.srt.layers.attention.dsa.head_gate import ( logits_head_gate_graph, scale_head_gate_graph, ) @@ -162,7 +161,6 @@ from sglang.kernels.ops.attention.dsv4 import fused_q_indexer_rope_first_quant @register_custom_op(mutates_args=["topk_indices"]) - @register_split_op() def broadcast_indexer_topk_from_rank0_(topk_indices: torch.Tensor) -> None: _broadcast_indexer_topk_from_rank0_impl(topk_indices) @@ -194,10 +192,7 @@ def _broadcast_indexer_topk_from_rank0( if topk_indices is None or not envs.SGLANG_DSA_TOPK_BROADCAST.get(): return topk_indices - if is_in_tc_piecewise_cuda_graph(): - broadcast_indexer_topk_from_rank0_(topk_indices) - else: - _broadcast_indexer_topk_from_rank0_impl(topk_indices) + _broadcast_indexer_topk_from_rank0_impl(topk_indices) return topk_indices @@ -1627,15 +1622,10 @@ def forward_cuda( if TYPE_CHECKING: assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool) - in_piecewise_or_breakable_cuda_graph = ( - _is_in_piecewise_or_breakable_cuda_graph() - ) + in_breakable_cuda_graph = is_in_breakable_cuda_graph() - # In piecewise/breakable CUDA graph mode, metadata is fetched inside - # custom ops via get_tc_piecewise_forward_context() to prevent Dynamo - # from guarding on forward_metadata identity, which changes each replay - # when init_forward_metadata creates a new ForwardMetadata object. - if not in_piecewise_or_breakable_cuda_graph: + # Eager replay reads refreshed metadata from the active backend. + if not in_breakable_cuda_graph: metadata = get_attn_backend().get_indexer_metadata(layer_id, forward_batch) if metadata is None: return None @@ -1652,7 +1642,7 @@ def forward_cuda( # Determine if should skip topk based on sequence length # We can only skip the logits computation if cuda graph is not involved skip_logits_computation = False - if not in_piecewise_or_breakable_cuda_graph: + if not in_breakable_cuda_graph: skip_logits_computation = self._should_skip_logits_computation( forward_batch ) @@ -1681,7 +1671,7 @@ def forward_cuda( if ( self._aiter_fused_fp8_active(forward_batch) - and not in_piecewise_or_breakable_cuda_graph + and not in_breakable_cuda_graph and not weights_proj_lora ): q_fp8, weights = self._aiter_fused_fp8_prepare_and_store( @@ -1689,16 +1679,13 @@ def forward_cuda( ) elif ( self.use_dsa_indexer_fusion - and not in_piecewise_or_breakable_cuda_graph + and not in_breakable_cuda_graph and forward_batch.attn_cp_metadata is None ): q_fp8, weights = self._fused_q_prepare_and_store( x, q_lora, positions, forward_batch, layer_id, act_quant ) - elif ( - is_graph_dsa_split_op_surface(forward_batch) - and not self.dsa_enable_prefill_cp - ): + elif is_dsa_bcg_prefill(forward_batch) and not self.dsa_enable_prefill_cp: # Default path for non-CP prefill under PCG/BCG: run the whole indexer # (q/k proj, head gate, k-cache store, topk) as a single eager split op # instead of capturing it piecemeal in the graph. The split op is @@ -1716,12 +1703,8 @@ def forward_cuda( topk_result = torch.empty( (0, self.index_topk), device=x.device, dtype=torch.int32 ) - graph_dispatch_fn = ( - bcg_dsa_indexer_prefill_split - if is_in_breakable_cuda_graph() - else pcg_dsa_indexer_prefill_split - ) - graph_dispatch_fn( + self._eager_indexer( + forward_batch=forward_batch, layer_id=layer_id, x=x, q_lora=q_lora, @@ -1779,7 +1762,7 @@ def forward_cuda( act_quant=act_quant, ) current_stream.wait_stream(self.alt_stream) - elif not in_piecewise_or_breakable_cuda_graph: + elif not in_breakable_cuda_graph: q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt) self._store_index_k_cache( forward_batch=forward_batch, @@ -1835,7 +1818,7 @@ def forward_cuda( else: x_for_gate = x - if in_piecewise_or_breakable_cuda_graph: + if in_breakable_cuda_graph: if self.use_dsa_indexer_fusion: weights = scale_head_gate_graph( weights_raw, @@ -1865,7 +1848,7 @@ def forward_cuda( # In piecewise/breakable CUDA graph, any access to seq_lens_cpu # creates a Dynamo shape guard. These graph modes never have empty # batches. - if not in_piecewise_or_breakable_cuda_graph: + if not in_breakable_cuda_graph: if forward_batch.seq_lens.numel() == 0: # this seems b/c max-pad, no worries? # if x.shape[0] != 0: @@ -1893,7 +1876,7 @@ def forward_cuda( # In-graph (PCG/BCG) non-CP prefill is handled earlier by the # graph DSA split-op dispatch, so only the eager path reaches # here. - assert not in_piecewise_or_breakable_cuda_graph, ( + assert not in_breakable_cuda_graph, ( "Internal error: in-graph DSA prefill must go through the " "graph DSA split-op dispatch" ) @@ -1909,3 +1892,95 @@ def forward_cuda( raise NotImplementedError("DSA indexer only supports CUDA, HIP, and NPU") topk_result = _broadcast_indexer_topk_from_rank0(topk_result) return maybe_capture_indexer_topk(layer_id, topk_result) + + @eager_on_graph + def _eager_indexer( + self, + forward_batch: ForwardBatch, + layer_id: int, + x: torch.Tensor, + q_lora: torch.Tensor, + positions: torch.Tensor, + topk_result: torch.Tensor, + ) -> None: + # Run projections, head gating, cache storage, and top-k in one eager region. + # Mutate the caller's padded output so the next segment keeps a stable address. + assert _is_cuda, "Internal error: DSA graph dispatch is only supported on CUDA" + from sglang.kernels.ops.attention.dsa.triton_kernel import act_quant + + metadata = get_attn_backend().get_indexer_metadata(layer_id, forward_batch) + + extend_num_tokens = forward_batch.extend_num_tokens + # Empty buffer encodes return_indices=False for graph dispatch. + return_indices = topk_result.numel() != 0 + k_only = not return_indices or ( + self._should_skip_logits_computation(forward_batch) + and not self.dsa_enable_prefill_cp + ) + if k_only: + self._forward_cuda_k_only( + x, + positions, + forward_batch, + layer_id, + act_quant, + metadata=metadata, + return_indices=return_indices, + num_tokens=extend_num_tokens, + topk_result=topk_result, + ) + return + + # Fused path stores K (no-Hadamard) and computes q_fp8 + head gate in the + # fused kernels, sliced to the unpadded count, on a single stream. + if self.use_dsa_indexer_fusion: + q_fp8, weights = self._fused_q_prepare_and_store( + x, + q_lora, + positions, + forward_batch, + layer_id, + act_quant, + num_tokens=extend_num_tokens, + enable_dual_stream=False, + ) + self._get_topk_ragged( + False, + forward_batch, + layer_id, + q_fp8, + weights, + metadata, + topk_result, + ) + return + + query, key, _ = self._get_q_k_bf16( + q_lora, + x, + positions, + enable_dual_stream=False, + forward_batch=forward_batch, + ) + q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt) + # Reuse the compiled head-gate util shared with the eager path. + weights = self._get_logits_head_gate(x, q_scale) + # Store K cache + ragged top-k, sliced to the unpadded count and writing into + # the static padded topk_result buffer (the graph contract). Mirrors the eager + # path's store + _get_topk_ragged. + self._store_index_k_cache( + forward_batch=forward_batch, + layer_id=layer_id, + key=key[:extend_num_tokens], + act_quant=act_quant, + out_cache_loc=forward_batch.out_cache_loc[:extend_num_tokens], + ) + self._get_topk_ragged( + False, + forward_batch, + layer_id, + q_fp8[:extend_num_tokens], + weights, + metadata, + topk_result, + ) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer_kpool.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer_kpool.py index 91f279ea13c1..6f2d4e5bacba 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer_kpool.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer_kpool.py @@ -59,6 +59,7 @@ ) from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( + eager_on_graph, is_in_breakable_cuda_graph, ) from sglang.srt.model_executor.runner_utils import capture_mode @@ -1634,10 +1635,6 @@ def forward_cuda( is_in_breakable_cuda_graph() and forward_batch.forward_mode.is_extend_without_speculative() ): - from sglang.srt.layers.attention.dsa.kpool_prefill_cuda_graph import ( - bcg_kpool_indexer_prefill_with_output, - ) - # K-pool prefill plans contain request-specific tensors and launch # counts. Like the ordinary DSA indexer, execute them eagerly and # bridge the result into a stable buffer for captured attention. @@ -1649,9 +1646,7 @@ def forward_cuda( dtype=torch.int32, device=x.device, ) - bcg_kpool_indexer_prefill_with_output( - self, x, q_lora, positions, output, layer_id - ) + self._eager_indexer(forward_batch, x, q_lora, positions, output, layer_id) return output if return_indices else None return self._forward_cuda_impl( x, q_lora, positions, forward_batch, layer_id, return_indices @@ -1828,3 +1823,58 @@ def compress_write(): "kpool indexer is only supported on CUDA and ROCm" ) return topk_result + + def _capture_stub_indexer( + self, + forward_batch: ForwardBatch, + x: torch.Tensor, + q_lora: torch.Tensor, + positions: torch.Tensor, + output: torch.Tensor, + layer_id: int, + ) -> None: + output.fill_(-1) + + @eager_on_graph(capture_stub=_capture_stub_indexer) + def _eager_indexer( + self, + forward_batch: ForwardBatch, + x: torch.Tensor, + q_lora: torch.Tensor, + positions: torch.Tensor, + output: torch.Tensor, + layer_id: int, + ) -> None: + # Metadata, write counts and cache destinations change between requests. + # Resolve the live batch inside the eager break, never from capture args. + n = forward_batch.extend_num_tokens + if n is None or not 0 <= n <= x.shape[0]: + raise ValueError(f"Invalid pooled-indexer prefill token count: {n}") + if n > q_lora.shape[0] or n > positions.shape[0]: + raise ValueError("Pooled-indexer prefill inputs have inconsistent rows") + return_indices = output.shape[0] != 0 + result = self._forward_cuda_impl( + x=x[:n], + q_lora=q_lora[:n], + positions=positions[:n], + forward_batch=forward_batch, + layer_id=layer_id, + return_indices=return_indices, + ) + if not return_indices: + return + num_logical = sum(forward_batch.extend_seq_lens_cpu) + if ( + result is None + or result.ndim != 2 + or not num_logical <= result.shape[0] <= n + or result.shape[1] != output.shape[1] + ): + raise ValueError( + "Pooled-indexer prefill returned an unexpected top-k shape: got " + f"{None if result is None else tuple(result.shape)}, expected " + f"between {num_logical} and {n} rows of width {output.shape[1]}" + ) + # The following captured attention segment reads this stable padded buffer. + output[:num_logical].copy_(result[:num_logical]) + output[num_logical:].fill_(-1) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_prefill_cuda_graph.py b/python/sglang/srt/layers/attention/dsa/dsa_prefill_cuda_graph.py deleted file mode 100644 index 92a661a186bf..000000000000 --- a/python/sglang/srt/layers/attention/dsa/dsa_prefill_cuda_graph.py +++ /dev/null @@ -1,192 +0,0 @@ -from __future__ import annotations - -import torch - -from sglang.srt.compilation.compilation_config import register_split_op -from sglang.srt.model_executor.forward_context import get_attn_backend -from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( - eager_on_graph, -) -from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( - is_in_breakable_cuda_graph, -) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, - is_in_tc_piecewise_cuda_graph, -) -from sglang.srt.utils import is_cuda, is_hip -from sglang.srt.utils.custom_op import register_custom_op - -_is_cuda = is_cuda() -_is_hip = is_hip() - -GRAPH_WEIGHTS_PROJ_LORA_ERROR = ( - "DSA indexer weights_proj LoRA is incompatible with " - "piecewise/breakable CUDA graph; remove the explicit " - "prefill cuda-graph backend override or drop " - "indexer.weights_proj from the LoRA target modules." -) - - -def _is_in_piecewise_or_breakable_cuda_graph() -> bool: - return is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph() - - -if _is_cuda or _is_hip: - - def _scale_head_gate_graph_fake_impl( - weights_raw: torch.Tensor, - n_heads_inv_sqrt: float, - softmax_scale: float, - q_scale: torch.Tensor, - ) -> torch.Tensor: - return torch.empty( - (weights_raw.shape[0], weights_raw.shape[1], q_scale.shape[-1]), - dtype=torch.float32, - device=weights_raw.device, - ) - - # In-graph (PCG/BCG) head gate for the fused path: weights_proj is folded - # into wk_weights_proj, so weights_raw is precomputed and there is no GEMM. - @register_custom_op(fake_impl=_scale_head_gate_graph_fake_impl) - def scale_head_gate_graph( - weights_raw: torch.Tensor, - n_heads_inv_sqrt: float, - softmax_scale: float, - q_scale: torch.Tensor, - ) -> torch.Tensor: - weights = weights_raw * n_heads_inv_sqrt - return weights.unsqueeze(-1) * q_scale * softmax_scale - - def _logits_head_gate_graph_fake_impl( - x: torch.Tensor, - weight: torch.Tensor, - n_heads_inv_sqrt: float, - softmax_scale: float, - q_scale: torch.Tensor, - ) -> torch.Tensor: - return torch.empty( - (x.shape[0], weight.shape[0], q_scale.shape[-1]), - dtype=torch.float32, - device=x.device, - ) - - # In-graph (PCG/BCG) head gate for the NON-prefill path - @register_custom_op(fake_impl=_logits_head_gate_graph_fake_impl) - def logits_head_gate_graph( - x: torch.Tensor, - weight: torch.Tensor, - n_heads_inv_sqrt: float, - softmax_scale: float, - q_scale: torch.Tensor, - ) -> torch.Tensor: - out = torch.mm(x, weight.t(), out_dtype=torch.float32) - weights = out * n_heads_inv_sqrt - weights = weights.unsqueeze(-1) * q_scale * softmax_scale - return weights - - -@register_custom_op(mutates_args=["topk_result"]) -@register_split_op() -def pcg_dsa_indexer_prefill_split( - layer_id: int, - x: torch.Tensor, - q_lora: torch.Tensor, - positions: torch.Tensor, - topk_result: torch.Tensor, -) -> None: - # Default in-graph indexer path for non-CP prefill: runs the whole indexer - # (q/k proj, head gate, k-cache store, topk) as one eager split op. PCG calls - # this as a split op; BCG uses the explicit eager wrapper below. - # - # Output contract (differs from the eager `forward` path): a split op returns - # None, so results are delivered only by mutating `topk_result` in place. The - # call site pre-allocates it at a static, padded shape and a downstream - # captured graph reads it at a fixed address; eager code instead allocates - # and returns a fresh, naturally-sized tensor each call. - assert _is_cuda, "Internal error: DSA graph dispatch is only supported on CUDA" - from sglang.kernels.ops.attention.dsa.triton_kernel import act_quant - - forward_context = get_tc_piecewise_forward_context() - forward_batch = forward_context.forward_batch - indexer = forward_context.dsa_indexers[layer_id] - metadata = get_attn_backend().get_indexer_metadata(layer_id, forward_batch) - - extend_num_tokens = forward_batch.extend_num_tokens - # Empty buffer encodes return_indices=False for graph dispatch. - return_indices = topk_result.numel() != 0 - k_only = not return_indices or ( - indexer._should_skip_logits_computation(forward_batch) - and not indexer.dsa_enable_prefill_cp - ) - if k_only: - indexer._forward_cuda_k_only( - x, - positions, - forward_batch, - layer_id, - act_quant, - metadata=metadata, - return_indices=return_indices, - num_tokens=extend_num_tokens, - topk_result=topk_result, - ) - return - - # Fused path stores K (no-Hadamard) and computes q_fp8 + head gate in the - # fused kernels, sliced to the unpadded count. Single stream: the split op is - # captured, so the dual-stream overlap is disabled. - if indexer.use_dsa_indexer_fusion: - q_fp8, weights = indexer._fused_q_prepare_and_store( - x, - q_lora, - positions, - forward_batch, - layer_id, - act_quant, - num_tokens=extend_num_tokens, - enable_dual_stream=False, - ) - indexer._get_topk_ragged( - False, - forward_batch, - layer_id, - q_fp8, - weights, - metadata, - topk_result, - ) - return - - query, key, _ = indexer._get_q_k_bf16( - q_lora, - x, - positions, - enable_dual_stream=False, - forward_batch=forward_batch, - ) - q_fp8, q_scale = act_quant(query, indexer.block_size, indexer.scale_fmt) - # Reuse the compiled head-gate util shared with the eager path. - weights = indexer._get_logits_head_gate(x, q_scale) - # Store K cache + ragged top-k, sliced to the unpadded count and writing into - # the static padded topk_result buffer (the graph contract). Mirrors the eager - # path's store + _get_topk_ragged. - indexer._store_index_k_cache( - forward_batch=forward_batch, - layer_id=layer_id, - key=key[:extend_num_tokens], - act_quant=act_quant, - out_cache_loc=forward_batch.out_cache_loc[:extend_num_tokens], - ) - indexer._get_topk_ragged( - False, - forward_batch, - layer_id, - q_fp8[:extend_num_tokens], - weights, - metadata, - topk_result, - ) - - -bcg_dsa_indexer_prefill_split = eager_on_graph(True)(pcg_dsa_indexer_prefill_split) diff --git a/python/sglang/srt/layers/attention/dsa/head_gate.py b/python/sglang/srt/layers/attention/dsa/head_gate.py new file mode 100644 index 000000000000..264a12a0f083 --- /dev/null +++ b/python/sglang/srt/layers/attention/dsa/head_gate.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +import torch + +from sglang.srt.utils import is_cuda, is_hip +from sglang.srt.utils.custom_op import register_custom_op + +_is_cuda = is_cuda() +_is_hip = is_hip() + + +if _is_cuda or _is_hip: + + def _scale_head_gate_graph_fake_impl( + weights_raw: torch.Tensor, + n_heads_inv_sqrt: float, + softmax_scale: float, + q_scale: torch.Tensor, + ) -> torch.Tensor: + return torch.empty( + (weights_raw.shape[0], weights_raw.shape[1], q_scale.shape[-1]), + dtype=torch.float32, + device=weights_raw.device, + ) + + # In-graph head gate for the fused path: weights_proj is folded + # into wk_weights_proj, so weights_raw is precomputed and there is no GEMM. + @register_custom_op(fake_impl=_scale_head_gate_graph_fake_impl) + def scale_head_gate_graph( + weights_raw: torch.Tensor, + n_heads_inv_sqrt: float, + softmax_scale: float, + q_scale: torch.Tensor, + ) -> torch.Tensor: + weights = weights_raw * n_heads_inv_sqrt + return weights.unsqueeze(-1) * q_scale * softmax_scale + + def _logits_head_gate_graph_fake_impl( + x: torch.Tensor, + weight: torch.Tensor, + n_heads_inv_sqrt: float, + softmax_scale: float, + q_scale: torch.Tensor, + ) -> torch.Tensor: + return torch.empty( + (x.shape[0], weight.shape[0], q_scale.shape[-1]), + dtype=torch.float32, + device=x.device, + ) + + # In-graph head gate for the NON-prefill path + @register_custom_op(fake_impl=_logits_head_gate_graph_fake_impl) + def logits_head_gate_graph( + x: torch.Tensor, + weight: torch.Tensor, + n_heads_inv_sqrt: float, + softmax_scale: float, + q_scale: torch.Tensor, + ) -> torch.Tensor: + out = torch.mm(x, weight.t(), out_dtype=torch.float32) + weights = out * n_heads_inv_sqrt + weights = weights.unsqueeze(-1) * q_scale * softmax_scale + return weights diff --git a/python/sglang/srt/layers/attention/dsa/kpool_prefill_cuda_graph.py b/python/sglang/srt/layers/attention/dsa/kpool_prefill_cuda_graph.py deleted file mode 100644 index 385607af6fe2..000000000000 --- a/python/sglang/srt/layers/attention/dsa/kpool_prefill_cuda_graph.py +++ /dev/null @@ -1,70 +0,0 @@ -"""Breakable prefill bridge for the request-dependent pooled-key indexer.""" - -import torch - -from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( - eager_on_graph, -) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, -) - - -def _kpool_indexer_prefill_with_output( - indexer, - x: torch.Tensor, - q_lora: torch.Tensor, - positions: torch.Tensor, - output: torch.Tensor, - layer_id: int, -) -> None: - # Metadata, write counts and cache destinations change between requests. - # Resolve the live batch inside the eager break, never from capture args. - forward_batch = get_tc_piecewise_forward_context().forward_batch - n = forward_batch.extend_num_tokens - if n is None or not 0 <= n <= x.shape[0]: - raise ValueError(f"Invalid pooled-indexer prefill token count: {n}") - if n > q_lora.shape[0] or n > positions.shape[0]: - raise ValueError("Pooled-indexer prefill inputs have inconsistent rows") - return_indices = output.shape[0] != 0 - result = indexer._forward_cuda_impl( - x=x[:n], - q_lora=q_lora[:n], - positions=positions[:n], - forward_batch=forward_batch, - layer_id=layer_id, - return_indices=return_indices, - ) - if not return_indices: - return - num_logical = sum(forward_batch.extend_seq_lens_cpu) - if ( - result is None - or result.ndim != 2 - or not num_logical <= result.shape[0] <= n - or result.shape[1] != output.shape[1] - ): - raise ValueError( - "Pooled-indexer prefill returned an unexpected top-k shape: got " - f"{None if result is None else tuple(result.shape)}, expected " - f"between {num_logical} and {n} rows of width {output.shape[1]}" - ) - # The following captured attention segment reads this stable padded buffer. - output[:num_logical].copy_(result[:num_logical]) - output[num_logical:].fill_(-1) - - -def _kpool_indexer_prefill_capture_stub( - indexer, - x: torch.Tensor, - q_lora: torch.Tensor, - positions: torch.Tensor, - output: torch.Tensor, - layer_id: int, -) -> None: - output.fill_(-1) - - -bcg_kpool_indexer_prefill_with_output = eager_on_graph( - True, capture_stub=_kpool_indexer_prefill_capture_stub -)(_kpool_indexer_prefill_with_output) diff --git a/python/sglang/srt/layers/attention/dsa/utils.py b/python/sglang/srt/layers/attention/dsa/utils.py index 2ae70ec78e58..58bb0096e22c 100644 --- a/python/sglang/srt/layers/attention/dsa/utils.py +++ b/python/sglang/srt/layers/attention/dsa/utils.py @@ -9,9 +9,6 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.runtime_context import ( get_disagg, get_memory, @@ -135,15 +132,12 @@ def is_dsa_prefill_cp_interleave(): is_dsa_prefill_cp_round_robin_split = is_dsa_prefill_cp_interleave -# Structural surface where the graph DSA split-op dispatch (DSA indexer) and the -# MLA BMM-into-attention fusion apply: a non-speculative extend (prefill) running -# inside a piecewise/breakable CUDA graph. Both fusions are now on by default on -# this surface (no feature flag); each adds its own extra carve-outs at its call -# site (e.g. the indexer also excludes DSA prefill context parallelism). -def is_graph_dsa_split_op_surface(forward_batch: "ForwardBatch") -> bool: +# CUDA BCG prefill enables the DSA indexer eager region and MLA BMM-attention +# fusion. Each caller applies its own additional eligibility checks. +def is_dsa_bcg_prefill(forward_batch: "ForwardBatch") -> bool: return ( is_cuda() - and (is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph()) + and is_in_breakable_cuda_graph() and forward_batch.forward_mode.is_extend_without_speculative() ) diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 84a21fc5c960..5bc90338f167 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -426,13 +426,9 @@ def __init__( from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( is_in_breakable_cuda_graph, ) - from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, - ) from sglang.srt.utils import get_device_sm, is_blackwell self._is_in_breakable_cuda_graph = is_in_breakable_cuda_graph - self._is_in_tc_piecewise_cuda_graph = is_in_tc_piecewise_cuda_graph self._get_device_sm = get_device_sm self._is_blackwell = is_blackwell @@ -3595,12 +3591,11 @@ def set_dsa_prefill_impl(self, forward_batch: Optional[ForwardBatch] = None): """ # Hoisted in __init__ (import cost is per-call otherwise). is_in_breakable_cuda_graph = self._is_in_breakable_cuda_graph - is_in_tc_piecewise_cuda_graph = self._is_in_tc_piecewise_cuda_graph get_device_sm = self._get_device_sm is_blackwell = self._is_blackwell # Decide MHA vs MLA - if is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph(): + if is_in_breakable_cuda_graph(): # Can't branch on seq_lens_cpu in graph replay, force MHA off to # guarantee correctness. self.use_mha = False diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index f0bac3d4108a..6fc696cffd43 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -48,9 +48,6 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.runtime_context import ( get_exec, get_parallel, @@ -598,7 +595,6 @@ def _can_use_nonpaged_indexer( if ( get_parallel().attn_cp_size != 1 or self.hisparse_coordinator is not None - or is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph() ): return False diff --git a/python/sglang/srt/layers/attention/dsv4/metadata.py b/python/sglang/srt/layers/attention/dsv4/metadata.py index bd850d38b342..de2d8dc3658f 100644 --- a/python/sglang/srt/layers/attention/dsv4/metadata.py +++ b/python/sglang/srt/layers/attention/dsv4/metadata.py @@ -18,9 +18,6 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.model_executor.runner_utils.capture_mode import get_is_capture_mode from sglang.srt.utils import is_hip, is_sm120_supported, is_xpu @@ -310,7 +307,6 @@ def _mqa_logits_budget(self, *, num_rows: int) -> Optional[int]: get_is_capture_mode() or torch.cuda.is_current_stream_capturing() or is_in_breakable_cuda_graph() - or is_in_tc_piecewise_cuda_graph() ): return None return mqa_logits_budget_bytes( diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index e4d318207cc6..23b5fc99a49e 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -36,15 +36,7 @@ from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.mem_cache.memory_pool import KVWriteLoc -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.speculative.spec_info import SpecInput, SpecInputType from sglang.srt.speculative.spec_utils import ( draft_kv_indices_buffer_width, @@ -488,10 +480,7 @@ def __init__( model_runner.model_config.head_dim, model_runner.model_config.v_head_dim, ) - if ( - head_dims in cutlass_supported_head_dims - and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - ): + if head_dims in cutlass_supported_head_dims: fmha_backend = "cutlass" self.prefill_wrapper_ragged = BatchPrefillWithRaggedKVCacheWrapper( self.workspace_buffer, "NHD", backend=fmha_backend @@ -1006,11 +995,7 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): use_ragged = False extend_no_prefix = False else: - use_ragged = ( - not self.enable_deterministic - and not is_in_tc_piecewise_cuda_graph() - and not self.use_paged - ) + use_ragged = (not self.enable_deterministic) and (not self.use_paged) extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu) # Process multi-item scoring in attention backend instead of ForwardBatch diff --git a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py index 329223f87700..e35da7fd2519 100644 --- a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py @@ -36,9 +36,6 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.speculative.spec_info import SpecInput, SpecInputType from sglang.srt.speculative.spec_utils import ( draft_kv_indices_buffer_width, @@ -441,13 +438,9 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): prefix_lens = forward_batch.extend_prefix_lens extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu) use_ragged = ( - not get_exec().kernel.flashinfer_mla_disable_ragged - and extend_no_prefix - # Captured prefill (tc_piecewise or breakable) must use paged - # prefill: it stays compatible with prefix cache, and the ragged - # wrapper rejects the absorbed-MLA head dims (qk=576, vo=512). - and not is_in_tc_piecewise_cuda_graph() - and not is_in_breakable_cuda_graph() + (not get_exec().kernel.flashinfer_mla_disable_ragged) + and (extend_no_prefix) + and (not is_in_breakable_cuda_graph()) ) # build host indptr/len arrays for eager DRAFT_EXTEND_V2 fast plan path diff --git a/python/sglang/srt/layers/attention/graph_utils.py b/python/sglang/srt/layers/attention/graph_utils.py new file mode 100644 index 000000000000..de168a86745b --- /dev/null +++ b/python/sglang/srt/layers/attention/graph_utils.py @@ -0,0 +1,79 @@ +# Copyright 2023-2024 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Token slicing and temporary batch state for graph attention calls.""" + +from contextlib import contextmanager + +import torch + + +def _zero_padded_tokens(output, actual_tokens): + if actual_tokens is not None: + output[actual_tokens:].zero_() + + +@contextmanager +def attention_input_scope(forward_batch, output, query_tokens, kv_tokens, kwargs): + """Narrow per-token operands and restore the caller's batch even on failure.""" + kwargs = dict(kwargs) + for name in ( + "q_rope", + "topk_indices", + "rel_bias", + "q_descale", + "idx_q", + "mxfp8_norm_rope_positions", + "mxfp8_norm_rope_temp_scale", + ): + if kwargs.get(name) is not None: + kwargs[name] = kwargs[name][:query_tokens] + for name in ("k_rope", "k_descale", "v_descale", "idx_k", "idx_v"): + if kwargs.get(name) is not None: + kwargs[name] = kwargs[name][:kv_tokens] + if kwargs.get("aux_tensors") is not None: + kwargs["aux_tensors"] = [ + tensor[:query_tokens] for tensor in kwargs["aux_tensors"] + ] + cache_loc, positions = forward_batch.out_cache_loc, forward_batch.positions + previous_output = forward_batch._attn_output + forward_batch.out_cache_loc = cache_loc[:query_tokens] + if positions is not None: + forward_batch.positions = positions[:query_tokens] + forward_batch._attn_output = output[:query_tokens] + try: + yield kwargs + finally: + forward_batch.out_cache_loc, forward_batch.positions = cache_loc, positions + forward_batch._attn_output = previous_output + + +def allocate_attention_outputs(layer, query, value, index_query): + """Allocate graph-owned buffers before entering the eager attention region.""" + rows = query.shape[0] + sparse = index_query is not None + dtype = query.dtype if sparse or value is None else value.dtype + if dtype in (torch.float8_e4m3fn, torch.float8_e5m2): + dtype = torch.bfloat16 + shape = ( + (rows, layer.tp_q_head_num * layer.v_head_dim) + if sparse or layer.qk_head_dim != layer.v_head_dim + else query.shape + ) + output = torch.empty(shape, dtype=dtype, device=query.device) + index_output = ( + query.new_empty((rows, index_query.shape[1] * index_query.shape[2])) + if sparse + else None + ) + return output, index_output diff --git a/python/sglang/srt/layers/attention/hpc_ops_backend.py b/python/sglang/srt/layers/attention/hpc_ops_backend.py index 22f40a8073f8..2ac4d1faa751 100644 --- a/python/sglang/srt/layers/attention/hpc_ops_backend.py +++ b/python/sglang/srt/layers/attention/hpc_ops_backend.py @@ -40,19 +40,17 @@ from sglang.kernels.ops.kvcache.trtllm_mha_page_table import ( build_trtllm_mha_page_table, ) -from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.model_executor.forward_context import get_attn_backend +from sglang.srt.model_executor.forward_context import ( + get_attn_backend, + is_in_full_prefill_graph, +) from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( eager_on_graph, is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, -) -from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -392,7 +390,7 @@ def fused_qk_rope_store_kv_fp8( head_dim]. The subsequent ``self.attn(...)`` call must pass ``save_kv_cache=False`` (K/V are already written here). - Under a captured prefill graph (breakable / tc_piecewise) this routes + Under a captured prefill graph (breakable / full) this routes through a graph-splitting op — like attention itself — so the Python Q-scale hand-off between this op and attention stays alive at replay. @@ -409,29 +407,23 @@ def fused_qk_rope_store_kv_fp8( device=qkv.device, ) - if is_extend and get_tc_piecewise_forward_context() is not None: - # Captured prefill graph: run through the splitting op (eager at - # capture AND at replay, keeping the metadata hand-off alive). - if is_in_breakable_cuda_graph(): - breakable_hpc_ops_fp8_rope_store_kv( - qkv, - cos_sin_cache, - out_q, - layer.layer_id, - qk_norm_policy, - q_norm_weight=q_norm_weight, - k_norm_weight=k_norm_weight, - ) - else: - hpc_ops_fp8_rope_store_kv( - qkv, - cos_sin_cache, - out_q, - layer.layer_id, - qk_norm_policy, - q_norm_weight=q_norm_weight, - k_norm_weight=k_norm_weight, - ) + if is_extend and ( + is_in_full_prefill_graph() + or ( + is_in_breakable_cuda_graph() + and forward_batch.forward_mode.is_extend_without_speculative() + ) + ): + self._eager_rope_store_kv( + qkv, + cos_sin_cache, + out_q, + layer, + forward_batch, + qk_norm_policy, + q_norm_weight=q_norm_weight, + k_norm_weight=k_norm_weight, + ) else: self._run_fp8_rope_store_kv( layer=layer, @@ -656,44 +648,37 @@ def forward_decode( ) return o.view(-1, layer.tp_q_head_num * layer.head_dim) + @eager_on_graph + def _eager_rope_store_kv( + self, + qkv: torch.Tensor, + cos_sin_cache: torch.Tensor, + out_q: torch.Tensor, + layer, + forward_batch: ForwardBatch, + qk_norm_policy: int, + *, + q_norm_weight: Optional[torch.Tensor] = None, + k_norm_weight: Optional[torch.Tensor] = None, + ) -> None: + """Graph-splitting wrapper for the fused QKNorm+RoPE+FP8-quant+StoreKV op. -@register_custom_op(mutates_args=["out_q"]) -@register_split_op() -def hpc_ops_fp8_rope_store_kv( - qkv: torch.Tensor, - cos_sin_cache: torch.Tensor, - out_q: torch.Tensor, - layer_id: int, - qk_norm_policy: int, - *, - q_norm_weight: Optional[torch.Tensor] = None, - k_norm_weight: Optional[torch.Tensor] = None, -) -> None: - """Graph-splitting wrapper for the fused QKNorm+RoPE+FP8-quant+StoreKV op. - - Like ``unified_attention_with_output``, this runs eagerly between captured - prefill-graph segments (at capture and at every replay), so the Python - hand-off of the dynamic Q scales to the following attention op stays - alive. ``out_q`` is preallocated by the captured segment and mutated in - place, which is what stitches the surrounding graph segments together. - """ - context = get_tc_piecewise_forward_context() - forward_batch = context.forward_batch - attention_layer = context.attention_layers[layer_id] - real_num_tokens = forward_batch.global_num_token_non_padded_cpu - - backend = get_attn_backend() - backend._run_fp8_rope_store_kv( - layer=attention_layer, - forward_batch=forward_batch, - qkv=qkv[:real_num_tokens], - cos_sin_cache=cos_sin_cache, - q_norm_weight=q_norm_weight, - k_norm_weight=k_norm_weight, - qk_norm_policy=qk_norm_policy, - is_extend=True, - out_q=out_q[:real_num_tokens], - ) - - -breakable_hpc_ops_fp8_rope_store_kv = eager_on_graph(True)(hpc_ops_fp8_rope_store_kv) + Like ``RadixAttention._eager_attention``, this runs eagerly between captured + prefill-graph segments (at capture and at every replay), so the Python + hand-off of the dynamic Q scales to the following attention op stays + alive. ``out_q`` is preallocated by the captured segment and mutated in + place, which is what stitches the surrounding graph segments together. + """ + real_num_tokens = forward_batch.global_num_token_non_padded_cpu + + get_attn_backend()._run_fp8_rope_store_kv( + layer=layer, + forward_batch=forward_batch, + qkv=qkv[:real_num_tokens], + cos_sin_cache=cos_sin_cache, + q_norm_weight=q_norm_weight, + k_norm_weight=k_norm_weight, + qk_norm_policy=qk_norm_policy, + is_extend=True, + out_q=out_q[:real_num_tokens], + ) diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index bfe20646affb..a71d046ebcb1 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -61,9 +61,6 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.runtime_context import ( get_buffer, get_parallel, @@ -718,10 +715,8 @@ def init_mha_chunk_metadata( ) -> None: has_prefix = any(forward_batch.extend_prefix_lens_cpu) fallback_to_flashinfer_impl = ( - (self.disable_chunked_prefix_cache and has_prefix) - or is_in_tc_piecewise_cuda_graph() - or is_in_breakable_cuda_graph() - ) + self.disable_chunked_prefix_cache and has_prefix + ) or is_in_breakable_cuda_graph() if fallback_to_flashinfer_impl: super().init_mha_chunk_metadata( forward_batch, disable_flashinfer_ragged=True @@ -831,10 +826,8 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): # Also fallback to flashinfer MLA backend under a captured prefill graph has_prefix = any(forward_batch.extend_prefix_lens_cpu) fallback_to_flashinfer_impl = ( - (self.disable_chunked_prefix_cache and has_prefix) - or is_in_tc_piecewise_cuda_graph() - or is_in_breakable_cuda_graph() - ) + self.disable_chunked_prefix_cache and has_prefix + ) or is_in_breakable_cuda_graph() if fallback_to_flashinfer_impl: super().init_forward_metadata(forward_batch) diff --git a/python/sglang/srt/layers/cp/bcg.py b/python/sglang/srt/layers/cp/bcg.py index 6631c9d48ade..56556ac19c64 100644 --- a/python/sglang/srt/layers/cp/bcg.py +++ b/python/sglang/srt/layers/cp/bcg.py @@ -277,8 +277,6 @@ def execute_prefill_cp_bcg( assert cp_input is not None model = runner.model_runner.model with runner._prefill_forward_context( - static_forward_batch, - num_tokens=static_num_tokens, raw_num_tokens=raw_num_tokens, ): local_output = runner.backend.replay( diff --git a/python/sglang/srt/layers/layer_boundary/adapters/attention.py b/python/sglang/srt/layers/layer_boundary/adapters/attention.py index 3767408fadea..837538f7f987 100644 --- a/python/sglang/srt/layers/layer_boundary/adapters/attention.py +++ b/python/sglang/srt/layers/layer_boundary/adapters/attention.py @@ -29,11 +29,6 @@ enable_moe_dense_fully_dp, ) from sglang.srt.layers.moe import get_moe_a2a_backend -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.runtime_context import get_forward, get_parallel, get_spec from sglang.srt.utils import is_cuda, is_npu @@ -108,7 +103,6 @@ def init_context(self, q_lora_rank, is_dsa, is_mhc=False): and not is_dp_attention_enabled() and get_moe_a2a_backend().is_none() and not enable_moe_dense_fully_dp() - and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) and get_spec().speculative_algorithm != "EAGLE3" ) if get_parallel().enable_attn_tp_input_scattered: diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 6a4d049b653e..89c4642528d2 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -27,11 +27,6 @@ rms_norm_batch_invariant, ) from sglang.srt.environ import envs -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils import ( cpu_has_amx_support, @@ -790,8 +785,6 @@ def forward_musa( post_residual_addition: Optional[torch.Tensor] = None, quant_linear: Optional[nn.Module] = None, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE): - return self.forward_native(x, residual, post_residual_addition) if not x.is_contiguous(): x = x.contiguous() diff --git a/python/sglang/srt/layers/moe/ep_moe/layer.py b/python/sglang/srt/layers/moe/ep_moe/layer.py index d445bfc71e3c..115e308b4719 100644 --- a/python/sglang/srt/layers/moe/ep_moe/layer.py +++ b/python/sglang/srt/layers/moe/ep_moe/layer.py @@ -19,7 +19,6 @@ ) from sglang.srt.layers.moe.fused_moe_triton.layer import ( FusedMoE, - moe_forward_piecewise_cuda_graph_impl, ) from sglang.srt.layers.moe.token_dispatcher.deepep import ( DeepEPLLCombineInput, @@ -39,9 +38,6 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.utils import get_bool_env_var, is_hip, is_npu if TYPE_CHECKING: @@ -179,7 +175,20 @@ def __init__( f"DeepEP {self.deepep_mode} mode requires deep_gemm" ) - def _a2a_forward_with_output_impl( + def _a2a_forward_capture_stub( + self, + hidden_states: torch.Tensor, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + router_logits: torch.Tensor, + output: torch.Tensor, + ) -> None: + # Capture pass only: record the buffer address, skip the + # rank-coupled a2a. Warmup and replay run the real body. + output.zero_() + + @eager_on_graph(capture_stub=_a2a_forward_capture_stub) + def _eager_a2a_forward( self, hidden_states: torch.Tensor, topk_weights: torch.Tensor, @@ -200,22 +209,6 @@ def _a2a_forward_with_output_impl( finally: set_is_extend_in_batch(saved_is_extend_in_batch) - def _a2a_forward_capture_stub( - self, - hidden_states: torch.Tensor, - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - router_logits: torch.Tensor, - output: torch.Tensor, - ) -> None: - # Capture pass only: record the buffer address, skip the - # rank-coupled a2a. Warmup and replay run the real body. - output.zero_() - - a2a_forward_with_output = eager_on_graph( - True, capture_stub=_a2a_forward_capture_stub - )(_a2a_forward_with_output_impl) - def forward( self, hidden_states: torch.Tensor, @@ -227,7 +220,7 @@ def forward( "Only standard topk output is supported for breakable cuda graph" ) output = torch.empty_like(hidden_states) - self.a2a_forward_with_output( + self._eager_a2a_forward( hidden_states, topk_output.topk_weights, topk_output.topk_ids, @@ -235,19 +228,7 @@ def forward( output, ) return output - if is_in_tc_piecewise_cuda_graph(): - assert TopKOutputChecker.format_is_standard(topk_output), ( - "Only standard topk output is supported for piecewise cuda graph" - ) - return moe_forward_piecewise_cuda_graph_impl( - hidden_states, - topk_output.topk_weights, - topk_output.topk_ids, - topk_output.router_logits, - self.layer_id, - ) - else: - return self.forward_impl(hidden_states, topk_output) + return self.forward_impl(hidden_states, topk_output) def forward_impl( self, diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index 739ad7958366..1ac69b48cc0b 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -43,11 +43,7 @@ StandardDispatcher, ) from sglang.srt.layers.moe.topk import ( - BypassedTopKOutput, - StandardTopKOutput, - TopKConfig, TopKOutput, - TopKOutputChecker, ) from sglang.srt.layers.moe.utils import ( DispatcherOutputDtype, @@ -67,10 +63,6 @@ from sglang.srt.layers.quantization.fp8_utils import quantize_block_fp8_weight_to_mxfp4 from sglang.srt.layers.quantization.modelopt_quant import ModelOptNvFp4FusedMoEMethod from sglang.srt.layers.quantization.unquant import UnquantizedFusedMoEMethod -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight from sglang.srt.runtime_context import ( get_exec, @@ -87,7 +79,6 @@ is_npu, round_up, ) -from sglang.srt.utils.custom_op import register_custom_op _is_hip = is_hip() _is_cpu_amx_available = cpu_has_amx_support() @@ -1528,36 +1519,9 @@ def forward( from sglang.srt.hardware_backend.npu.moe.fuseep import forward_fuseep return forward_fuseep(self, hidden_states, topk_output) - if is_in_tc_piecewise_cuda_graph(): - if TopKOutputChecker.format_is_standard(topk_output): - return moe_forward_piecewise_cuda_graph_impl( - hidden_states, - topk_output.topk_weights, - topk_output.topk_ids, - topk_output.router_logits, - self.layer_id, - ) - elif TopKOutputChecker.format_is_bypassed(topk_output): - return fused_moe_bypassed_piecewise_cuda_graph_impl( - hidden_states, - topk_output.router_logits, - topk_output.topk_config.top_k, - topk_output.topk_config.topk_group, - topk_output.topk_config.num_expert_group, - topk_output.topk_config.correction_bias, - topk_output.topk_config.renormalize, - self.layer_id, - topk_output.topk_config.allow_routed_experts_capture, - ) - else: - # Make sure there is torch lib op registration for the whole moe layer - return self.forward_impl( - hidden_states, topk_output, pre_quant_input=pre_quant_input - ) - else: - return self.forward_impl( - hidden_states, topk_output, pre_quant_input=pre_quant_input - ) + return self.forward_impl( + hidden_states, topk_output, pre_quant_input=pre_quant_input + ) def forward_impl( self, @@ -1809,49 +1773,3 @@ def materialize_gguf_weights(self) -> None: stacked = torch.stack(weight_list, dim=0) param.materialize(stacked.shape, dtype=stacked.dtype) param.data.copy_(stacked) - - -@register_custom_op(out_shape="hidden_states") -def moe_forward_piecewise_cuda_graph_impl( - hidden_states: torch.Tensor, - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - router_logits: torch.Tensor, - layer_id: int, -) -> torch.Tensor: - # only standard topk output is supported for piecewise cuda graph - topk_output = StandardTopKOutput( - topk_weights=topk_weights, topk_ids=topk_ids, router_logits=router_logits - ) - forward_context = get_tc_piecewise_forward_context() - moe_layer = forward_context.moe_layers[layer_id] - return moe_layer.forward_impl(hidden_states, topk_output) - - -@register_custom_op(out_shape="hidden_states") -def fused_moe_bypassed_piecewise_cuda_graph_impl( - hidden_states: torch.Tensor, - router_logits: torch.Tensor, - top_k: int, - topk_group: Optional[int], - num_expert_group: Optional[int], - correction_bias: Optional[torch.Tensor], - renormalize: bool, - layer_id: int, - allow_routed_experts_capture: bool, -) -> torch.Tensor: - topk_output = BypassedTopKOutput( - hidden_states=hidden_states, - router_logits=router_logits, - topk_config=TopKConfig( - top_k=top_k, - topk_group=topk_group, - num_expert_group=num_expert_group, - correction_bias=correction_bias, - renormalize=renormalize, - allow_routed_experts_capture=allow_routed_experts_capture, - ), - ) - forward_context = get_tc_piecewise_forward_context() - moe_layer = forward_context.moe_layers[layer_id] - return moe_layer.forward_impl(hidden_states, topk_output) diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index fd86e97283fd..8adb6a9e8a6e 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -16,7 +16,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( fp8_dtype, fp8_max, - fp8_min, is_fp8_fnuz, per_token_group_quant_fp8, scaled_fp8_quant, @@ -2135,33 +2134,12 @@ def apply_fp8_linear( # On XPU, sgl-kernel-xpu's native quant kernels require output_q # to exactly match input's shape; padded output isn't supported. num_token_padding = None - # For static per-tensor activation scales when using inductor compiler, - # use pure PyTorch ops instead of the opaque sgl_kernel quant kernel. - # Inductor fuses these with surrounding ops (RMSNorm, residual add), - # eliminating a separate kernel launch per linear layer. - # weight_scale shape does not matter here -- it is only used in the - # GEMM epilogue, not in the activation quant fusion. Only activates when - # cuda_graph_config[prefill].tc_compiler=inductor; eager PCG and - # decode both use the faster custom kernel. - - if ( - input_scale is not None - and input_scale.numel() == 1 - and get_exec().graph.cuda_graph_config.prefill.tc_compiler == "inductor" - ): - qinput = ( - (input_2d * input_scale.reciprocal()) - .clamp(min=fp8_min, max=fp8_max) - .to(fp8_dtype) - ) - x_scale = input_scale - else: - qinput, x_scale = scaled_fp8_quant( - input_2d, - input_scale, - num_token_padding=num_token_padding, - use_per_token_if_dynamic=use_per_token_if_dynamic, - ) + qinput, x_scale = scaled_fp8_quant( + input_2d, + input_scale, + num_token_padding=num_token_padding, + use_per_token_if_dynamic=use_per_token_if_dynamic, + ) if ( input_scale is not None and channelwise_cutlass diff --git a/python/sglang/srt/layers/quantization/marlin_utils.py b/python/sglang/srt/layers/quantization/marlin_utils.py index e972584d5419..9dd46860d665 100644 --- a/python/sglang/srt/layers/quantization/marlin_utils.py +++ b/python/sglang/srt/layers/quantization/marlin_utils.py @@ -17,15 +17,11 @@ unpack_cols, ) from sglang.srt.utils import get_device_capability, is_cuda -from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: from sglang.srt.layers.linear import LinearBase from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, -) _is_cuda = is_cuda() @@ -423,8 +419,7 @@ def maybe_warn_marlin_atomic_add_env(): if torch.compiler.is_dynamo_compiling(): return # TODO(yiyun): Need to add sglang's MARLIN_USE_ATOMIC_ADD: bool = False - if True: - return + return # if envs.VLLM_MARLIN_USE_ATOMIC_ADD: # return logger.info_once( @@ -485,44 +480,25 @@ def apply_gptq_marlin_linear( dtype=input.dtype, ) - forward_context = get_tc_piecewise_forward_context() - if forward_context is None: - output = gptq_marlin_gemm( - reshaped_x, - None, - weight, - weight_scale, - None, - weight_zp, - g_idx, - g_idx_sort_indices, - workspace, - wtype, - size_m=reshaped_x.shape[0], - size_n=output_size_per_partition, - size_k=input_size_per_partition, - is_k_full=is_k_full, - use_atomic_add=use_atomic_add, - use_fp32_reduce=use_fp32_reduce, - is_zp_float=False, - ) - else: - output = unified_apply_gptq_marlin_gemm_with_wtype( - input=reshaped_x, - weight=weight, - weight_scale=weight_scale, - weight_zp=weight_zp, - g_idx=g_idx, - g_idx_sort_indices=g_idx_sort_indices, - workspace=workspace, - wtype_id=wtype.id, - output_size_per_partition=output_size_per_partition, - input_size_per_partition=input_size_per_partition, - is_k_full=is_k_full, - use_atomic_add=use_atomic_add, - use_fp32_reduce=use_fp32_reduce, - is_zp_float=False, - ) + output = gptq_marlin_gemm( + reshaped_x, + None, + weight, + weight_scale, + None, + weight_zp, + g_idx, + g_idx_sort_indices, + workspace, + wtype, + size_m=reshaped_x.shape[0], + size_n=output_size_per_partition, + size_k=input_size_per_partition, + is_k_full=is_k_full, + use_atomic_add=use_atomic_add, + use_fp32_reduce=use_fp32_reduce, + is_zp_float=False, + ) if bias is not None: output.add_(bias) # In-place add @@ -555,86 +531,8 @@ def apply_awq_marlin_linear( dtype=input.dtype, ) - forward_context = get_tc_piecewise_forward_context() - if forward_context is None: - output = gptq_marlin_gemm( - reshaped_x, - None, - weight, - weight_scale, - None, - weight_zp, - g_idx, - g_idx_sort_indices, - workspace, - quant_type, - size_m=reshaped_x.shape[0], - size_n=output_size_per_partition, - size_k=input_size_per_partition, - use_atomic_add=use_atomic_add, - use_fp32_reduce=use_fp32_reduce, - is_zp_float=False, - ) - else: - output = unified_apply_gptq_marlin_gemm( - input=reshaped_x, - weight=weight, - weight_scale=weight_scale, - weight_zp=weight_zp, - g_idx=g_idx, - g_idx_sort_indices=g_idx_sort_indices, - workspace=workspace, - output_size_per_partition=output_size_per_partition, - input_size_per_partition=input_size_per_partition, - use_atomic_add=use_atomic_add, - use_fp32_reduce=use_fp32_reduce, - is_zp_float=False, - ) - - if bias is not None: - output.add_(bias) # In-place add - - return output.reshape(out_shape) - - -def fake_unified_apply_gptq_marlin_gemm( - input: torch.Tensor, - weight: torch.Tensor, - weight_scale: torch.Tensor, - weight_zp: torch.Tensor, - g_idx: torch.Tensor, - g_idx_sort_indices: torch.Tensor, - workspace: torch.Tensor, - output_size_per_partition: int, - input_size_per_partition: int, - use_atomic_add: bool, - use_fp32_reduce: bool, - is_zp_float: bool, -) -> torch.Tensor: - return input.new_empty( - (input.shape[0], output_size_per_partition), dtype=input.dtype - ) - - -@register_custom_op(fake_impl=fake_unified_apply_gptq_marlin_gemm) -def unified_apply_gptq_marlin_gemm( - input: torch.Tensor, - weight: torch.Tensor, - weight_scale: torch.Tensor, - weight_zp: torch.Tensor, - g_idx: torch.Tensor, - g_idx_sort_indices: torch.Tensor, - workspace: torch.Tensor, - output_size_per_partition: int, - input_size_per_partition: int, - use_atomic_add: bool, - use_fp32_reduce: bool, - is_zp_float: bool, -) -> torch.Tensor: - quant_config = get_tc_piecewise_forward_context().quant_config - quant_type = quant_config.quant_type - return gptq_marlin_gemm( - input, + output = gptq_marlin_gemm( + reshaped_x, None, weight, weight_scale, @@ -644,77 +542,15 @@ def unified_apply_gptq_marlin_gemm( g_idx_sort_indices, workspace, quant_type, - size_m=input.shape[0], + size_m=reshaped_x.shape[0], size_n=output_size_per_partition, size_k=input_size_per_partition, use_atomic_add=use_atomic_add, use_fp32_reduce=use_fp32_reduce, - is_zp_float=is_zp_float, - ) - - -def fake_unified_apply_gptq_marlin_gemm_with_wtype( - input: torch.Tensor, - weight: torch.Tensor, - weight_scale: torch.Tensor, - weight_zp: torch.Tensor, - g_idx: torch.Tensor, - g_idx_sort_indices: torch.Tensor, - workspace: torch.Tensor, - wtype_id: int, - output_size_per_partition: int, - input_size_per_partition: int, - is_k_full: bool, - use_atomic_add: bool, - use_fp32_reduce: bool, - is_zp_float: bool, -) -> torch.Tensor: - return input.new_empty( - (input.shape[0], output_size_per_partition), dtype=input.dtype + is_zp_float=False, ) + if bias is not None: + output.add_(bias) # In-place add -@register_custom_op(fake_impl=fake_unified_apply_gptq_marlin_gemm_with_wtype) -def unified_apply_gptq_marlin_gemm_with_wtype( - input: torch.Tensor, - weight: torch.Tensor, - weight_scale: torch.Tensor, - weight_zp: torch.Tensor, - g_idx: torch.Tensor, - g_idx_sort_indices: torch.Tensor, - workspace: torch.Tensor, - wtype_id: int, - output_size_per_partition: int, - input_size_per_partition: int, - is_k_full: bool, - use_atomic_add: bool, - use_fp32_reduce: bool, - is_zp_float: bool, -) -> torch.Tensor: - # Reconstruct ScalarType from id - wtype = None - for attr_name in dir(scalar_types): - if not attr_name.startswith("_"): - st = getattr(scalar_types, attr_name) - if hasattr(st, "id") and st.id == wtype_id: - wtype = st - break - return gptq_marlin_gemm( - input, - None, - weight, - weight_scale, - None, - weight_zp, - g_idx, - g_idx_sort_indices, - workspace, - wtype, - size_m=input.shape[0], - size_n=output_size_per_partition, - size_k=input_size_per_partition, - is_k_full=is_k_full, - use_atomic_add=use_atomic_add, - use_fp32_reduce=use_fp32_reduce, - is_zp_float=is_zp_float, - ) + return output.reshape(out_shape) diff --git a/python/sglang/srt/layers/radix_attention.py b/python/sglang/srt/layers/radix_attention.py index 6840c21e2f08..a1be5a23e095 100644 --- a/python/sglang/srt/layers/radix_attention.py +++ b/python/sglang/srt/layers/radix_attention.py @@ -15,66 +15,26 @@ from __future__ import annotations -from contextlib import contextmanager -from contextvars import ContextVar from enum import Enum from typing import TYPE_CHECKING, Optional import torch from torch import nn -from sglang.srt.compilation.compilation_config import register_split_op -from sglang.srt.model_executor.forward_context import get_attn_backend +from sglang.srt.layers.attention.graph_utils import ( + _zero_padded_tokens, + allocate_attention_outputs, + attention_input_scope, +) +from sglang.srt.model_executor.forward_context import ( + get_attn_backend, + get_forward_context, + is_in_full_prefill_graph, +) from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( eager_on_graph, is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, -) -from sglang.srt.utils.common import is_hip -from sglang.srt.utils.custom_op import register_custom_op - -_is_hip = is_hip() - -# When set, RadixAttention.forward runs the attention backend eagerly instead of -# routing through the tc-piecewise split op. A caller already inside a -# breakable-CUDA-graph eager break (e.g. Inkling wrapping norm+attn+sconv in one eager -# region for multi-seq-correct short-conv metadata) sets this so the attn does not -# start a nested break (which would assert on the ended segment). Default off. -_force_eager_attn: ContextVar[bool] = ContextVar("_force_eager_attn", default=False) - - -@contextmanager -def force_eager_attention(): - token = _force_eager_attn.set(True) - try: - yield - finally: - _force_eager_attn.reset(token) - - -def _zero_padded_pcg_tail(buf: torch.Tensor, context) -> None: - """Zero the padded tail ``buf`` leaves as torch.empty garbage under PCG - replay, so NaN/Inf cannot reach residual / MoE routing / allreduce.""" - pcg_static_tokens = context.num_tokens - actual_tokens = context.raw_num_tokens - if ( - pcg_static_tokens is not None - and actual_tokens is not None - and pcg_static_tokens > actual_tokens - ): - first_dim = buf.shape[0] - elems_per_token = buf.numel() // first_dim - buf.view(first_dim, elems_per_token)[actual_tokens:].zero_() - - -def _zero_skipped_attn_outputs(*bufs: Optional[torch.Tensor]) -> None: - """Zero outputs when an idle DP rank skips attention work.""" - for buf in bufs: - if buf is not None: - buf.zero_() - if TYPE_CHECKING: from sglang.srt.layers.quantization.base_config import QuantizationConfig @@ -82,24 +42,14 @@ def _zero_skipped_attn_outputs(*bufs: Optional[torch.Tensor]) -> None: class AttentionType(Enum): - """ - Attention type. - Use string to be compatible with `torch.compile`. - """ + """String-valued attention types, compatible with torch.compile.""" - # Decoder attention between previous layer Q/K/V DECODER = "decoder" - # Decoder bidirectional attention between image tokens DECODER_BIDIRECTIONAL = "decoder_bidirectional" - # Encoder attention between previous layer Q/K/V ENCODER_ONLY = "encoder_only" class RadixAttention(nn.Module): - """ - The attention layer implementation. - """ - def __init__( self, num_heads: int, @@ -137,9 +87,6 @@ def __init__( self.v_scale = None self.k_scale_float = None self.v_scale_float = None - # MiniMax-M3 fp8 attention-GEMM scales (fp8 attn-GEMM mode): main q and - # lightning-indexer q/k/v. No checkpoint loader populates them yet; - # None means unit scale. self.q_scale_float = None self.idx_q_scale_float = None self.idx_k_scale_float = None @@ -167,510 +114,101 @@ def forward( **kwargs, ): if k is not None: - # For cross-layer sharing, kv can be None assert v is not None + k = k.view( + -1, + self.tp_k_head_num, + self.v_head_dim if "k_rope" in kwargs else self.qk_head_dim, + ) if "k_rope" not in kwargs: - k = k.view(-1, self.tp_k_head_num, self.qk_head_dim) v = v.view(-1, self.tp_v_head_num, self.v_head_dim) - else: - k = k.view(-1, self.tp_k_head_num, self.v_head_dim) - - context = get_tc_piecewise_forward_context() - if ( + breakable_cg = is_in_breakable_cuda_graph() + full_cg = is_in_full_prefill_graph() + use_eager_attention = ( self.use_prefill_attention_wrapper and forward_batch.forward_mode.is_extend() - and context is not None - # ``_force_eager_attn`` is only set inside Inkling's eager - # norm+attn+sconv region, never during tc-piecewise capture. Reading - # the ContextVar under the fullgraph torch.compile trace is - # untraceable ("Unsupported method call: ContextVar.get"), so - # short-circuit it while compiling -- force-eager is always off there. - and (torch.compiler.is_compiling() or not _force_eager_attn.get()) + and (breakable_cg or full_cg) + ) + # Sparse BCG already owns its backend-specific capture path. + if use_eager_attention and not ( + breakable_cg and kwargs.get("idx_q") is not None ): - if kwargs.get("idx_q") is not None: - if is_in_breakable_cuda_graph(): - return get_attn_backend().forward( - q, k, v, self, forward_batch, save_kv_cache, **kwargs - ) - idx_q = kwargs["idx_q"] - idx_k = kwargs["idx_k"] - idx_v = kwargs.get("idx_v") - attn_out = q.new_empty( - (q.shape[0], self.tp_q_head_num * self.v_head_dim) - ) - idx_out = q.new_empty((q.shape[0], idx_q.shape[1] * idx_q.shape[2])) - unified_sparse_attention_with_output( - q, - k, - v, - attn_out, - idx_out, - idx_q, - idx_k, - save_kv_cache, - self.layer_id, - idx_v=idx_v, - ) - return idx_out, attn_out - # Output dtype follows v (the model dtype) when available: qk-norm - # may emit q in a different dtype without changing the dtype the - # backend writes. FP8 q/v (e.g. mxfp8 KV-cache attention) still - # produce a bf16 attention output; sizing the buffer off an fp8 - # dtype would silently cast-copy the result to fp8. - out_dtype = v.dtype if v is not None else q.dtype - if out_dtype in (torch.float8_e4m3fn, torch.float8_e5m2): - out_dtype = torch.bfloat16 - if self.qk_head_dim != self.v_head_dim: - output = q.new_empty( - (q.shape[0], self.tp_q_head_num * self.v_head_dim), - dtype=out_dtype, - ) - else: - output = torch.empty_like(q, dtype=out_dtype) - if any( - key in kwargs - for key in ( - "score_mod", - "aux_tensors", - "rel_bias", - "return_lse", - "q_descale", - "k_descale", - "v_descale", - "mxfp8_norm_rope_positions", - ) - ): - # A score_mod callable, aux_tensors, rel_bias, mxfp8 descale - # tensors, or the mxfp8 deferred norm/RoPE operands can't cross - # the unified_attention_with_output custom-op schema; route this - # backend's extend attention through the plain eager path. - if is_in_breakable_cuda_graph(): - lse = breakable_attention_with_output_extra_kwargs( - q, k, v, output, save_kv_cache, self.layer_id, kwargs - ) - else: - lse = attention_with_output_extra_kwargs( - q, k, v, output, save_kv_cache, self.layer_id, kwargs - ) - if kwargs.get("return_lse") or forward_batch.mha_return_lse: - assert lse is not None - return output.view(-1, self.tp_q_head_num, self.v_head_dim), lse - return output - # Chunked-prefix MHA needs LSE to merge independently normalized - # suffix and cached-prefix attention states. - return_lse = bool(forward_batch.mha_return_lse) - mha_companion_layers = context.mha_companion_layers - use_mha_companion = ( - mha_companion_layers is not None - and mha_companion_layers[self.layer_id] is self + is_sparse = kwargs.get("idx_q") is not None + output, idx_output = allocate_attention_outputs( + self, q, v, kwargs.get("idx_q") ) - if is_in_breakable_cuda_graph(): - op = ( - breakable_unified_attention_with_output_and_lse - if return_lse - else breakable_unified_attention_with_output - ) - else: - op = ( - unified_attention_with_output_and_lse - if return_lse - else unified_attention_with_output - ) - lse = op( + lse = self._eager_attention( q, k, v, output, + forward_batch, save_kv_cache, - self.layer_id, - use_mha_companion=use_mha_companion, - key_value_num_tokens=key_value_num_tokens, + key_value_num_tokens, + idx_output, **kwargs, ) - if return_lse: + if is_sparse: + return idx_output, output + if kwargs.get("return_lse") or forward_batch.mha_return_lse: return output.view(-1, self.tp_q_head_num, self.v_head_dim), lse return output - else: - return get_attn_backend().forward( - q, - k, - v, - self, - forward_batch, - save_kv_cache, - **kwargs, - ) - - -def _unified_attention_with_output_impl( - query: torch.Tensor, - key: Optional[torch.Tensor], - value: Optional[torch.Tensor], - output: torch.Tensor, - save_kv_cache: bool, - layer_id: int, - use_mha_companion: bool, - return_lse: bool, - *, - key_value_num_tokens: Optional[int] = None, - q_rope: Optional[torch.Tensor] = None, - k_rope: Optional[torch.Tensor] = None, - sinks: Optional[torch.Tensor] = None, - attn_sink: Optional[torch.Tensor] = None, - # MLA / TRT-LLM / NSA paths pass these through RadixAttention.forward(**kwargs); - # they must appear in the schema when --cuda-graph-backend-prefill=tc_piecewise is on. - cos_sin_cache: Optional[torch.Tensor] = None, - is_neox: Optional[bool] = None, - llama_4_scaling: Optional[torch.Tensor] = None, - topk_indices: Optional[torch.Tensor] = None, -) -> Optional[torch.Tensor]: - context = get_tc_piecewise_forward_context() - forward_batch = context.forward_batch - attention_layers = context.attention_layers - attention_layer = attention_layers[layer_id] - real_query_num_tokens = forward_batch.global_num_token_non_padded_cpu - # Ordinary PCG attention pads Q/K/V to the same token bucket. Prefix MHA - # instead supplies a fixed-capacity K/V chunk whose extent is independent - # of the suffix queries, so its caller must preserve that separate extent. - if key_value_num_tokens is None: - key_value_num_tokens = real_query_num_tokens - - if real_query_num_tokens == 0: - _zero_skipped_attn_outputs(output) - if return_lse: - # unified_attention_with_output_and_lse asserts a tensor comes back. - # Match _unified_attention_with_output_and_lse_fake's meta shape and - # the padded LSE the normal path returns below (padded row count, - # i.e. query before narrowing). - return query.new_zeros( - (query.shape[0], query.shape[1]), dtype=torch.float32 - ) - return None - - query = query[:real_query_num_tokens] - if key is not None: - key = key[:key_value_num_tokens] - if value is not None: - value = value[:key_value_num_tokens] - - # DeepSeek MLA has two RadixAttention instances per layer (attn_mqa and - # attn_mha) that share the same layer_id. Preserve the calling instance's - # identity through the custom-op boundary; save_kv_cache is not an identity - # signal because absorbed MLA can also disable a redundant cache store. - if use_mha_companion: - assert context.mha_companion_layers is not None - attention_layer = context.mha_companion_layers[layer_id] - assert attention_layer is not None - - kwargs = {} - if q_rope is not None: - kwargs["q_rope"] = q_rope[:real_query_num_tokens] - if k_rope is not None: - kwargs["k_rope"] = k_rope[:key_value_num_tokens] - if sinks is not None: - kwargs["sinks"] = sinks - if attn_sink is not None: - kwargs["attn_sink"] = attn_sink - if cos_sin_cache is not None: - kwargs["cos_sin_cache"] = cos_sin_cache - if is_neox is not None: - kwargs["is_neox"] = is_neox - if llama_4_scaling is not None: - kwargs["llama_4_scaling"] = llama_4_scaling - if topk_indices is not None: - kwargs["topk_indices"] = topk_indices[:real_query_num_tokens] - - original_out_cache_loc = forward_batch.out_cache_loc - original_positions = forward_batch.positions - # Keep the original ForwardBatch object and only narrow cache locations for - # this backend call so model/backend state is still written to the same batch. - forward_batch.out_cache_loc = original_out_cache_loc[:real_query_num_tokens] - if original_positions is not None: - forward_batch.positions = original_positions[:real_query_num_tokens] - - # Store pre-allocated output for FA backend to write directly into. - # Must slice to real_query_num_tokens to match the narrowed query shape — - # the FA kernel validates out.size(0) == q.size(0). - forward_batch._attn_output = output[:real_query_num_tokens] - - ret = get_attn_backend().forward( - query, - key, - value, - attention_layer, - forward_batch, - save_kv_cache, - **kwargs, - ) - forward_batch.out_cache_loc = original_out_cache_loc - forward_batch.positions = original_positions - - lse = None - if return_lse: - assert isinstance(ret, tuple) - ret, lse, *_ = ret - else: - assert isinstance(ret, torch.Tensor) - - if ret.data_ptr() != output.data_ptr(): - output[:real_query_num_tokens].view(ret.shape).copy_(ret) + return get_attn_backend().forward( + q, k, v, self, forward_batch, save_kv_cache, **kwargs + ) - # During PCG replay the attention backend writes only the narrowed - # real-token slice (output[:real_query_num_tokens]) and leaves padded positions - # as uninitialized torch.empty garbage. Zero them so garbage (NaN/Inf) does - # not propagate through residual connections, MoE routing, and allreduce. - # This affects every backend that varlen-writes under PCG, not just ROCm. - # Use context.raw_num_tokens (pre-padding count from PCG runner) instead of - # forward_batch.extend_num_tokens, which is None for TARGET_VERIFY batches. - _zero_padded_pcg_tail(output, context) - if lse is not None and lse.shape[0] != output.shape[0]: - padded_lse = lse.new_zeros((output.shape[0], *lse.shape[1:])) - padded_lse[:real_query_num_tokens].copy_(lse) - lse = padded_lse - return lse - - -@register_custom_op(mutates_args=["output"]) -@register_split_op() -def unified_attention_with_output( - query: torch.Tensor, - key: Optional[torch.Tensor], - value: Optional[torch.Tensor], - output: torch.Tensor, - save_kv_cache: bool, - layer_id: int, - *, - use_mha_companion: bool = False, - key_value_num_tokens: Optional[int] = None, - q_rope: Optional[torch.Tensor] = None, - k_rope: Optional[torch.Tensor] = None, - sinks: Optional[torch.Tensor] = None, - attn_sink: Optional[torch.Tensor] = None, - cos_sin_cache: Optional[torch.Tensor] = None, - is_neox: Optional[bool] = None, - llama_4_scaling: Optional[torch.Tensor] = None, - topk_indices: Optional[torch.Tensor] = None, -) -> None: - _unified_attention_with_output_impl( - query, - key, - value, - output, - save_kv_cache, - layer_id, - use_mha_companion, - False, - key_value_num_tokens=key_value_num_tokens, - q_rope=q_rope, - k_rope=k_rope, - sinks=sinks, - attn_sink=attn_sink, - cos_sin_cache=cos_sin_cache, - is_neox=is_neox, - llama_4_scaling=llama_4_scaling, - topk_indices=topk_indices, - ) - - -def _unified_attention_with_output_and_lse_fake( - query: torch.Tensor, *args, **kwargs -) -> torch.Tensor: - return query.new_empty((query.shape[0], query.shape[1]), dtype=torch.float32) - - -@register_custom_op( - mutates_args=["output"], fake_impl=_unified_attention_with_output_and_lse_fake -) -@register_split_op() -def unified_attention_with_output_and_lse( - query: torch.Tensor, - key: Optional[torch.Tensor], - value: Optional[torch.Tensor], - output: torch.Tensor, - save_kv_cache: bool, - layer_id: int, - *, - use_mha_companion: bool = False, - key_value_num_tokens: Optional[int] = None, - q_rope: Optional[torch.Tensor] = None, - k_rope: Optional[torch.Tensor] = None, - sinks: Optional[torch.Tensor] = None, - attn_sink: Optional[torch.Tensor] = None, - cos_sin_cache: Optional[torch.Tensor] = None, - is_neox: Optional[bool] = None, - llama_4_scaling: Optional[torch.Tensor] = None, - topk_indices: Optional[torch.Tensor] = None, -) -> torch.Tensor: - lse = _unified_attention_with_output_impl( - query, - key, - value, + @eager_on_graph + def _eager_attention( + self, + q, + k, + v, output, - save_kv_cache, - layer_id, - use_mha_companion, - True, - key_value_num_tokens=key_value_num_tokens, - q_rope=q_rope, - k_rope=k_rope, - sinks=sinks, - attn_sink=attn_sink, - cos_sin_cache=cos_sin_cache, - is_neox=is_neox, - llama_4_scaling=llama_4_scaling, - topk_indices=topk_indices, - ) - assert lse is not None - return lse - - -@register_custom_op(mutates_args=["attn_out", "idx_out"]) -@register_split_op() -def unified_sparse_attention_with_output( - query: torch.Tensor, - key: Optional[torch.Tensor], - value: Optional[torch.Tensor], - attn_out: torch.Tensor, - idx_out: torch.Tensor, - idx_q: torch.Tensor, - idx_k: torch.Tensor, - save_kv_cache: bool, - layer_id: int, - *, - idx_v: Optional[torch.Tensor] = None, -) -> None: - context = get_tc_piecewise_forward_context() - forward_batch = context.forward_batch - attention_layer = context.attention_layers[layer_id] - real_num_tokens = forward_batch.global_num_token_non_padded_cpu - - if real_num_tokens == 0: - _zero_skipped_attn_outputs(attn_out, idx_out) - return - - query = query[:real_num_tokens] - if key is not None: - key = key[:real_num_tokens] - if value is not None: - value = value[:real_num_tokens] - idx_q = idx_q[:real_num_tokens] - idx_k = idx_k[:real_num_tokens] - if idx_v is not None: - idx_v = idx_v[:real_num_tokens] - - original_out_cache_loc = forward_batch.out_cache_loc - forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens] - - ret_idx, ret_out = get_attn_backend().forward( - query, - key, - value, - attention_layer, - forward_batch, - save_kv_cache, - idx_q=idx_q, - idx_k=idx_k, - idx_v=idx_v, - ) - forward_batch.out_cache_loc = original_out_cache_loc - - attn_out[:real_num_tokens].view(ret_out.shape).copy_(ret_out) - # disable_value layers return ret_idx=None; the guard keeps idx_out's - # untouched real-token slice safe (model returns before index_o_proj). - if ret_idx is not None: - idx_out[:real_num_tokens].view(ret_idx.shape).copy_(ret_idx) - - for buf in (attn_out, idx_out): - _zero_padded_pcg_tail(buf, context) - return - - -breakable_unified_attention_with_output = eager_on_graph(True)( - unified_attention_with_output -) -breakable_unified_attention_with_output_and_lse = eager_on_graph(True)( - unified_attention_with_output_and_lse -) - - -def attention_with_output_extra_kwargs( - query: torch.Tensor, - key: Optional[torch.Tensor], - value: Optional[torch.Tensor], - output: torch.Tensor, - save_kv_cache: bool, - layer_id: int, - extra_kwargs: dict, -) -> Optional[torch.Tensor]: - """Breakable/tc_piecewise attention for backends whose forward needs kwargs - that cannot cross the ``unified_attention_with_output`` custom-op schema -- - a ``score_mod`` callable and/or ``aux_tensors`` (e.g. Inkling's relative-bias - fa4 attention), or the per-token mxfp8 deferred norm/RoPE operands. Plain - (not a custom op) so the callable passes through; still - runs eagerly between graph segments under BCG via the wrapper below. Mirrors - the real-token narrowing + padded-output write of - ``unified_attention_with_output``, and narrows per-token ``aux_tensors`` too. - """ - context = get_tc_piecewise_forward_context() - forward_batch = context.forward_batch - attention_layer = context.attention_layers[layer_id] - real_num_tokens = forward_batch.global_num_token_non_padded_cpu - - if real_num_tokens == 0: - _zero_skipped_attn_outputs(output) - return - - query = query[:real_num_tokens] - if key is not None: - key = key[:real_num_tokens] - if value is not None: - value = value[:real_num_tokens] - - kwargs = dict(extra_kwargs) - aux_tensors = kwargs.get("aux_tensors") - if aux_tensors is not None: - kwargs["aux_tensors"] = [t[:real_num_tokens] for t in aux_tensors] - for per_token_key in ( - "rel_bias", - "q_descale", - "k_descale", - "v_descale", - "mxfp8_norm_rope_positions", - "mxfp8_norm_rope_temp_scale", + forward_batch: ForwardBatch, + save_kv_cache: bool = True, + key_value_num_tokens: Optional[int] = None, + idx_output=None, + **kwargs, ): - t = kwargs.get(per_token_key) - if t is not None: - kwargs[per_token_key] = t[:real_num_tokens] - - original_out_cache_loc = forward_batch.out_cache_loc - forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens] - forward_batch._attn_output = output[:real_num_tokens] - - ret = get_attn_backend().forward( - query, key, value, attention_layer, forward_batch, save_kv_cache, **kwargs - ) - forward_batch.out_cache_loc = original_out_cache_loc - - return_lse = bool(kwargs.get("return_lse") or forward_batch.mha_return_lse) - if return_lse: - assert isinstance(ret, tuple) - ret, lse, *_ = ret - else: - assert isinstance(ret, torch.Tensor) + """Run real tokens, retaining padded outputs for the next graph segment.""" + rows = output.shape[0] + n = forward_batch.global_num_token_non_padded_cpu + kv_n = n if key_value_num_tokens is None else key_value_num_tokens + return_lse = bool(kwargs.get("return_lse") or forward_batch.mha_return_lse) lse = None - - if ret.data_ptr() != output.data_ptr(): - output[:real_num_tokens].view(ret.shape).copy_(ret) - - if _is_hip: - _zero_padded_pcg_tail(output, context) - if lse is not None and lse.shape[0] != output.shape[0]: - padded_lse = lse.new_zeros((output.shape[0], *lse.shape[1:])) - padded_lse[:real_num_tokens].copy_(lse) - lse = padded_lse - return lse - - -breakable_attention_with_output_extra_kwargs = eager_on_graph(True)( - attention_with_output_extra_kwargs -) + if n == 0: + output.zero_() + if idx_output is not None: + idx_output.zero_() + if return_lse: + lse = q.new_zeros((rows, q.shape[1]), dtype=torch.float32) + else: + with attention_input_scope( + forward_batch, output, n, kv_n, kwargs + ) as kwargs: + result = get_attn_backend().forward( + q[:n], + k[:kv_n] if k is not None else None, + v[:kv_n] if v is not None else None, + self, + forward_batch, + save_kv_cache, + **kwargs, + ) + if idx_output is not None: + idx_result, result = result + if idx_result is not None: + idx_output[:n].view(idx_result.shape).copy_(idx_result) + elif return_lse: + result, lse, *_ = result + if result.data_ptr() != output.data_ptr(): + output[:n].view(result.shape).copy_(result) + raw = get_forward_context().raw_num_tokens + _zero_padded_tokens(output, raw) + if idx_output is not None: + _zero_padded_tokens(idx_output, raw) + if lse is not None and lse.shape[0] != rows: + padded = lse.new_zeros((rows, *lse.shape[1:])) + padded[:n].copy_(lse) + lse = padded + return lse diff --git a/python/sglang/srt/layers/radix_linear_attention.py b/python/sglang/srt/layers/radix_linear_attention.py index 39201effbe85..c572eec8d3aa 100644 --- a/python/sglang/srt/layers/radix_linear_attention.py +++ b/python/sglang/srt/layers/radix_linear_attention.py @@ -20,16 +20,14 @@ import torch from torch import nn -from sglang.srt.compilation.compilation_config import register_split_op -from sglang.srt.model_executor.forward_context import get_attn_backend +from sglang.srt.model_executor.forward_context import ( + get_attn_backend, + is_in_full_prefill_graph, +) from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( eager_on_graph, is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, -) -from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: from sglang.srt.model_executor.forward_batch_info import ForwardBatch @@ -85,7 +83,13 @@ def forward( b: torch.Tensor, ) -> torch.Tensor: is_extend = forward_batch.forward_mode.is_extend() - if is_extend and get_tc_piecewise_forward_context() is not None: + if is_extend and ( + is_in_full_prefill_graph() + or ( + is_in_breakable_cuda_graph() + and forward_batch.forward_mode.is_extend_without_speculative() + ) + ): # Output shape from linear attention: (1, seq_len, num_v_heads, head_v_dim) seq_len = mixed_qkv.shape[0] output = torch.empty( @@ -93,22 +97,7 @@ def forward( dtype=mixed_qkv.dtype, device=mixed_qkv.device, ) - if is_in_breakable_cuda_graph(): - bcg_unified_linear_attention_with_output( - mixed_qkv, - a, - b, - output, - self.layer_id, - ) - else: - unified_linear_attention_with_output( - mixed_qkv, - a, - b, - output, - self.layer_id, - ) + self._eager_linear_attention(mixed_qkv, a, b, output, forward_batch) return output # Target verify rebuilds query_start_loc from the physical padded input, @@ -117,7 +106,7 @@ def forward( is_extend and not forward_batch.forward_mode.is_target_verify() ) real_num_tokens = ( - getattr(forward_batch, "global_num_token_non_padded_cpu", None) + forward_batch.global_num_token_non_padded_cpu if should_trim_padded_extend else None ) @@ -147,6 +136,13 @@ def forward( b=b, ) + def _capture_stub_linear_attention(self, mixed_qkv, a, b, output, forward_batch): + output.zero_() + + @eager_on_graph(capture_stub=_capture_stub_linear_attention) + def _eager_linear_attention(self, mixed_qkv, a, b, output, forward_batch): + _linear_attention_with_output_impl(mixed_qkv, a, b, output, self, forward_batch) + def _linear_attention_with_output_impl( mixed_qkv: torch.Tensor, @@ -190,60 +186,3 @@ def _linear_attention_with_output_impl( # Physical padding participates in following residual, router, expert/MoE, # and collective operations. Keep those inputs finite and deterministic. output[:, real_num_tokens:].zero_() - - -def _unified_linear_attention_with_output_impl( - mixed_qkv: torch.Tensor, - a: torch.Tensor, - b: torch.Tensor, - output: torch.Tensor, - layer_id: int, -) -> None: - """Eager implementation kept separate for backend-independent tests.""" - context = get_tc_piecewise_forward_context() - forward_batch = context.forward_batch - attention_layers = context.attention_layers - attention_layer = attention_layers[layer_id] - _linear_attention_with_output_impl( - mixed_qkv=mixed_qkv, - a=a, - b=b, - output=output, - attention_layer=attention_layer, - forward_batch=forward_batch, - ) - return - - -@register_custom_op(mutates_args=["output"]) -@register_split_op() -def unified_linear_attention_with_output( - mixed_qkv: torch.Tensor, - a: torch.Tensor, - b: torch.Tensor, - output: torch.Tensor, - layer_id: int, -) -> None: - """Custom op wrapper for linear attention computation only.""" - _unified_linear_attention_with_output_impl( - mixed_qkv=mixed_qkv, - a=a, - b=b, - output=output, - layer_id=layer_id, - ) - - -def _linear_attention_capture_stub( - mixed_qkv: torch.Tensor, - a: torch.Tensor, - b: torch.Tensor, - output: torch.Tensor, - layer_id: int, -) -> None: - output.zero_() - - -bcg_unified_linear_attention_with_output = eager_on_graph( - True, capture_stub=_linear_attention_capture_stub -)(unified_linear_attention_with_output) diff --git a/python/sglang/srt/model_executor/cuda_graph_config.py b/python/sglang/srt/model_executor/cuda_graph_config.py index 760659ede26a..042b7272d9fc 100644 --- a/python/sglang/srt/model_executor/cuda_graph_config.py +++ b/python/sglang/srt/model_executor/cuda_graph_config.py @@ -42,16 +42,14 @@ class Backend: FULL = "full" BREAKABLE = "breakable" - TC_PIECEWISE = "tc_piecewise" DISABLED = "disabled" - ALL = (FULL, BREAKABLE, TC_PIECEWISE, DISABLED) + ALL = (FULL, BREAKABLE, DISABLED) ALLOWED_BACKENDS_PER_PHASE = { Phase.DECODE: ( Backend.FULL, Backend.BREAKABLE, - Backend.TC_PIECEWISE, Backend.DISABLED, ), # full for prefill captures one whole-forward graph per num_tokens @@ -61,26 +59,23 @@ class Backend: Phase.PREFILL: ( Backend.FULL, Backend.BREAKABLE, - Backend.TC_PIECEWISE, Backend.DISABLED, ), } # Per-phase settings schema. Keys other than backend are runner-level -# (read by any backend in that phase); tc_compiler is the lone -# backend-specific knob (only meaningful when backend == tc_piecewise). +# (read by any backend in that phase). # For prefill, bs carries aggregate-token capture buckets for every backend; # full_prefill_max_req separately controls Full's fixed request-slot count. # full_prefill_max_req and full_prefill_prefix_chunk_tokens are prefill-only and # only meaningful when backend == full. max_context_size is shared by the # breakable and full prefill body-capture backends. ALLOWED_KEYS_PER_PHASE = { - Phase.DECODE: ("backend", "max_bs", "bs", "tc_compiler"), + Phase.DECODE: ("backend", "max_bs", "bs"), Phase.PREFILL: ( "backend", "max_bs", "bs", - "tc_compiler", "max_context_size", "full_prefill_max_req", "full_prefill_prefix_chunk_tokens", @@ -96,8 +91,6 @@ class PhaseConfig: backend: str = Backend.DISABLED max_bs: Optional[int] = None bs: Optional[List[int]] = None - # Only meaningful when backend == tc_piecewise; ignored otherwise. - tc_compiler: str = "eager" # Effective for both full and breakable backends and currently only DSV4: # fixed maximum context length used by context-shaped prefill graph metadata. # Every token bucket shares this size; larger live contexts run eagerly. @@ -105,7 +98,7 @@ class PhaseConfig: # Only meaningful for the prefill phase with backend == full: max number of # request slots baked into each captured graph. Real bs <= full_prefill_max_req # reuses the graph (unused slots become zero-length sentinels); larger - # batches fall back to eager. Ignored by BCG and TC_PIECEWISE. None + # batches fall back to eager. Ignored by BCG. None # auto-derives chunked_prefill_size // 512. full_prefill_max_req: Optional[int] = None # Only meaningful for Full prefill CUDA graphs that capture a distinct @@ -120,15 +113,10 @@ class PhaseConfig: def default_prefill_backend() -> str: - """BCG (breakable) is the prefill default on CUDA only; other platforms - (HIP/NPU/...) keep tc_piecewise until BCG is validated there. Full-graph - prefill capture is opt-in per model architecture via the declarative - registry (see _inkling_overrides in arg_groups/overrides.py), not a global - default. Lazy import keeps this module's stdlib-only import invariant (see - module docstring).""" + """Enable BCG by default on CUDA; other platforms remain opt-in.""" from sglang.srt.utils import is_cuda - return Backend.BREAKABLE if is_cuda() else Backend.TC_PIECEWISE + return Backend.BREAKABLE if is_cuda() else Backend.DISABLED def with_phase(config: "CudaGraphConfig", phase: str, **changes) -> "CudaGraphConfig": @@ -188,6 +176,12 @@ def from_dict(cls, raw: Optional[Dict[str, Dict[str, Any]]]) -> "CudaGraphConfig if phase not in Phase.ALL or not isinstance(phase_settings, dict): continue phase_cfg = getattr(cfg, phase) + if phase_settings.get("backend") == "tc_piecewise": + raise ValueError( + "tc_piecewise was removed; select breakable, full, or disabled" + ) + if "tc_compiler" in phase_settings: + raise ValueError("tc_compiler was removed with tc_piecewise") allowed = ALLOWED_KEYS_PER_PHASE[phase] for key, value in phase_settings.items(): if key in allowed: diff --git a/python/sglang/srt/model_executor/forward_context.py b/python/sglang/srt/model_executor/forward_context.py index eb84508dfe94..e6ae56d92fd3 100644 --- a/python/sglang/srt/model_executor/forward_context.py +++ b/python/sglang/srt/model_executor/forward_context.py @@ -11,9 +11,6 @@ per-stream backend, frozen-KV MTP draft loop, TBO per-child dispatch) use dataclasses.replace and wrap the override scope with forward_context(). -Distinct from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph.TcPiecewiseForwardContext, -which collects compilation-time refs for the piecewise CUDA graph backend. - Concurrency: _current is a plain module-level global, not thread-local. This matches the global_server_args precedent and is safe because each forward runs synchronously on a single Python thread per worker process. If @@ -38,6 +35,9 @@ class ForwardContext: write time — use dataclasses.replace for per-call overrides.""" attn_backend: AttentionBackend + # Runner-owned graph policy; no module registries or hidden batch lookup. + full_graph: bool = False + raw_num_tokens: Optional[int] = None _current: Optional[ForwardContext] = None @@ -82,3 +82,7 @@ def forward_context(ctx: ForwardContext): yield finally: set_forward_context(prev) + + +def is_in_full_prefill_graph() -> bool: + return _current is not None and _current.full_graph diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 604a72f4ed7f..59c40f2930a5 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -131,9 +131,7 @@ is_post_capture_kv_active, ) from sglang.srt.model_executor.model_runner_components.layer_setup import ( - AttentionAndMoeLayers, ModelLayerInfo, - compute_attention_and_moe_layers, resolve_layer_indices, ) from sglang.srt.model_executor.model_runner_components.load_model_utils import ( @@ -1507,10 +1505,6 @@ def _decode_cuda_graph_runner_cls(self): return DecodeCudaGraphRunner - def get_cuda_graph_layers(self, layer_model) -> AttentionAndMoeLayers: - """Return the model layers used by prefill CUDA graph execution.""" - return compute_attention_and_moe_layers(layer_model) - def init_decode_cuda_graph(self): self.decode_cuda_graph_runner = None capture = capture_decode_graph(model_runner=self) diff --git a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py index 5151de063c25..316ce7bda7c3 100644 --- a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py @@ -35,6 +35,9 @@ ) from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput from sglang.srt.model_executor.hook_manager import register_forward_hooks +from sglang.srt.model_executor.model_runner_components.layer_setup import ( + compute_attention_layer_info, +) from sglang.srt.model_executor.runner import ( EagerRunner, PrefillCudaGraphRunner, @@ -63,31 +66,6 @@ _deep_gemm_layout_memory_budget_initialized = False -def _align_pipeline_layers(layers: list, layer_model) -> list: - has_start_layer = hasattr(layer_model, "start_layer") - has_end_layer = hasattr(layer_model, "end_layer") - assert has_start_layer == has_end_layer, ( - "pipeline layer ranges must define start_layer and end_layer together" - ) - start_layer = layer_model.start_layer if has_start_layer else 0 - end_layer = layer_model.end_layer if has_end_layer else len(layer_model.layers) - assert isinstance(start_layer, int) and isinstance(end_layer, int), ( - "pipeline layer ranges must define integer start_layer and end_layer" - ) - assert 0 <= start_layer <= end_layer <= len(layer_model.layers), ( - f"invalid pipeline layer range [{start_layer}, {end_layer}) for " - f"{len(layer_model.layers)} layers" - ) - if len(layers) == len(layer_model.layers): - return layers - assert len(layers) <= end_layer - start_layer, ( - f"found {len(layers)} layers in PP range [{start_layer}, {end_layer})" - ) - return ( - [None] * start_layer + layers + [None] * (len(layer_model.layers) - end_layer) - ) - - def has_standard_gqa_for_all_local_layers( *, attention_layer_count: int, start_layer: int, end_layer: int ) -> bool: @@ -95,54 +73,6 @@ def has_standard_gqa_for_all_local_layers( return attention_layer_count >= end_layer - start_layer -def index_attention_layers_by_global_id( - attention_layers: list[Any], - mha_companion_layers: list[Any], - layer_model=None, -) -> tuple[list[Any], list[Any]]: - """Pad PP-local attention metadata so global layer_id remains a valid index. - - Models that re-execute layers pre-expand these lists into position-indexed - lookup tables (the same layer at several positions); such tables are - returned unchanged. - """ - if len(attention_layers) != len(mha_companion_layers): - raise ValueError("attention and MHA companion metadata must be parallel") - populated = [layer for layer in attention_layers if layer is not None] - if not populated or any(not hasattr(layer, "layer_id") for layer in populated): - if layer_model is not None: - return ( - _align_pipeline_layers(attention_layers, layer_model), - _align_pipeline_layers(mha_companion_layers, layer_model), - ) - return attention_layers, mha_companion_layers - max_layer_id = max(int(layer.layer_id) for layer in populated) - indexed_attention = [None] * (max_layer_id + 1) - indexed_companions = [None] * (max_layer_id + 1) - has_reused_layers = False - for attention, companion in zip(attention_layers, mha_companion_layers): - if attention is None: - if companion is not None: - raise ValueError("MHA companion has no primary attention layer") - continue - layer_id = int(attention.layer_id) - if layer_id < 0: - raise ValueError(f"invalid or duplicate attention layer_id: {layer_id}") - if indexed_attention[layer_id] is not None: - if ( - indexed_attention[layer_id] is not attention - or indexed_companions[layer_id] is not companion - ): - raise ValueError(f"invalid or duplicate attention layer_id: {layer_id}") - has_reused_layers = True - continue - indexed_attention[layer_id] = attention - indexed_companions[layer_id] = companion - if has_reused_layers: - return attention_layers, mha_companion_layers - return indexed_attention, indexed_companions - - class GraphCapture(msgspec.Struct, frozen=True, kw_only=True): runner: Optional[BaseRunner] memory_phase: str @@ -490,24 +420,6 @@ def result( if model_runner.is_draft_worker and not force_for_draft_worker: return result(None) - # Skip prefill CG for EAGLE target on tc_piecewise when the fixed server - # capture ceiling is below FULL. EAGLE target prefill requests FULL, so a - # NULL or LAST graph is dead; capturing it can perturb FP4/TRTLLM-MoE - # state and corrupt decode replay (see #28386 and #28870). BCG and FullCG - # capture FULL for EAGLE targets in PrefillCudaGraphRunner.__init__, so - # they do not need this skip. - if ( - model_runner.spec_algorithm.is_eagle() - and not model_runner.is_draft_worker - and get_server_return_hidden_states_mode() < CaptureHiddenMode.FULL - and check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - ): - logger.info( - "Disable prefill CUDA graph for EAGLE target on tc_piecewise " - "to avoid FP4/MoE decode-replay corruption (#28386)." - ) - return result(eager_runner) - if ( model_runner.lora_manager is not None and not model_runner.lora_manager.supports_prefill_cuda_graph @@ -608,26 +520,12 @@ def result( ) return result(None) - ( - model_runner.attention_layers, - model_runner.moe_layers, - model_runner.moe_fusions, - model_runner.dsa_indexers, - model_runner.mha_companion_layers, - ) = model_runner.get_cuda_graph_layers(layer_model) - ( - model_runner.attention_layers, - model_runner.mha_companion_layers, - ) = index_attention_layers_by_global_id( - model_runner.attention_layers, - model_runner.mha_companion_layers, - layer_model, + attention_layer_count, model_runner.has_mha_companion_layers = ( + compute_attention_layer_info(layer_model) ) if not has_standard_gqa_for_all_local_layers( - attention_layer_count=sum( - layer is not None for layer in model_runner.attention_layers - ), + attention_layer_count=attention_layer_count, start_layer=model_runner.layer_info.start_layer, end_layer=model_runner.layer_info.end_layer, ): diff --git a/python/sglang/srt/model_executor/model_runner_components/layer_setup.py b/python/sglang/srt/model_executor/model_runner_components/layer_setup.py index 01d84532873b..e3a7a43788bf 100644 --- a/python/sglang/srt/model_executor/model_runner_components/layer_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/layer_setup.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, NamedTuple, Optional +from typing import TYPE_CHECKING, Any, Optional import msgspec from torch import nn @@ -9,30 +9,15 @@ from sglang.srt.configs.model_config import ModelConfig -class AttentionAndMoeLayers(NamedTuple): - attention_layers: list[Any] - moe_layers: list[Any] - moe_fusions: list[Any] - dsa_indexers: list[Any] - mha_companion_layers: list[Any] - - def _get_loop_num(hf_config: Any) -> int: # Nanbeige uses num_loops; IQuestLoopCoder uses loop_num. return int(getattr(hf_config, "loop_num", getattr(hf_config, "num_loops", 1)) or 1) -def compute_attention_and_moe_layers(layer_model: Any) -> AttentionAndMoeLayers: - attention_layers: list[Any] = [] - moe_layers: list[Any] = [] - moe_fusions: list[Any] = [] - dsa_indexers: list[Any] = [] - mha_companion_layers: list[Any] = [] - - # Loop models (Nanbeige / IQuestLoopCoder) store one RadixAttention per loop - # in a ModuleList. Prefill CUDA graph indexes by layer_id, so expand and - # reorder to a dense [0..N) list. - has_loop_attn = False +def compute_attention_layer_info(layer_model: Any) -> tuple[int, bool]: + """Count supported attention layers and detect MHA companions for graph gates.""" + attention_layer_count = 0 + has_mha_companion_layers = False layers = layer_model.layers if isinstance(layers, nn.ModuleDict): @@ -69,56 +54,19 @@ def compute_attention_and_moe_layers(layer_model: Any) -> AttentionAndMoeLayers: if hasattr(layer.mixer, "attn"): attn_layer = layer.mixer.attn elif hasattr(layer, "_forward_mamba"): - # Mamba layer with split op support - store the layer itself + # Mamba layer with graph support attn_layer = layer if isinstance(attn_layer, nn.ModuleList): - attention_layers.extend(attn_layer) - mha_companion_layers.extend([mha_companion_layer] * len(attn_layer)) - has_loop_attn = True - else: - # Keep these lists aligned with global layer ids. Pipeline-parallel - # models retain placeholders outside the local stage, while real - # attention modules use their global layer_id during graph replay. - attention_layers.append(attn_layer) - mha_companion_layers.append(mha_companion_layer) - - moe_block = None - moe_fusion = None - if hasattr(layer, "mlp") and hasattr(layer.mlp, "experts"): - moe_block = layer.mlp.experts - moe_fusion = layer.mlp - if hasattr(layer, "block_sparse_moe") and hasattr( - layer.block_sparse_moe, "experts" - ): - moe_block = layer.block_sparse_moe.experts - moe_fusion = layer.block_sparse_moe - if hasattr(layer, "moe") and hasattr(layer.moe, "experts"): - moe_block = layer.moe.experts - moe_fusion = layer.moe - # For NemotronH MoE layers using 'mixer' attribute - if hasattr(layer, "mixer") and hasattr(layer.mixer, "experts"): - moe_block = layer.mixer.experts - moe_fusion = layer.mixer - moe_layers.append(moe_block) - moe_fusions.append(moe_fusion) - # NSA indexers (None for layers without NSA) - dsa_indexer = None - if hasattr(layer, "self_attn") and hasattr(layer.self_attn, "indexer"): - dsa_indexer = layer.self_attn.indexer - dsa_indexers.append(dsa_indexer) - - # Reorder so attention_layers[i] matches RadixAttention.layer_id. - if has_loop_attn: - attention_layers.sort(key=lambda x: x.layer_id) - - return AttentionAndMoeLayers( - attention_layers, - moe_layers, - moe_fusions, - dsa_indexers, - mha_companion_layers, - ) + # Loop models have one attention module for each execution of a block. + attention_layer_count += sum(attn is not None for attn in attn_layer) + if len(attn_layer) and mha_companion_layer is not None: + has_mha_companion_layers = True + elif attn_layer is not None: + attention_layer_count += 1 + has_mha_companion_layers |= mha_companion_layer is not None + + return attention_layer_count, has_mha_companion_layers class _PPLayerRange(msgspec.Struct, frozen=True, kw_only=True): diff --git a/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py b/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py index 856281b73a63..a1454ab4f22f 100644 --- a/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py +++ b/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py @@ -340,10 +340,6 @@ def _get_unsupported_reason( basic_rules = ( (not options.is_cuda_platform or options.device != "cuda", "CUDA only"), (not options.cuda_graph_enabled, "CUDA graph capture is disabled"), - ( - options.prefill_cuda_graph_backend == Backend.TC_PIECEWISE, - "tc_piecewise prefill CUDA graphs are not supported", - ), (type(loader) is not DefaultModelLoader, "DefaultModelLoader only"), ( load_config.load_format diff --git a/python/sglang/srt/model_executor/runner/__init__.py b/python/sglang/srt/model_executor/runner/__init__.py index 9aa65fbf579b..dc2ac23e6b16 100644 --- a/python/sglang/srt/model_executor/runner/__init__.py +++ b/python/sglang/srt/model_executor/runner/__init__.py @@ -30,7 +30,7 @@ get_batch_sizes_to_capture, ) from sglang.srt.model_executor.runner.base_runner import BaseRunner # noqa: F401 -from sglang.srt.model_executor.runner.decode_cuda_graph_runner import ( +from sglang.srt.model_executor.runner.decode_cuda_graph_runner import ( # noqa: F401 DecodeCudaGraphRunner, ) from sglang.srt.model_executor.runner.eager_runner import EagerRunner # noqa: F401 @@ -38,9 +38,6 @@ PrefillCudaGraphRunner, ) from sglang.srt.model_executor.runner.shape_key import ShapeKey # noqa: F401 -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( # noqa: F401 - TCPCG_FAILURE_HINT, -) from sglang.srt.model_executor.runner_utils import ( # noqa: F401 DecodeInputBuffers, DeepEPCudaGraphRunnerAdapter, diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index a7b73754114c..4b87dff84dc5 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -19,8 +19,6 @@ torch.cuda.CUDAGraph per shape. - "breakable" — experimental, BreakableCudaGraphBackend: segmented capture (no torch.compile). - - "tc_piecewise" — not implemented for decode; logs a one-shot warning - and falls back to "full". """ from __future__ import annotations diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index acb438eb5770..b3afb64569e9 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -50,10 +50,6 @@ get_token_to_kv_pool, ) from sglang.srt.model_executor.runner.base_runner import BaseRunner -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - enable_tc_piecewise_cuda_graph, - set_tc_piecewise_forward_context, -) from sglang.srt.model_executor.runner_utils import ( maybe_publish_prefill_shared_read_done, ) @@ -338,34 +334,7 @@ def _execute_extend( else "extend" ) with device_timer_ctx(model_runner.device_timer, category): - pcg_runner = model_runner.prefill_cuda_graph_runner - if ( - _is_hip - and pcg_runner is not None - and not isinstance(pcg_runner, EagerRunner) - and not cp_active - ): - # HIP PCG eager fallback: enter the PCG context so Dynamo guards - # and PCG-specific MoE/attention paths stay consistent. - with ( - enable_tc_piecewise_cuda_graph(), - set_tc_piecewise_forward_context( - forward_batch, - model_runner.attention_layers, - getattr(model_runner.model, "quant_config", None), - model_runner.moe_layers, - model_runner.moe_fusions, - dsa_indexers=model_runner.dsa_indexers, - mha_companion_layers=model_runner.mha_companion_layers, - ), - ): - ret = model_runner.model.forward( - forward_batch.input_ids, - forward_batch.positions, - forward_batch, - **kwargs, - ) - elif cp_active: + if cp_active: ret = self._execute_extend_cp(forward_batch, kwargs) else: ret = model_runner.model.forward( diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index 1b506b826c7e..765a5c230d51 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -28,11 +28,6 @@ zero-length sentinels. bs > slots falls back to eager. Attention metadata is refreshed out-of-graph against the slot-padded batch before capture/replay. - - "tc_piecewise" — TcPiecewiseCudaGraphBackend: torch.compile - wraps the model; per-shape compiled/captured pieces live - in torch.compile's internal cache. Multi-request prefill - is supported. - - "disabled" — handled at the model_runner level; runner not constructed. """ from __future__ import annotations @@ -112,10 +107,6 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( BCG_FAILURE_HINT, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - TCPCG_FAILURE_HINT, - set_tc_piecewise_forward_context, -) from sglang.srt.model_executor.runner_utils import ( maybe_publish_prefill_shared_read_done, ) @@ -246,11 +237,9 @@ class _ChunkedPrefixCaptureBuffers: def prefill_failure_msg(backend_name: str) -> str: """Render PREFILL_CUDA_GRAPH_CAPTURE_FAILED_MSG with a backend-specific numbered suggestion list. The runner is only constructed for BREAKABLE - or TC_PIECEWISE; other values fall back to a generic OOM-style list.""" + ; other values fall back to a generic OOM-style list.""" if backend_name == Backend.BREAKABLE: hint = BCG_FAILURE_HINT - elif backend_name == Backend.TC_PIECEWISE: - hint = TCPCG_FAILURE_HINT else: hint = ( "1. disable the prefill CUDA graph by --cuda-graph-backend-prefill=disabled\n" @@ -319,7 +308,7 @@ def __init__(self, model_runner: ModelRunner): ) ) # bs in prefill carries the captured shape (token count for - # tc_piecewise) — one shape knob per phase. + # full) — one shape knob per phase. capture_tokens = prefill_config.bs assert capture_tokens is not None, "cuda_graph_config[prefill].bs is not set" self.capture_num_tokens = sorted(capture_tokens) @@ -333,13 +322,6 @@ def __init__(self, model_runner: ModelRunner): max_context_size=self.max_context_size, table_width=model_runner.req_to_token_pool.req_to_token.shape[1], ) - if ( - self.prefill_backend_name == Backend.TC_PIECEWISE - and self.max_context_size is not None - ): - # TODO(SYChen123): Plumb max_seq_len_override through TcPiecewise - # metadata preparation before enabling the fixed context limit here. - self._ignore_max_context_size("tc_piecewise prefill CUDA graph") # --- capture modes -------------------------------------------- self.capture_forward_mode = ForwardMode.EXTEND @@ -418,47 +400,23 @@ def __init__(self, model_runner: ModelRunner): source=self.buffers, ) - self.attention_layers = self.model_runner.attention_layers - self.mha_companion_layers = self.model_runner.mha_companion_layers - self.has_mha_companion_layers = any( - layer is not None for layer in self.mha_companion_layers - ) - self.moe_layers = self.model_runner.moe_layers - self.moe_fusions = self.model_runner.moe_fusions - self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None) + self.has_mha_companion_layers = self.model_runner.has_mha_companion_layers self.dp_size = get_parallel().dp_size self.require_mlp_tp_gather = require_mlp_tp_gather() self.require_attn_tp_gather = require_attn_tp_gather() # --- backend --------------------------------------------------- - # TcPiecewise resolves by running a compile pass that calls back into - # capture_prepare / _run_forward, so these fields must exist first. self._prefill_static_buffers: Optional[Dict[str, torch.Tensor]] = None self.static_draft_hidden_states: Optional[torch.Tensor] = None self.layer_model = None self._capture_req_slots = 1 - # Same rationale: _run_compile_pass runs a dummy _run_forward before - # resolve_prefill_backend returns, and that forward reads - # self._is_full_backend. The compile-pass backend is never Full, so - # default False; the assignment below sets the real value once the - # backend type is known. self._is_full_backend = False # Same ordering requirement: capture_prepare reads this. self._capture_lora = False self.enable_cp_bcg_capture = False self.prefill_cp_bcg_input: Optional[PrefillCPBCGInput] = None - # TcPiecewise does its compile pass during backend construction. - # Wrap only that path with the prefill CUDA graph failure hint. - try: - self.backend = resolve_prefill_backend(self) - except RuntimeError as e: - if self.prefill_backend_name != Backend.TC_PIECEWISE: - raise - raise RuntimeError( - f"Capture prefill CUDA graph failed: {e}\n" - f"{prefill_failure_msg(self.prefill_backend_name)}" - ) from e + self.backend = resolve_prefill_backend(self) self._is_full_backend = isinstance(self.backend, FullCudaGraphBackend) if self._is_full_backend: @@ -580,7 +538,7 @@ def __init__(self, model_runner: ModelRunner): # contract under BCG: capture-time builds a per-bucket metadata # object the backend then refreshes in place at replay. We honor # the contract only when the backend is Breakable; FullCG and - # TC_PIECEWISE use the eager init_forward_metadata path. + # Breakable uses the eager init_forward_metadata path. if isinstance(self.backend, BreakableCudaGraphBackend): self.use_captured_attn_metadata = model_runner.attn_backend.use_captured_forward_metadata_for_breakable_cuda_graph else: @@ -661,15 +619,9 @@ def _next_token_logits_buffer(self, rows: int) -> Optional[torch.Tensor]: self.model_runner.model_config.vocab_size, rows=rows ) - def _uses_eager_prefill_tail(self) -> bool: - return self.prefill_backend_name in (Backend.BREAKABLE, Backend.FULL) - def _prefill_logits_buffer_rows(self, forward_batch: ForwardBatch) -> int: if not forward_batch.return_logprob: return forward_batch.batch_size - assert self._uses_eager_prefill_tail(), ( - "Prefill return_logprob requires an eager logits tail." - ) global_num_tokens = forward_batch.global_num_tokens_for_logprob_cpu if global_num_tokens is not None: @@ -725,42 +677,27 @@ def _capture_pp_proxy_tensors(self, num_tokens: int) -> Optional[PPProxyTensors] @contextmanager def _prefill_forward_context( self, - forward_batch: ForwardBatch, *, - num_tokens: Optional[int] = None, raw_num_tokens: Optional[int] = None, ): - with ( - forward_context( - ForwardContext(attn_backend=self.model_runner.attn_backend) - ), - set_tc_piecewise_forward_context( - forward_batch, - self.attention_layers, - self.quant_config, - self.moe_layers, - self.moe_fusions, - dsa_indexers=self.dsa_indexers, - mha_companion_layers=self.mha_companion_layers, - num_tokens=num_tokens, - raw_num_tokens=raw_num_tokens, + with forward_context( + ForwardContext( + attn_backend=self.model_runner.attn_backend, full_graph=self._is_full_backend, - ), + raw_num_tokens=raw_num_tokens, + ) ): yield @torch.no_grad() def _run_forward(self, forward_batch: ForwardBatch, num_tokens: int): - """Run forward inside the prefill set_tc_piecewise_forward_context. + """Capture the transformer body with prefill graph policy. BCG path: captures only the inner layer_model.forward (transformer stack), excluding the outer model.forward tail (logits_processor / pooler). The captured output is bs=1 hidden states; replay then runs the outer tail eagerly with live multi-req metadata. - TC_PIECEWISE path: captures the outer model.forward; torch.compile - FX-traces produce bs-invariant kernels. - ``@torch.no_grad`` mirrors the decorator on the outer ``*ForCausalLM.forward``. For BCG, calling ``layer_model.forward`` directly skips that decorator, so we apply it here — without it @@ -777,111 +714,26 @@ def _run_forward(self, forward_batch: ForwardBatch, num_tokens: int): ) set_is_extend_in_batch(False) - with self._prefill_forward_context(forward_batch): + with self._prefill_forward_context(): pp_proxy_tensors = self._capture_pp_proxy_tensors(num_tokens) - if self._uses_eager_prefill_tail(): - # BCG / Full: capture the transformer body only. - positions = self._get_layer_model_positions(forward_batch) - input_ids = forward_batch.input_ids - kwargs = _build_layer_model_forward_kwargs( - self.layer_model, forward_batch, pp_proxy_tensors - ) - if pp_proxy_tensors is not None: - input_ids = None - for embeds_name in ("input_embeds", "inputs_embeds"): - if embeds_name in kwargs: - kwargs[embeds_name] = None - break - return self.layer_model.forward( - input_ids, - positions, - forward_batch, - **kwargs, - ) - # tc_piecewise: compile/capture the outer model.forward path. - pp_kwargs = self.model_runner._pp_kwargs(pp_proxy_tensors) - return self.model_runner.model.forward( - forward_batch.input_ids, - forward_batch.positions, - forward_batch, - **pp_kwargs, + # BCG / Full: capture the transformer body only. + positions = self._get_layer_model_positions(forward_batch) + input_ids = forward_batch.input_ids + kwargs = _build_layer_model_forward_kwargs( + self.layer_model, forward_batch, pp_proxy_tensors ) - - def _run_dummy_forward(self, num_tokens: int) -> None: - """Build a dummy ForwardBatch at this shape, init attn metadata, - run forward once. Used by TcPiecewiseCudaGraphBackend.prepare - for both the JIT-activate forward (single shape, before - torch.compile install) and the compile-loop pass (every shape, - inside enable_torch_compile_warmup). - """ - fb, attn_backend = self.capture_prepare(num_tokens) - attn_backend.init_forward_metadata(fb) - self._run_forward(fb, num_tokens) - - def run_dummy_multimodal_deepstack_forward( - self, language_model: torch.nn.Module, num_tokens: int - ) -> bool: - """Warm the tensor-valued deepstack branch before serving requests. - - The regular PCG dummy is text-only. Qwen3-VL only provides - ``input_deepstack_embeds`` after visual encoding, so leaving this - branch cold makes the first image request synchronously recompile the - language model. The model/signature checks keep this a no-op for - non-deepstack architectures. - """ - if ( - "input_deepstack_embeds" - not in inspect.signature(language_model.forward).parameters - ): - return False - - num_deepstack = getattr(self.model_runner.model, "num_deepstack_embeddings", 0) - if num_deepstack <= 0: - return False - - hidden_size = ( - getattr(getattr(language_model, "config", None), "hidden_size", None) - or self.model_runner.model_config.hidden_size - ) - fb, attn_backend = self.capture_prepare(num_tokens) - attn_backend.init_forward_metadata(fb) - deepstack_embeds = torch.zeros( - (num_tokens, hidden_size * num_deepstack), - dtype=self.model_runner.dtype, - device=self.device, - ) - torch._dynamo.maybe_mark_dynamic(deepstack_embeds, 0) - - fb.dp_local_start_pos = fb.dp_local_num_tokens = None - set_dp_buffer_len( - fb.global_dp_buffer_len, - num_tokens, - fb.dp_padding_mode.is_max_len(), - fb.global_num_tokens_cpu, - ) - set_is_extend_in_batch(False) - - with ( - forward_context( - ForwardContext(attn_backend=self.model_runner.attn_backend) - ), - set_tc_piecewise_forward_context( - fb, - self.attention_layers, - self.quant_config, - self.moe_layers, - self.moe_fusions, - dsa_indexers=self.dsa_indexers, - ), - ): - language_model.forward( - fb.input_ids, - self._get_layer_model_positions(fb), - fb, - input_embeds=fb.input_embeds, - input_deepstack_embeds=deepstack_embeds, + if pp_proxy_tensors is not None: + input_ids = None + for embeds_name in ("input_embeds", "inputs_embeds"): + if embeds_name in kwargs: + kwargs[embeds_name] = None + break + return self.layer_model.forward( + input_ids, + positions, + forward_batch, + **kwargs, ) - return True def _has_inactive_dp_rank(self, forward_batch: ForwardBatch) -> bool: # DSV4 DP attention / DeepEP collectives need every DP rank to enter @@ -1142,7 +994,7 @@ def _init_forward_metadata_for_capture( """Capture-time metadata init for the BCG-with-captured-metadata contract. For opt-in backends (DSV4), call the BCG-specific entry and stash the returned per-bucket metadata object; otherwise fall - back to the generic eager init that BCG/TC_PIECEWISE use today.""" + back to the generic eager init that BCG use today.""" attn_backend = self.model_runner.attn_backend with forward_context(ForwardContext(attn_backend=attn_backend)): if not self.use_captured_attn_metadata: @@ -1263,7 +1115,6 @@ def can_replay_locally( # flag is FullCG-only, so this is inert for the BreakableCG vote path. if self._has_uncapturable_chunked_prefix(prefix_lens): return False - # tc_piecewise captures with ForwardMode.EXTEND and spec_info=None. if is_target_verify: return False if ( @@ -1271,8 +1122,6 @@ def can_replay_locally( and self.capture_hidden_mode < capture_hidden_mode ): return False - if return_logprob and not self._uses_eager_prefill_tail(): - return False if self.max_context_size is not None: if ( batch_max_context_len is None @@ -2013,8 +1862,6 @@ def replay_layer_forward(*args, **layer_kwargs): tail_batch.mm_input_embeds = forward_batch.mm_input_embeds try: with self._prefill_forward_context( - static_forward_batch, - num_tokens=static_num_tokens, raw_num_tokens=raw_num_tokens, ): return self.model_runner.model.forward( @@ -2026,27 +1873,6 @@ def replay_layer_forward(*args, **layer_kwargs): finally: self.layer_model.forward = original_layer_forward - def _execute_tc_piecewise( - self, - static_forward_batch: ForwardBatch, - static_num_tokens: int, - raw_num_tokens: int, - **kwargs, - ): - assert self.max_context_size is None, ( - "tc_piecewise replay does not support a fixed prefill context size" - ) - with self._prefill_forward_context( - static_forward_batch, - num_tokens=static_num_tokens, - raw_num_tokens=raw_num_tokens, - ): - return self.backend.replay( - ShapeKey(size=static_num_tokens), - static_forward_batch, - **kwargs, - ) - def _trim_logits_output( self, output: LogitsProcessorOutput ) -> LogitsProcessorOutput: @@ -2126,7 +1952,7 @@ def execute( raw_num_tokens, **kwargs, ) - elif self._uses_eager_prefill_tail(): + else: output = self._execute_body_capture( forward_batch, static_forward_batch, @@ -2135,11 +1961,4 @@ def execute( shape_key, **kwargs, ) - else: - output = self._execute_tc_piecewise( - static_forward_batch, - static_num_tokens, - raw_num_tokens, - **kwargs, - ) return self._finalize_execute_output(output) diff --git a/python/sglang/srt/model_executor/runner_backend/__init__.py b/python/sglang/srt/model_executor/runner_backend/__init__.py index b50f0205352a..4a288ea40eb3 100644 --- a/python/sglang/srt/model_executor/runner_backend/__init__.py +++ b/python/sglang/srt/model_executor/runner_backend/__init__.py @@ -9,8 +9,6 @@ - FullCudaGraphBackend — single torch.cuda.CUDAGraph per shape. - BreakableCudaGraphBackend — segmented capture with eager break markers; no torch.compile. - - TcPiecewiseCudaGraphBackend — torch.compile-driven piecewise - capture; FX-splits the model at attention layers. """ from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import ( # noqa: F401 @@ -22,9 +20,6 @@ from sglang.srt.model_executor.runner_backend.full_cuda_graph_backend import ( # noqa: F401 FullCudaGraphBackend, ) -from sglang.srt.model_executor.runner_backend.tc_piecewise_cuda_graph_backend import ( # noqa: F401 - TcPiecewiseCudaGraphBackend, -) from sglang.srt.model_executor.runner_backend.utils import ( # noqa: F401 resolve_decode_backend, resolve_prefill_backend, diff --git a/python/sglang/srt/model_executor/runner_backend/breakable_cuda_graph_backend.py b/python/sglang/srt/model_executor/runner_backend/breakable_cuda_graph_backend.py index b36c6373d253..8e3c6329b798 100644 --- a/python/sglang/srt/model_executor/runner_backend/breakable_cuda_graph_backend.py +++ b/python/sglang/srt/model_executor/runner_backend/breakable_cuda_graph_backend.py @@ -126,9 +126,7 @@ def capture_one( post_warmup_hook() graph = BreakableCUDAGraph(self.deduped_cuda_graph) - captured_fn = ( - eager_on_graph(True)(forward_fn) if self._debug_eager else forward_fn - ) + captured_fn = eager_on_graph(forward_fn) if self._debug_eager else forward_fn size = shape_key.size if self._shared_output_buffer is None: capacity_rows = self._cuda_graph_runner.cuda_graph_output_capacity_rows( @@ -266,7 +264,7 @@ def replay( **kwargs, ) -> Any: with graph_pool_replay_scope(): - self._graphs[shape_key].replay() + self._graphs[shape_key].replay(static_forward_batch) return self._outputs[shape_key] def cleanup(self) -> None: diff --git a/python/sglang/srt/model_executor/runner_backend/utils.py b/python/sglang/srt/model_executor/runner_backend/utils.py index cb9449f417b0..c9a50be993ec 100644 --- a/python/sglang/srt/model_executor/runner_backend/utils.py +++ b/python/sglang/srt/model_executor/runner_backend/utils.py @@ -34,9 +34,6 @@ from sglang.srt.model_executor.runner_backend.full_cuda_graph_backend import ( FullCudaGraphBackend, ) -from sglang.srt.model_executor.runner_backend.tc_piecewise_cuda_graph_backend import ( - TcPiecewiseCudaGraphBackend, -) from sglang.srt.platforms import current_platform from sglang.srt.runtime_context import get_exec @@ -47,9 +44,6 @@ logger = logging.getLogger(__name__) -# Track first occurrence of each fallback warning to avoid log spam. -_TC_PIECEWISE_DECODE_FALLBACK_LOGGED = False - def resolve_decode_backend( cuda_graph_runner: BaseCudaGraphRunner, @@ -90,14 +84,9 @@ def resolve_decode_backend( enable_memory_saver=enable_memory_saver, debug_eager=get_exec().graph.debug_cuda_graph, ) - if backend_name == Backend.TC_PIECEWISE: - global _TC_PIECEWISE_DECODE_FALLBACK_LOGGED - if not _TC_PIECEWISE_DECODE_FALLBACK_LOGGED: - logger.warning( - "cuda_graph_config decode='tc_piecewise' is not yet implemented; " - "falling back to 'full'." - ) - _TC_PIECEWISE_DECODE_FALLBACK_LOGGED = True + + if backend_name != Backend.FULL: + raise ValueError(f"Unsupported decode graph backend: {backend_name}") full_backend_cls = None if current_platform.is_out_of_tree(): @@ -116,9 +105,8 @@ def resolve_prefill_backend( cuda_graph_runner: BaseCudaGraphRunner, ) -> BaseCudaGraphBackend: """Pick a backend instance from cuda_graph_config['prefill']['backend'].""" - model_runner = cuda_graph_runner.model_runner cfg = get_exec().graph.cuda_graph_config - backend_name = cfg.prefill.backend if cfg is not None else Backend.TC_PIECEWISE + backend_name = cfg.prefill.backend if cfg is not None else Backend.BREAKABLE if backend_name == Backend.BREAKABLE: return BreakableCudaGraphBackend( @@ -132,5 +120,4 @@ def resolve_prefill_backend( enable_memory_saver=get_exec().features.enable_memory_saver, reuse_output_buffer=True, ) - # Default: tc_piecewise. - return TcPiecewiseCudaGraphBackend(cuda_graph_runner) + raise ValueError(f"Unsupported prefill graph backend: {backend_name}") diff --git a/python/sglang/srt/model_executor/runner_backend_utils/__init__.py b/python/sglang/srt/model_executor/runner_backend_utils/__init__.py index d9d24bc49812..ee30e9168237 100644 --- a/python/sglang/srt/model_executor/runner_backend_utils/__init__.py +++ b/python/sglang/srt/model_executor/runner_backend_utils/__init__.py @@ -3,8 +3,7 @@ Subpackages: - breakable_cuda_graph: BreakableCUDAGraph + capture context, eager_on_graph decorator, is_in_breakable_cuda_graph flag. - - piecewise_cuda_graph: shared piecewise context manager - (set_tc_piecewise_forward_context, is_in_tc_piecewise_cuda_graph). + Backends in cuda_graph_backend/ import from here. Runners do not. """ diff --git a/python/sglang/srt/model_executor/runner_backend_utils/breakable_cuda_graph/breakable_cuda_graph.py b/python/sglang/srt/model_executor/runner_backend_utils/breakable_cuda_graph/breakable_cuda_graph.py index 0f4847ca218a..6c9320b1e1c5 100644 --- a/python/sglang/srt/model_executor/runner_backend_utils/breakable_cuda_graph/breakable_cuda_graph.py +++ b/python/sglang/srt/model_executor/runner_backend_utils/breakable_cuda_graph/breakable_cuda_graph.py @@ -22,6 +22,8 @@ buffers to keep break-point tensors at stable addresses. """ +import functools +import inspect import threading import warnings from contextvars import ContextVar @@ -166,7 +168,9 @@ def _weak_ref_if_tensor(x): if torch.is_tensor(x): if x.numel() == 0 or x.device.type == "cpu": return x - from sglang.srt.compilation.weak_ref_tensor import weak_ref_tensors + from sglang.srt.model_executor.runner_backend_utils.weak_ref_tensor import ( + weak_ref_tensors, + ) return weak_ref_tensors(x) if isinstance(x, tuple): @@ -216,11 +220,21 @@ def _copy_output(dst: Any, src: Any) -> Any: return src -def eager_on_graph(enable: bool, capture_stub: Optional[Callable] = None): +def eager_on_graph( + fn: Optional[Callable] = None, *, capture_stub: Optional[Callable] = None +): + """Record an eager call between captured segments. + + A named ``forward_batch`` argument is rebound to the prepared serving batch + on replay. All other arguments retain their capture-time identity; tensors + must therefore use the static buffers owned by the graph runner. + """ + def decorator(inner: Callable): - if not enable: - return inner + signature = inspect.signature(inner) + has_batch = "forward_batch" in signature.parameters + @functools.wraps(inner) def wrapper(*args, **kwargs): capture = _current_capture_var.get() if capture is None: @@ -256,9 +270,31 @@ def wrapper(*args, **kwargs): # captured segment. Keep a strong reference so replay can safely # copy fresh eager output into that bridge buffer. captured_output = output + replay_args = ( + signature.bind(*captured_args, **captured_kwargs) if has_batch else None + ) - def replay_fn(): - new_out = captured_inner(*captured_args, **captured_kwargs) + if replay_args is not None: + # Batches own request metadata and tensors. Keep only static tensor + # operands between calls, never a captured or previous serving batch. + replay_args.arguments["forward_batch"] = None + captured_args, captured_kwargs = (), {} + + def replay_fn(forward_batch): + if has_batch: + if forward_batch is None: + raise ValueError( + "This eager region requires a replay ForwardBatch" + ) + replay_args.arguments["forward_batch"] = forward_batch + try: + new_out = captured_inner( + *replay_args.args, **replay_args.kwargs + ) + finally: + replay_args.arguments["forward_batch"] = None + else: + new_out = captured_inner(*captured_args, **captured_kwargs) return _copy_output(captured_output, new_out) capture.cuda_graph._break_fns.append(replay_fn) @@ -269,7 +305,7 @@ def replay_fn(): return wrapper - return decorator + return decorator(fn) if fn is not None else decorator class BreakableCUDAGraph: @@ -278,17 +314,17 @@ class BreakableCUDAGraph: def __init__(self, deduped_cuda_graph=None) -> None: self._segments: list[Any] = [] - self._break_fns: list[Callable[[], Any]] = [] + self._break_fns: list[Callable[[Any], Any]] = [] self._deduped_cuda_graph = deduped_cuda_graph - def replay(self) -> None: + def replay(self, forward_batch=None) -> None: stream = get_device_module().current_stream() token = _current_stream_var.set(stream) try: for i, seg in enumerate(self._segments): seg.replay() if i < len(self._break_fns): - self._break_fns[i]() + self._break_fns[i](forward_batch) finally: _current_stream_var.reset(token) @@ -408,8 +444,7 @@ def _end_current_segment(self) -> None: self._current_graph_needs_instantiate = False -@eager_on_graph(True) +@eager_on_graph def break_graph() -> None: """Insert a graph break. The @eager_on_graph decorator does the actual segment split; this function body intentionally does nothing.""" - pass diff --git a/python/sglang/srt/model_executor/runner_backend_utils/breakable_cuda_graph/context.py b/python/sglang/srt/model_executor/runner_backend_utils/breakable_cuda_graph/context.py index 16f80ee445b9..ffafbed3446b 100644 --- a/python/sglang/srt/model_executor/runner_backend_utils/breakable_cuda_graph/context.py +++ b/python/sglang/srt/model_executor/runner_backend_utils/breakable_cuda_graph/context.py @@ -52,9 +52,8 @@ def enable_breakable_cuda_graph(): BCG_FAILURE_HINT = ( - "1. change to tc_piecewise by --cuda-graph-backend-prefill=tc_piecewise\n" - "2. disable the prefill CUDA graph by --cuda-graph-backend-prefill=disabled\n" - "3. if it is an OOM problem, set --mem-fraction-static to a smaller value " + "1. disable the prefill CUDA graph by --cuda-graph-backend-prefill=disabled\n" + "2. if it is an OOM problem, set --mem-fraction-static to a smaller value " "(e.g., 0.8 or 0.7) or set --cuda-graph-max-bs-prefill to a smaller value " "(e.g., 2048)\n" ) diff --git a/python/sglang/srt/compilation/weak_ref_tensor.py b/python/sglang/srt/model_executor/runner_backend_utils/weak_ref_tensor.py similarity index 100% rename from python/sglang/srt/compilation/weak_ref_tensor.py rename to python/sglang/srt/model_executor/runner_backend_utils/weak_ref_tensor.py diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py index 3ceab0d8772d..a0fd74977f72 100644 --- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py @@ -5,9 +5,6 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods import ( AttnForwardMethod, ) @@ -106,10 +103,10 @@ def _support_mha_one_shot(attn, forward_batch, backend_name): def _handle_attention_backend(attn, forward_batch, backend_name): - # Captured prefill (tc_piecewise or breakable) must keep a single attention + # Captured prefill (full or breakable) must keep a single attention # path: pin the absorbed MLA method — MHA one-shot/chunked shapes vary with # kv-len and cannot be captured. - if is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph(): + if is_in_breakable_cuda_graph(): return AttnForwardMethod.MLA # Strategy CP gathers latent KV in the backend's absorbed MLA path; @@ -169,7 +166,7 @@ def handle_attention_fa4(attn, forward_batch): def handle_attention_trtllm_mla(attn, forward_batch): - if is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph(): + if is_in_breakable_cuda_graph(): return AttnForwardMethod.MLA sum_extend_prefix_lens = _get_sum_extend_prefix_lens(forward_batch) @@ -191,7 +188,7 @@ def handle_attention_aiter(attn, forward_batch): # During PCG/BCG capture on ROCm, aiter fp8 MLA prefill has no capture # kernels; route through the MHA path (radix_attention swaps attn_mqa for # its attn_mha companion) so capture/replay use valid head/dim metadata. - if is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph(): + if is_in_breakable_cuda_graph(): return AttnForwardMethod.MHA if forward_batch.forward_mode.is_extend_without_speculative(): if not _support_mha_one_shot(attn, forward_batch, "aiter"): @@ -237,7 +234,7 @@ def _can_use_triton_dense_fp8_prefill(attn, forward_batch) -> bool: def handle_attention_triton(attn, forward_batch): - if is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph(): + if is_in_breakable_cuda_graph(): return AttnForwardMethod.MLA # when deterministic inference is enabled, use MLA diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index 8db5a0d51925..06c67b1eb5ba 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -11,10 +11,9 @@ per_tensor_quant_mla_fp8, per_token_group_quant_mla_deep_gemm_masked_fp8, ) -from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.environ import envs from sglang.srt.layers import deep_gemm_wrapper -from sglang.srt.layers.attention.dsa.utils import is_graph_dsa_split_op_surface +from sglang.srt.layers.attention.dsa.utils import is_dsa_bcg_prefill from sglang.srt.layers.attention.dsa_backend import prepare_kv_for_attention from sglang.srt.layers.dcp import ( all_gather_kv_cache_for_mla_extend, @@ -24,7 +23,6 @@ ) from sglang.srt.layers.layer_boundary import get_attn_tp_context from sglang.srt.layers.logits_processor import get_in_autotune_dummy_run -from sglang.srt.layers.radix_attention import unified_attention_with_output from sglang.srt.lora.deepseek_mla_correction import ( apply_q_correction as apply_kv_b_lora_q_correction, ) @@ -38,6 +36,7 @@ from sglang.srt.model_executor.forward_context import ( get_attn_backend, get_token_to_kv_pool, + is_in_full_prefill_graph, ) from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( eager_on_graph, @@ -45,10 +44,6 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.models.deepseek_common.utils import ( FORWARD_ABSORB_CORE_ATTENTION_BACKENDS, _is_cpu, @@ -62,7 +57,6 @@ maybe_capture_indexer_topk, ) from sglang.srt.utils import BumpAllocator -from sglang.srt.utils.custom_op import register_custom_op logger = logging.getLogger(__name__) _SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get() @@ -159,10 +153,9 @@ def _can_fuse_bmm_into_attention( ) -> bool: if getattr(self, "_kimi_split_gguf_kv_b", False): return False - # Shared activation surface with the DSA indexer graph dispatch - # (in piecewise/breakable graph + non-speculative extend). Like the indexer - # dispatch, this fusion is on by default on that surface. - if not is_graph_dsa_split_op_surface(forward_batch): + # Like the DSA indexer eager region, this fusion is enabled for + # non-speculative CUDA BCG prefill. + if not is_dsa_bcg_prefill(forward_batch): return False if not self.use_dsa: return False @@ -253,11 +246,7 @@ def _q8kv8_born_fp8_q_backend( return None # Graph/compile surfaces run their own dispatch; the python-side # stash handshake is eager-only. - if is_graph_dsa_split_op_surface(forward_batch): - return None - if get_tc_piecewise_forward_context() is not None: - return None - if is_in_breakable_cuda_graph(): + if is_in_breakable_cuda_graph() or is_in_full_prefill_graph(): return None if get_is_capture_mode(): return None @@ -697,12 +686,7 @@ def forward_absorb_core( llama_4_scaling=llama_4_scaling, ) if fusion_plan is not None: - bmm_attention_fn = ( - bcg_mla_bmm_then_unified_attention - if is_in_breakable_cuda_graph() - else mla_bmm_then_unified_attention - ) - bmm_attention_fn( + self._eager_bmm_attention( fusion_plan.q_nope_t, self.w_kc, fusion_plan.q_nope_out_buf, @@ -710,7 +694,7 @@ def forward_absorb_core( k_nope, fusion_plan.attn_output_buf, save_kv_cache, - self.layer_id, + forward_batch, q_pe, k_pe, cos_sin_cache=extra_args.get("cos_sin_cache"), @@ -882,26 +866,18 @@ def forward_absorb_core( ) attn_bmm_output = attn_bmm_output.transpose(0, 1).flatten(1, 2) else: - if is_in_tc_piecewise_cuda_graph(): - # torch dynamo requires out= op was called where output tensor was non-contiguous - attn_bmm_output = ( - torch.bmm(attn_output.transpose(0, 1), self.w_vc) - .transpose(0, 1) - .flatten(1, 2) - ) - else: - attn_bmm_output = torch.empty( - (attn_output.shape[0], self.num_local_heads * self.v_head_dim), - dtype=attn_output.dtype, - device=attn_output.device, - ) - torch.bmm( - attn_output.transpose(0, 1), - self.w_vc, - out=attn_bmm_output.view( - -1, self.num_local_heads, self.v_head_dim - ).transpose(0, 1), - ) + attn_bmm_output = torch.empty( + (attn_output.shape[0], self.num_local_heads * self.v_head_dim), + dtype=attn_output.dtype, + device=attn_output.device, + ) + torch.bmm( + attn_output.transpose(0, 1), + self.w_vc, + out=attn_bmm_output.view( + -1, self.num_local_heads, self.v_head_dim + ).transpose(0, 1), + ) if _SGLANG_EXPERIMENTAL_LORA_OPTI: from sglang.srt.lora.trtllm_lora_temp.deepseek_mla_correction import ( kv_b_lora_v_apply, @@ -954,53 +930,37 @@ def _fuse_rope_for_trtllm_mla( and get_attn_backend().data_type == torch.float8_e4m3fn ) - -# Fuses the absorb BMM (`q_nope @ w_kc`) with `unified_attention_with_output` -# into one eager split op under both PCG and BCG. Without this, the bf16 -# fallback BMM is captured alone in its own single-kernel CUDA graph submodule, -# paying per-submodule host overhead with no fusion benefit. -# -# `q_nope_out_view` aliases `q_nope_out_buf` (transposed). The op writes -# `q_nope_out_buf` via `torch.bmm(..., out=...)` and then reads through -# `q_nope_out_view`, so the alias's storage is mutated too. Declare it in -# `mutates_args` to keep the schema honest. -@register_custom_op( - mutates_args=["q_nope_out_buf", "q_nope_out_view", "attn_output_buf"] -) -@register_split_op() -def mla_bmm_then_unified_attention( - q_nope_t: torch.Tensor, - w_kc: torch.Tensor, - q_nope_out_buf: torch.Tensor, - q_nope_out_view: torch.Tensor, - k_nope: torch.Tensor, - attn_output_buf: torch.Tensor, - save_kv_cache: bool, - layer_id: int, - q_pe: torch.Tensor, - k_pe: torch.Tensor, - cos_sin_cache: Optional[torch.Tensor] = None, - is_neox: Optional[bool] = None, - llama_4_scaling: Optional[torch.Tensor] = None, - topk_indices: Optional[torch.Tensor] = None, -) -> None: - torch.bmm(q_nope_t, w_kc, out=q_nope_out_buf) - unified_attention_with_output( - q_nope_out_view, - k_nope, - k_nope, - attn_output_buf, - save_kv_cache, - layer_id, - q_rope=q_pe, - k_rope=k_pe, - cos_sin_cache=cos_sin_cache, - is_neox=is_neox, - llama_4_scaling=llama_4_scaling, - topk_indices=topk_indices, - ) - - -bcg_mla_bmm_then_unified_attention = eager_on_graph(True)( - mla_bmm_then_unified_attention -) + @eager_on_graph + def _eager_bmm_attention( + self, + q_nope_t: torch.Tensor, + w_kc: torch.Tensor, + q_nope_out_buf: torch.Tensor, + q_nope_out_view: torch.Tensor, + k_nope: torch.Tensor, + attn_output_buf: torch.Tensor, + save_kv_cache: bool, + forward_batch: ForwardBatch, + q_pe: torch.Tensor, + k_pe: torch.Tensor, + cos_sin_cache: Optional[torch.Tensor] = None, + is_neox: Optional[bool] = None, + llama_4_scaling: Optional[torch.Tensor] = None, + topk_indices: Optional[torch.Tensor] = None, + ) -> None: + torch.bmm(q_nope_t, w_kc, out=q_nope_out_buf) + self.attn_mqa._eager_attention.__wrapped__( + self.attn_mqa, + q_nope_out_view, + k_nope, + k_nope, + attn_output_buf, + forward_batch, + save_kv_cache, + q_rope=q_pe, + k_rope=k_pe, + cos_sin_cache=cos_sin_cache, + is_neox=is_neox, + llama_4_scaling=llama_4_scaling, + topk_indices=topk_indices, + ) diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py index d6e8728c164b..8881da3777ed 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py @@ -45,9 +45,6 @@ from sglang.srt.mem_cache.hisparse_memory_pool import HiSparseDSATokenToKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_context import get_token_to_kv_pool -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mla import ( _select_local_dcp_heads_for_autotune, is_dcp_mla_decode_phase, @@ -254,27 +251,18 @@ def rocm_absorb_v_bmm( transpose_bm_in=True, dtype=torch.bfloat16, ) - elif not is_in_tc_piecewise_cuda_graph(): - # Same (batch, heads, dim) layout as the quantized paths above, so the - # post-GEMM flatten is a view. Skipped under piecewise: torch dynamo - # rejects out= with a non-contiguous output tensor. - _bmm_buf = torch.empty( - attn_output.shape[0], - attn.num_local_heads, - attn.w_vc.shape[2], - device=attn_output.device, - dtype=torch.bfloat16, - ) - torch.bmm( - attn_output.to(torch.bfloat16).transpose(0, 1), - _absorb_weight_bf16(attn.w_vc, attn.w_scale), - out=_bmm_buf.transpose(0, 1), - ) - else: - attn_bmm_output = torch.bmm( - attn_output.to(torch.bfloat16).transpose(0, 1), - _absorb_weight_bf16(attn.w_vc, attn.w_scale), - ) + _bmm_buf = torch.empty( + attn_output.shape[0], + attn.num_local_heads, + attn.w_vc.shape[2], + device=attn_output.device, + dtype=torch.bfloat16, + ) + torch.bmm( + attn_output.to(torch.bfloat16).transpose(0, 1), + _absorb_weight_bf16(attn.w_vc, attn.w_scale), + out=_bmm_buf.transpose(0, 1), + ) if _bmm_buf is not None: # _bmm_buf is already (batch, heads, dim) contiguous diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 0181eb7b91d2..0473f2d54a25 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -141,21 +141,12 @@ VocabParallelEmbedding, get_embedding_tp_kwargs, ) -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.models.deepseek_common.attention_backend_handler import ( AttentionBackendRegistry, resolve_rocm_forward_method, @@ -212,7 +203,6 @@ make_layers, use_intel_amx_backend, ) -from sglang.srt.utils.custom_op import register_custom_op if _use_aiter: from sglang.srt.layers.rocm_linear_utils import aiter_dsv3_router_gemm @@ -783,19 +773,22 @@ def __init__( fc1_n = self.shared_experts.gate_up_proj.output_size_per_partition if ( - get_platform().is_sm100 - and isinstance( - self.shared_experts.gate_up_proj.quant_method, - ModelOptFp4LinearMethod, + (get_platform().is_sm100) + and ( + isinstance( + self.shared_experts.gate_up_proj.quant_method, + ModelOptFp4LinearMethod, + ) ) - and self.shared_experts.gate_up_proj.quant_method.quant_mode == "w4a4" - and isinstance( - self.shared_experts.down_proj.quant_method, - ModelOptFp4LinearMethod, + and (self.shared_experts.gate_up_proj.quant_method.quant_mode == "w4a4") + and ( + isinstance( + self.shared_experts.down_proj.quant_method, + ModelOptFp4LinearMethod, + ) ) - and fc1_n % 128 == 0 - and self.shared_experts.swiglu_limit is None - and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) + and (fc1_n % 128 == 0) + and (self.shared_experts.swiglu_limit is None) ): self.shared_experts.gate_up_proj._interleave_for_swiglu_fusion = True self.shared_experts._enable_nvfp4_gemm_swiglu_fusion = True @@ -897,10 +890,22 @@ def get_moe_weights(self): ) ] + def _forward_moe_dual_stream_graph( + self, hidden_states, fuse_mlp_allreduce, mlp_reduce_scatter + ): + with get_forward().scoped( + fuse_mlp_allreduce=fuse_mlp_allreduce, + mlp_reduce_scatter=mlp_reduce_scatter, + flashinfer_trtllm_bypass=True, + lora_batch_layout=LoRABatchLayout.TP_GLOBAL, + defer_moe_finalize=False, + ): + return self.forward_normal_dual_stream(hidden_states) + def _can_dual_stream_graph(self, hidden_states: torch.Tensor) -> bool: return ( _enable_pcg_dsv2_dual_stream - and (is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph()) + and is_in_breakable_cuda_graph() and get_moe_runner_backend().is_flashinfer_trtllm() and self.alt_stream is not None and self.num_fused_shared_experts == 0 @@ -940,9 +945,8 @@ def forward( if not self._enable_a2a_moe: if self._can_dual_stream_graph(hidden_states): fwd = get_forward() - return dsv2_flashinfer_moe_dual_stream_graph( + return self._forward_moe_dual_stream_graph( hidden_states, - self.layer_id, fwd.fuse_mlp_allreduce, fwd.mlp_reduce_scatter, ) @@ -1795,10 +1799,6 @@ def _maybe_quant_moe_input_once( return None if not self._moe_quant_once_enabled(): return None - if is_in_tc_piecewise_cuda_graph(): - # The piecewise MoE op quantizes internally; a pre-quant here - # would be dead work. - return None from sglang.kernels.ops.quantization.fp8_kernel import ( sglang_per_token_group_quant_fp8_row_padded, ) @@ -1840,12 +1840,10 @@ def _compute_routed_mxfp8_prequant_enabled(self) -> Tuple[bool, str]: def _should_quant_routed_input_mxfp8(self, hidden_states: torch.Tensor) -> bool: return ( # Capture-only: graph-pool tensors need no record_stream. - torch.cuda.is_current_stream_capturing() - # The piecewise TC graph's MoE op drops pre_quant_input. - and not is_in_tc_piecewise_cuda_graph() - and hidden_states.shape[0] > 0 - and hidden_states.dtype == torch.bfloat16 - and self._routed_mxfp8_prequant_static_enabled + (torch.cuda.is_current_stream_capturing()) + and (hidden_states.shape[0] > 0) + and (hidden_states.dtype == torch.bfloat16) + and (self._routed_mxfp8_prequant_static_enabled) ) def op_gate(self, state): @@ -3095,11 +3093,7 @@ def forward( aux_hidden_states = AuxHiddenStatePacker(len(self.layers_to_capture)) for i in range(normal_start_layer, normal_end_layer): # NOTE: torch dynamo does not support graph break in context manager - ctx = ( - nullcontext() - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else get_global_expert_distribution_recorder().with_current_layer(i) - ) + ctx = get_global_expert_distribution_recorder().with_current_layer(i) with ctx: layer = self.layers[i] (hidden_states, topk_indices) = layer( @@ -3442,31 +3436,4 @@ class DeepseekV32ForCausalLM(DeepseekV2ForCausalLM): pass -@register_custom_op(out_shape="hidden_states") -def dsv2_flashinfer_moe_dual_stream_graph( - hidden_states: torch.Tensor, - layer_id: int, - fuse_mlp_allreduce: bool, - mlp_reduce_scatter: bool, -) -> torch.Tensor: - forward_context = get_tc_piecewise_forward_context() - assert forward_context is not None - assert forward_context.moe_fusions is not None - - moe_fusion = forward_context.moe_fusions[layer_id] - assert moe_fusion is not None - # Custom-op execution happens outside the caller's Python scope under - # torch.compile. Carry graph-varying control state as scalar operands and - # republish it for the nested MoE/linear consumers. - with get_forward().scoped( - fuse_mlp_allreduce=fuse_mlp_allreduce, - mlp_reduce_scatter=mlp_reduce_scatter, - flashinfer_trtllm_bypass=True, - lora_batch_layout=LoRABatchLayout.TP_GLOBAL, - # The op's Tensor schema cannot carry a MoeFinalizeHandoff. - defer_moe_finalize=False, - ): - return moe_fusion.forward_normal_dual_stream(hidden_states) - - EntryClass = [DeepseekV2ForCausalLM, DeepseekV3ForCausalLM, DeepseekV32ForCausalLM] diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index c28eb599fe04..9f61d285b1eb 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -43,7 +43,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( sglang_per_token_group_quant_fp8, ) -from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, @@ -127,11 +126,6 @@ ) from sglang.srt.managers.schedule_batch import MM_PAD_SHIFT_VALUE, MultimodalInputs from sglang.srt.mem_cache.memory_pool import RadixAttention -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, PPProxyTensors, @@ -150,9 +144,6 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, -) from sglang.srt.model_loader.utils import maybe_executor_submit, should_async_load from sglang.srt.model_loader.weight_utils import ( RUNAI_STREAMER_TENSOR_ATTR, @@ -198,7 +189,6 @@ log_info_on_rank0, make_layers, ) -from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.hf_transformers_utils import get_rope_config # NPU-only: bind torch_npu here so _compute_q_b / _forward_prepare can call @@ -657,89 +647,6 @@ def _apply_gguf_grouped_wo_a( from sglang.srt.model_executor.forward_batch_info import ForwardBatch -@register_custom_op(mutates_args=["output"]) -@register_split_op() -def deepseek_v4_attention_with_output( - query: torch.Tensor, - key_value: torch.Tensor, - output: torch.Tensor, - layer_id: int, - compress_ratio: int, - attn_sink: torch.Tensor, - save_kv_cache: bool, -) -> None: - context = get_tc_piecewise_forward_context() - forward_batch = context.forward_batch - attention_layers = context.attention_layers - attention_layer = attention_layers[layer_id] - real_num_tokens = forward_batch.global_num_token_non_padded_cpu - - if real_num_tokens == 0: - output.zero_() - return - - query = query[:real_num_tokens] - key_value = key_value[:real_num_tokens] - - original_out_cache_loc = forward_batch.out_cache_loc - forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens] - - attn_backend = get_attn_backend() - try: - ret = attn_backend.forward( - q=query, - k=key_value, - v=key_value, - layer=attention_layer, - forward_batch=forward_batch, - compress_ratio=compress_ratio, - attn_sink=attn_sink, - save_kv_cache=save_kv_cache, - ) - finally: - forward_batch.out_cache_loc = original_out_cache_loc - - assert output[:real_num_tokens].numel() == ret.numel(), ( - f"Output tensor element mismatch: {output[:real_num_tokens].numel()} != {ret.numel()}" - ) - - output[:real_num_tokens].view(ret.shape).copy_(ret) - output[real_num_tokens:].zero_() - return - - -bcg_deepseek_v4_attention_with_output = eager_on_graph(True)( - deepseek_v4_attention_with_output -) - - -def deepseek_v4_low_ratio_sources(layer, x, q_lora, positions) -> None: - # The compressor and prefill indexer sync with the host, like the attention. - forward_batch = get_tc_piecewise_forward_context().forward_batch - real_num_tokens = forward_batch.global_num_token_non_padded_cpu - if real_num_tokens == 0: - return - get_attn_backend().forward_low_ratio_sources( - layer=layer, - x=x[:real_num_tokens], - q_lora=q_lora[:real_num_tokens], - positions=positions[:real_num_tokens], - forward_batch=forward_batch, - ) - - -bcg_deepseek_v4_low_ratio_sources = eager_on_graph(True)(deepseek_v4_low_ratio_sources) - - -def deepseek_v4_engram_hash_ids(hasher, input_ids: torch.Tensor) -> torch.Tensor: - # The hasher reads per-request rows, so it cannot run inside the CUDA graph. - forward_batch = get_tc_piecewise_forward_context().forward_batch - return hasher(input_ids, forward_batch) - - -bcg_deepseek_v4_engram_hash_ids = eager_on_graph(True)(deepseek_v4_engram_hash_ids) - - class MqaAttentionBase(nn.Module): # Class-level default for subclasses that read it without running __init__. wo_a_fp8: bool = False @@ -2042,7 +1949,7 @@ def _forward_prepare( and is_in_breakable_cuda_graph() and not getattr(attn_backend, "low_ratio_prefill_graph", False) ): - bcg_deepseek_v4_low_ratio_sources(self, x, q_lora, positions) + self._eager_low_ratio_sources(x, q_lora, positions, forward_batch) else: attn_backend.forward_low_ratio_sources( layer=self, @@ -2345,11 +2252,11 @@ def forward( o = attn_q.new_empty( (*attn_q.shape[:-1], self.attn_mqa.v_head_dim), ) - bcg_deepseek_v4_attention_with_output( + self._eager_attention( attn_q, attn_k, o, - self.attn_mqa.layer_id, + forward_batch, self.compress_ratio, attn_sink, save_kv_cache, @@ -2581,6 +2488,66 @@ def op_attn(self, state): x_quant=state.pop("attn_x_quant"), ) + @eager_on_graph + def _eager_attention( + self, + query: torch.Tensor, + key_value: torch.Tensor, + output: torch.Tensor, + forward_batch: ForwardBatch, + compress_ratio: int, + attn_sink: torch.Tensor, + save_kv_cache: bool, + ) -> None: + real_num_tokens = forward_batch.global_num_token_non_padded_cpu + + if real_num_tokens == 0: + output.zero_() + return + + query = query[:real_num_tokens] + key_value = key_value[:real_num_tokens] + + original_out_cache_loc = forward_batch.out_cache_loc + forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens] + + attn_backend = get_attn_backend() + try: + ret = attn_backend.forward( + q=query, + k=key_value, + v=key_value, + layer=self.attn_mqa, + forward_batch=forward_batch, + compress_ratio=compress_ratio, + attn_sink=attn_sink, + save_kv_cache=save_kv_cache, + ) + finally: + forward_batch.out_cache_loc = original_out_cache_loc + + assert output[:real_num_tokens].numel() == ret.numel(), ( + f"Output tensor element mismatch: {output[:real_num_tokens].numel()} != {ret.numel()}" + ) + + output[:real_num_tokens].view(ret.shape).copy_(ret) + output[real_num_tokens:].zero_() + return + + @eager_on_graph + def _eager_low_ratio_sources(self, x, q_lora, positions, forward_batch) -> None: + # The compressor and prefill indexer sync with the host, like the attention. + real_num_tokens = forward_batch.global_num_token_non_padded_cpu + if real_num_tokens == 0: + return + get_attn_backend().forward_low_ratio_sources( + layer=self, + x=x[:real_num_tokens], + q_lora=q_lora[:real_num_tokens], + positions=positions[:real_num_tokens], + forward_batch=forward_batch, + ) + @contextmanager def _every_row_routed(forward_batch: ForwardBatch, num_rows: int): @@ -4415,9 +4382,7 @@ def _forward_layers_hc_pre_from_prev( elif ( forward_batch.forward_mode.is_extend() and is_in_breakable_cuda_graph() ): - hash_ids = bcg_deepseek_v4_engram_hash_ids( - self.engram_hasher, input_ids - ) + hash_ids = self._eager_engram_hash(input_ids, forward_batch) else: hash_ids = self.engram_hasher(input_ids, forward_batch) tail = None @@ -4475,11 +4440,7 @@ def _forward_layers_hc_pre_from_prev( if tail is not None and i < self.late_layer_start: aux = tail.rows(aux) dspark_aux_hidden_states.append(aux.mean(dim=1)) - ctx = ( - nullcontext() - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else get_global_expert_distribution_recorder().with_current_layer(i) - ) + ctx = get_global_expert_distribution_recorder().with_current_layer(i) next_norm = None next_input = [] # The next layer can consume a collapsed input only if no Engram @@ -4774,11 +4735,7 @@ def forward( for i in range(self.start_layer, self.end_layer): layer = self.layers[i] last_layer = layer - ctx = ( - nullcontext() - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else get_global_expert_distribution_recorder().with_current_layer(i) - ) + ctx = get_global_expert_distribution_recorder().with_current_layer(i) with ctx: hidden_states, prev_residual, prev_post, prev_comb = layer( positions=positions, @@ -4836,6 +4793,13 @@ def forward( return hidden_states, pre_hc_head + @eager_on_graph + def _eager_engram_hash( + self, input_ids: torch.Tensor, forward_batch: ForwardBatch + ) -> torch.Tensor: + # The hasher reads per-request rows, so it cannot run inside the CUDA graph. + return self.engram_hasher(input_ids, forward_batch) + class DeepseekV4ForCausalLM(nn.Module): supports_cuda_vmm_feature_transport = True diff --git a/python/sglang/srt/models/glm5_next.py b/python/sglang/srt/models/glm5_next.py index 5fc8c9bc5b49..7a27ac51e14b 100644 --- a/python/sglang/srt/models/glm5_next.py +++ b/python/sglang/srt/models/glm5_next.py @@ -70,11 +70,6 @@ general_mm_embed_routine, ) from sglang.srt.managers.schedule_batch import MultimodalDataItem, MultimodalInputs -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ( ForwardBatch, PPProxyTensors, @@ -1149,11 +1144,7 @@ def forward( topk_indices = None for i in range(normal_start_layer, normal_end_layer): # NOTE: torch dynamo does not support graph break in context manager - ctx = ( - nullcontext() - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else get_global_expert_distribution_recorder().with_current_layer(i) - ) + ctx = get_global_expert_distribution_recorder().with_current_layer(i) with ctx: def capture_output(aux_hidden_state, *, owned=False): diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py index 5218e5d76dae..08199750dc3f 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -64,10 +64,6 @@ VocabParallelEmbedding, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.model_loader.weight_utils import ( RUNAI_STREAMER_TENSOR_ATTR, default_weight_loader, @@ -92,7 +88,6 @@ is_npu, make_layers, ) -from sglang.srt.utils.custom_op import register_custom_op _is_cpu = is_cpu() _is_npu = is_npu() @@ -329,12 +324,9 @@ def forward_normal( else: router_input = hidden_states - if is_in_tc_piecewise_cuda_graph(): - final_hidden_states = moe_impl(self.layer_id, hidden_states) - else: - router_logits, _ = self.router(router_input) - topk_output = self.topk(router_input, router_logits) - final_hidden_states = self.experts(hidden_states, topk_output) + router_logits, _ = self.router(router_input) + topk_output = self.topk(router_input, router_logits) + final_hidden_states = self.experts(hidden_states, topk_output) final_hidden_states = reduce_moe_output(final_hidden_states) @@ -353,16 +345,6 @@ def forward_normal( return ans -@register_custom_op(out_shape="hidden_states") -def moe_impl(layer_id: int, hidden_states: torch.Tensor) -> torch.Tensor: - forward_context = get_tc_piecewise_forward_context() - moe_fusion = forward_context.moe_fusions[layer_id] - router_logits, _ = moe_fusion.router(hidden_states) - topk_output = moe_fusion.topk(hidden_states, router_logits) - final_hidden_states = moe_fusion.experts(hidden_states, topk_output) - return final_hidden_states - - class GptOssAttention(nn.Module): def __init__( self, diff --git a/python/sglang/srt/models/inkling.py b/python/sglang/srt/models/inkling.py index e9eeb9e616e4..a6f74d2800a5 100644 --- a/python/sglang/srt/models/inkling.py +++ b/python/sglang/srt/models/inkling.py @@ -4,6 +4,7 @@ import logging import re from array import array +from contextlib import contextmanager from typing import Iterable, Optional, Set, Tuple import torch @@ -20,7 +21,6 @@ from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.moe import get_moe_runner_backend from sglang.srt.layers.quantization.base_config import QuantizationConfig -from sglang.srt.layers.radix_attention import force_eager_attention from sglang.srt.layers.utils import get_layer_id from sglang.srt.layers.vocab_parallel_embedding import ( ParallelLMHead, @@ -40,9 +40,6 @@ eager_on_graph, is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, -) from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.inkling_common.attn import ( InklingAttention, @@ -151,6 +148,16 @@ def _is_unsupported_mm_weight_name(name: str) -> bool: ) +@contextmanager +def force_eager_attention(layer): + previous = layer.use_prefill_attention_wrapper + layer.use_prefill_attention_wrapper = False + try: + yield + finally: + layer.use_prefill_attention_wrapper = previous + + class InklingDecoderLayer(nn.Module): def __init__( self, @@ -264,8 +271,6 @@ def __init__( # ONE eager break; only mlp_norm + MoE stay captured. Outside a capture # these wrappers just run inline. `_breakable_mlp_sconv` runs the final # layer's deferred mlp_sconv after the layer loop. - self._breakable_attn_group = eager_on_graph(True)(self._attn_group_impl) - self._breakable_mlp_sconv = eager_on_graph(True)(self._mlp_sconv_impl) def _attn_block( self, @@ -361,11 +366,13 @@ def _attn_block( # then restore the full buffer on forward_batch (shared across layers/replays). orig_out_cache_loc = forward_batch.out_cache_loc forward_batch.out_cache_loc = orig_out_cache_loc[: hs.shape[0]] - with force_eager_attention(): - hs = self.attn( - hs, positions, forward_batch, log_scaling_tau=log_scaling_tau - ) - forward_batch.out_cache_loc = orig_out_cache_loc + try: + with force_eager_attention(self.attn.attn): + hs = self.attn( + hs, positions, forward_batch, log_scaling_tau=log_scaling_tau + ) + finally: + forward_batch.out_cache_loc = orig_out_cache_loc else: hs = self.attn( hs, @@ -387,7 +394,8 @@ def _attn_block( hs = all_gather_hidden(hs, self.attn_tp_group) return hs, res - def _attn_group_impl( + @eager_on_graph + def _eager_attn_group( self, hidden_states: torch.Tensor, residual: Optional[torch.Tensor], @@ -396,12 +404,12 @@ def _attn_group_impl( residual_out: torch.Tensor, prev_mlp_sconv: Optional[ShortConvolution], log_scaling_tau: Optional[torch.Tensor], + forward_batch: ForwardBatch, ) -> None: """Eager break: run `_attn_block` on the REAL (non-padded) tokens with the LIVE forward_batch and write the result into the padded output buffers. Mutates attn_out / residual_out and returns None (the eager_on_graph copy-back is per-tensor, not per-tuple, so outputs must be pre-allocated buffers).""" - forward_batch = get_tc_piecewise_forward_context().forward_batch n = forward_batch.global_num_token_non_padded_cpu # log_scaling_tau is per-token, so narrow it to match the real tokens too. hs, res = self._attn_block( @@ -417,15 +425,16 @@ def _attn_group_impl( if attn_out.shape[0] != n: torch._foreach_zero_((attn_out[n:], residual_out[n:])) - def _mlp_sconv_impl( + @eager_on_graph + def _eager_mlp_sconv( self, hidden_states: torch.Tensor, positions: torch.Tensor, out: torch.Tensor, + forward_batch: ForwardBatch, ) -> None: """Eager break for the final layer's deferred mlp_sconv: run on the real tokens with the live forward_batch, write the padded output buffer.""" - forward_batch = get_tc_piecewise_forward_context().forward_batch n = forward_batch.global_num_token_non_padded_cpu y = self.mlp_sconv(hidden_states[:n], positions[:n], forward_batch) if self.scattered_sconv: @@ -463,30 +472,24 @@ def forward( if forward_batch.forward_mode.is_idle(): return hidden_states, residual - # The eager group reads the LIVE forward_batch from the tc_piecewise context - # (the only hook evaluated at BCG replay time — the break's captured args are - # frozen at capture-bucket shapes). Only the prefill BCG runner installs that - # context; the decode breakable backend sets is_in_breakable_cuda_graph() but - # NOT the context, so gate on both and otherwise fall through to the inline - # path below (which uses the passed forward_batch — correct for decode). + # Prefill replay rebinds the group's forward_batch to the live prepared batch. if ( is_in_breakable_cuda_graph() - and get_tc_piecewise_forward_context() is not None + and forward_batch.forward_mode.is_extend_without_speculative() ): # BCG prefill path: the AR fusion is decode-only, so partials never # reach (or leave) this branch. assert not prev_mlp_partial and not fuse_ar_sconv and not fuse_attn_ar # BCG: {prev mlp_sconv, attn_norm, attn, attn_sconv} run eagerly (one - # break under capture); mlp_norm + MoE stay captured. (The live - # forward_batch inside the break is read from the shared tc_piecewise - # context, which the prefill BCG runner populates at capture and replay.) + # break under capture); mlp_norm + MoE stay captured. The decorator + # supplies the live prepared forward_batch at replay. # Under scattered sconv the group's INPUT can be the previous layer's # [T, H/P] MoE shard while its OUTPUT is post-all-gather [T, H], so # size the output buffers explicitly (residual is always [T, H]). out_shape = (hidden_states.shape[0], self.attn_norm.weight.shape[0]) attn_out = hidden_states.new_empty(out_shape) residual_out = hidden_states.new_empty(out_shape) - self._breakable_attn_group( + self._eager_attn_group( hidden_states, residual, positions, @@ -494,6 +497,7 @@ def forward( residual_out, prev_mlp_sconv, log_scaling_tau, + forward_batch=forward_batch, ) hidden_states, residual = self.mlp_norm(attn_out, residual_out) del attn_out @@ -909,13 +913,11 @@ def forward( if self._dflash_layers_to_capture else hidden_states ) - # Same gate as the per-layer group: the eager break needs the tc_piecewise - # context (installed only by the prefill BCG runner) to read the live - # forward_batch at replay; else run inline with the passed forward_batch. + # Match the per-layer prefill eager region. scattered = self.layers[-1].scattered_sconv if ( is_in_breakable_cuda_graph() - and get_tc_piecewise_forward_context() is not None + and forward_batch.forward_mode.is_extend_without_speculative() ): # Under scattered sconv the input is the last MoE's [T, H/P] # shard; the break's output buffer is post-all-gather [T, H]. @@ -925,8 +927,11 @@ def forward( else hidden_states.shape ) mlp_sconv_out = hidden_states.new_empty(out_shape) - self.layers[-1]._breakable_mlp_sconv( - hidden_states, positions, mlp_sconv_out + self.layers[-1]._eager_mlp_sconv( + hidden_states, + positions, + mlp_sconv_out, + forward_batch=forward_batch, ) hidden_states = mlp_sconv_out else: diff --git a/python/sglang/srt/models/inkling_common/kernels/comm.py b/python/sglang/srt/models/inkling_common/kernels/comm.py index fa0fc14c61d0..57768d94f53a 100644 --- a/python/sglang/srt/models/inkling_common/kernels/comm.py +++ b/python/sglang/srt/models/inkling_common/kernels/comm.py @@ -7,6 +7,9 @@ import torch from sglang.srt.environ import envs +from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( + is_in_breakable_cuda_graph, +) from sglang.srt.runtime_context import get_exec from sglang.srt.utils import is_cuda @@ -710,10 +713,9 @@ def scattered_ar_sconv_fusable( if not (fm.is_extend() or fm.is_decode()): return False # Prefill scope: the BCG runner's eager-break sites are not wired (its - # baked flags would disagree with the break bodies) and tc_piecewise's - # FX pieces can't carry the cross-layer producer/consumer contract, so - # both fall back to the unfused chain. The FULL prefill CUDA-graph - # backend (context.full_graph -- the whole model captured uniformly) IS + # baked flags would disagree with the break bodies), so it falls back + # to the unfused chain. The FULL prefill CUDA-graph + # backend (the whole model captured uniformly) IS # supported: the kernel is capture-safe (barrier epochs advance across # replays; validated capture+replay) and all its metadata (qsl/si/ # cache_mask/safe_idx/track rows) is recomputed in-graph from the @@ -721,12 +723,8 @@ def scattered_ar_sconv_fusable( # rows only write pad rows of the OUT region (the eager tail slices # [:raw]), and sentinel request slots have qlen == 0 so the in-kernel # cache update/track skip them. - from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, - ) - tc_ctx = get_tc_piecewise_forward_context() - if tc_ctx is not None and not tc_ctx.full_graph: + if is_in_breakable_cuda_graph() and fm.is_extend_without_speculative(): return False comm = group.torch_symm_mem_comm if ( @@ -1050,15 +1048,11 @@ def fullwidth_ar_sconv_fusable( return False if num_tokens < _INKLING_AR_FW_MIN_TOKENS: return False - # Same prefill-runner scope as the scattered gate: BCG / tc_piecewise + # Same prefill-runner scope as the scattered gate: BCG # pieces can't carry the cross-layer producer contract; the FULL prefill # CUDA-graph backend is supported (capture-safe kernel, in-graph metadata). - from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, - ) - tc_ctx = get_tc_piecewise_forward_context() - if tc_ctx is not None and not tc_ctx.full_graph: + if is_in_breakable_cuda_graph(): return False comm = group.torch_symm_mem_comm if ( diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py index b04ac8b40cac..a8f075005b82 100644 --- a/python/sglang/srt/models/kimi_k3.py +++ b/python/sglang/srt/models/kimi_k3.py @@ -1033,17 +1033,12 @@ def _can_overlap_shared_experts_npu(self, hidden_states: torch.Tensor) -> bool: from sglang.srt.batch_overlap.two_batch_overlap import ( MaybeTboDeepEPDispatcher, ) - from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, - ) # The hooks must surround the complete dispatch, including its receive # wait. Fused EP bypasses these hooks. An eager/piecewise graph break # must not split the side-stream event record from its wait. - return ( - isinstance(self.experts.dispatcher, MaybeTboDeepEPDispatcher) - and not is_in_breakable_cuda_graph() - and not is_in_tc_piecewise_cuda_graph() + return (isinstance(self.experts.dispatcher, MaybeTboDeepEPDispatcher)) and ( + not is_in_breakable_cuda_graph() ) def _forward_unfused( diff --git a/python/sglang/srt/models/minimax_m2.py b/python/sglang/srt/models/minimax_m2.py index 0beecf8eef2a..0958bd70af82 100644 --- a/python/sglang/srt/models/minimax_m2.py +++ b/python/sglang/srt/models/minimax_m2.py @@ -16,7 +16,6 @@ """Inference-only MiniMax M2 model compatible with HuggingFace weights.""" import logging -from contextlib import nullcontext from functools import lru_cache from typing import Any, Dict, Iterable, Optional, Set, Tuple, Union @@ -66,11 +65,6 @@ ParallelLMHead, VocabParallelEmbedding, ) -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import ( default_weight_loader, @@ -646,12 +640,8 @@ def op_select_experts(self, state): hidden_states = state.hidden_states_mlp_input if router_logits is not None: - ctx = ( - nullcontext() - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else get_global_expert_distribution_recorder().with_current_layer( - self.layer_id - ) + ctx = get_global_expert_distribution_recorder().with_current_layer( + self.layer_id ) with ctx: state.topk_weights_local, state.topk_idx_local, _ = self.topk( @@ -684,12 +674,8 @@ def op_dispatch_a(self, state): def op_dispatch_b(self, state): """Dispatch B operation for TBO - complete async dispatch""" if self.ep_size > 1: - ctx = ( - nullcontext() - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else get_global_expert_distribution_recorder().with_current_layer( - self.layer_id - ) + ctx = get_global_expert_distribution_recorder().with_current_layer( + self.layer_id ) with ctx: state.dispatch_output = self.experts.deepep_dispatcher.dispatch_b( @@ -1158,11 +1144,7 @@ def forward( ) else: for i in range(self.start_layer, self.end_layer): - ctx = ( - nullcontext() - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else get_global_expert_distribution_recorder().with_current_layer(i) - ) + ctx = get_global_expert_distribution_recorder().with_current_layer(i) with ctx: layer = self.layers[i] hidden_states = layer( diff --git a/python/sglang/srt/models/minimax_m3.py b/python/sglang/srt/models/minimax_m3.py index f08b2c1a8d7e..3f47a3a1b73a 100644 --- a/python/sglang/srt/models/minimax_m3.py +++ b/python/sglang/srt/models/minimax_m3.py @@ -16,7 +16,6 @@ """Inference-only MiniMax M3 model compatible with HuggingFace weights.""" import logging -from contextlib import nullcontext from typing import Iterable, List, Optional, Set, Tuple, Union import torch @@ -73,11 +72,6 @@ ParallelLMHead, VocabParallelEmbedding, ) -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_context import ( get_forward_context, @@ -1499,11 +1493,7 @@ def forward( else: for i in range(self.start_layer, self.end_layer): # NOTE: torch dynamo does not support graph break in context manager - ctx = ( - nullcontext() - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else get_global_expert_distribution_recorder().with_current_layer(i) - ) + ctx = get_global_expert_distribution_recorder().with_current_layer(i) with ctx: layer = self.layers[i] hidden_states = layer( diff --git a/python/sglang/srt/models/nemotron_h.py b/python/sglang/srt/models/nemotron_h.py index 4a0cc531c65c..d2dc65b6a86c 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -22,7 +22,6 @@ import torch from torch import nn -from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.configs import NemotronHConfig from sglang.srt.configs.nemotron_h import ATTENTION, MAMBA, MLP, MOE from sglang.srt.layers.activation import ReLU2 @@ -67,10 +66,6 @@ eager_on_graph, is_in_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, - is_in_tc_piecewise_cuda_graph, -) from sglang.srt.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name, @@ -86,7 +81,6 @@ is_cuda, make_layers, ) -from sglang.srt.utils.custom_op import register_custom_op from sglang.utils import logger _is_cuda = is_cuda() @@ -529,18 +523,40 @@ def _forward_mixer( ) -> torch.Tensor: if is_in_breakable_cuda_graph(): output = torch.empty_like(hidden_states) - breakable_nemotron_mamba2_with_output( - hidden_states, output, self.layer_id, skip_reduce - ) - return output - if is_in_tc_piecewise_cuda_graph(): - output = torch.empty_like(hidden_states) - nemotron_mamba2_with_output( - hidden_states, output, self.layer_id, skip_reduce - ) + self._eager_mamba(hidden_states, output, forward_batch, skip_reduce) return output return self._forward_mamba(hidden_states, forward_batch) + @eager_on_graph + def _eager_mamba( + self, + hidden_states: torch.Tensor, + output: torch.Tensor, + forward_batch: ForwardBatch, + fuse_mlp_allreduce: bool = False, + ) -> None: + """Run Mamba2 on live tokens between graph segments.""" + # In graph mode, hidden_states may be padded to the + # captured graph size. Slice to actual token count for Mamba forward. + attn_backend = get_attn_backend() + metadata = attn_backend.linear_attn_backend.forward_metadata + num_actual_tokens = metadata.num_prefill_tokens + ( + metadata.num_decodes * metadata.draft_token_num + if metadata.is_target_verify + else metadata.num_decodes + ) + if hidden_states.shape[0] != num_actual_tokens: + hidden_states = hidden_states[:num_actual_tokens] + + # Replay runs outside the caller's scope; restore its collective policy. + with get_forward().scoped(fuse_mlp_allreduce=fuse_mlp_allreduce): + ret = self._forward_mamba(hidden_states, forward_batch) + + # Copy result back; output may be larger (padded) so only fill actual tokens + output[:num_actual_tokens].view(ret.shape).copy_(ret) + if output.shape[0] != num_actual_tokens: + output[num_actual_tokens:].zero_() + class NemotronHAttention(nn.Module): def __init__( @@ -1144,47 +1160,3 @@ class NemotronHPuzzleForCausalLM(NemotronHForCausalLM): EntryClass = [NemotronHForCausalLM, NemotronHPuzzleForCausalLM] - - -@register_custom_op(mutates_args=["output"]) -@register_split_op() -def nemotron_mamba2_with_output( - hidden_states: torch.Tensor, - output: torch.Tensor, - layer_id: int, - fuse_mlp_allreduce: bool = False, -) -> None: - """Split op for Mamba2 forward in piecewise CUDA graph mode.""" - context = get_tc_piecewise_forward_context() - forward_batch = context.forward_batch - attention_layers = context.attention_layers - mamba_layer = attention_layers[layer_id] - - # In piecewise CUDA graph mode, hidden_states may be padded to the - # captured graph size. Slice to actual token count for Mamba forward. - attn_backend = get_attn_backend() - metadata = attn_backend.linear_attn_backend.forward_metadata - num_actual_tokens = metadata.num_prefill_tokens + ( - metadata.num_decodes * metadata.draft_token_num - if metadata.is_target_verify - else metadata.num_decodes - ) - if hidden_states.shape[0] != num_actual_tokens: - hidden_states = hidden_states[:num_actual_tokens] - - # This function is an opaque custom op under torch.compile. The caller's - # ForwardFlags scope is Python control-plane state and is no longer active - # when the compiled graph invokes this implementation. Carry the scalar - # across the graph boundary and republish it for RowParallelLinear. - with get_forward().scoped(fuse_mlp_allreduce=fuse_mlp_allreduce): - ret = mamba_layer._forward_mamba(hidden_states, forward_batch) - - # Copy result back; output may be larger (padded) so only fill actual tokens - output[:num_actual_tokens].view(ret.shape).copy_(ret) - if output.shape[0] != num_actual_tokens: - output[num_actual_tokens:].zero_() - - -breakable_nemotron_mamba2_with_output = eager_on_graph(True)( - nemotron_mamba2_with_output -) diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index 8755df797ef8..f7fe9691b7d8 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -19,7 +19,6 @@ """Inference-only Qwen2MoE model compatible with HuggingFace weights.""" import logging -from contextlib import nullcontext from typing import Any, Dict, Iterable, List, Optional, Tuple, Union import torch @@ -83,11 +82,6 @@ ParallelLMHead, VocabParallelEmbedding, ) -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( @@ -1195,11 +1189,7 @@ def forward( ) else: for i in range(self.start_layer, self.end_layer): - ctx = ( - nullcontext() - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else get_global_expert_distribution_recorder().with_current_layer(i) - ) + ctx = get_global_expert_distribution_recorder().with_current_layer(i) with ctx: layer = self.layers[i] hidden_states = layer( diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index 2939fc608d82..4ec54abe3bc1 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -22,11 +22,6 @@ from sglang.srt.layers.rotary_embedding.mrope import MRotaryEmbedding from sglang.srt.layers.utils import PPMissingLayer, get_layer_id from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_context import get_token_to_kv_pool from sglang.srt.model_loader.weight_utils import ( @@ -413,16 +408,7 @@ def forward( hidden_states = self.ffn_boundary.prepare( hidden_states, forward_batch, - cache=( - [self.mlp.gate_up_proj.weight, self.mlp.down_proj.weight] - if _is_npu - and check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - and ( - hasattr(self.mlp.gate_up_proj, "weight") - and hasattr(self.mlp.down_proj, "weight") - ) - else None - ), + cache=(None), ) with self.ffn_boundary.exit(forward_batch) as ffn_exit: hidden_states = self.mlp(hidden_states, forward_batch=forward_batch) diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py index fac16f30a38a..6ba92aba74ae 100644 --- a/python/sglang/srt/models/qwen3_5.py +++ b/python/sglang/srt/models/qwen3_5.py @@ -88,11 +88,6 @@ from sglang.srt.layers.rotary_embedding import get_rope from sglang.srt.layers.utils import PPMissingLayer, get_layer_id from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import ( @@ -808,11 +803,7 @@ def _forward_input_proj(self, hidden_states: torch.Tensor): fused_out[:, self._fused_in_proj_qkvz_width :], ) - if ( - _is_cpu - or _is_npu - or check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - ): + if (_is_cpu) or (_is_npu): DUAL_STREAM_TOKEN_THRESHOLD = 0 else: DUAL_STREAM_TOKEN_THRESHOLD = 1024 @@ -855,10 +846,7 @@ def _forward_input_proj_fused_quant_amd(self, hidden_states): hs_qkvz = _select_fused_ar_input_for_linear(hidden_states, self.in_proj_qkvz) seq_len = hs_bf16.shape[0] - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE): - DUAL_STREAM_TOKEN_THRESHOLD = 0 - else: - DUAL_STREAM_TOKEN_THRESHOLD = 1024 + DUAL_STREAM_TOKEN_THRESHOLD = 1024 if ( self.alt_stream is not None diff --git a/python/sglang/srt/models/qwen3_next.py b/python/sglang/srt/models/qwen3_next.py index ccfd7c5656f4..2c8507909bd0 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -7,7 +7,6 @@ from torch import nn from sglang.kernels.ops.attention.fla.fused_norm_gate import FusedRMSNormGated -from sglang.kernels.ops.attention.fla.layernorm_gated import RMSNorm as RMSNormGated from sglang.kernels.ops.attention.triton_gdn_fused_proj import ( fused_qkvzba_split_reshape_cat, ) @@ -45,11 +44,6 @@ ParallelLMHead, VocabParallelEmbedding, ) -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import ( @@ -196,32 +190,16 @@ def __init__( set_weight_attrs(self.A_log, {"weight_loader": sharded_weight_loader(0)}) set_weight_attrs(self.dt_bias, {"weight_loader": sharded_weight_loader(0)}) - self.norm = ( - RMSNormGated( - self.head_v_dim, - eps=self.layer_norm_epsilon, - group_size=None, - norm_before_gate=True, - device=torch.get_device_module().current_device(), - dtype=config.torch_dtype, - **( - {"activation": self.output_gate_type} - if self.output_gate_type is not None - else {} - ), - ) - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else FusedRMSNormGated( - self.head_v_dim, - eps=self.layer_norm_epsilon, - activation=( - self.output_gate_type - if self.output_gate_type is not None - else self.activation - ), - device=torch.get_device_module().current_device(), - dtype=config.torch_dtype, - ) + self.norm = FusedRMSNormGated( + self.head_v_dim, + eps=self.layer_norm_epsilon, + activation=( + self.output_gate_type + if self.output_gate_type is not None + else self.activation + ), + device=torch.get_device_module().current_device(), + dtype=config.torch_dtype, ) self.out_proj = RowParallelLinear( @@ -383,11 +361,7 @@ def fix_query_key_value_ordering( return query, key, value, z, b, a def _forward_input_proj(self, hidden_states: torch.Tensor): - if ( - _is_cpu - or _is_npu - or check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - ): + if (_is_cpu) or (_is_npu): DUAL_STREAM_TOKEN_THRESHOLD = 0 else: DUAL_STREAM_TOKEN_THRESHOLD = 1024 diff --git a/python/sglang/srt/models/transformers.py b/python/sglang/srt/models/transformers.py index 46b8bd32f238..8dbe4a42e563 100644 --- a/python/sglang/srt/models/transformers.py +++ b/python/sglang/srt/models/transformers.py @@ -509,16 +509,6 @@ def _transformers_moe_forward_fake( fake_impl=_transformers_moe_forward_fake, ) -try: - from sglang.srt.compilation.compilation_config import SPLIT_OPS - - _MOE_SPLIT_OP = "sglang.transformers_moe_forward" - if _MOE_SPLIT_OP not in SPLIT_OPS: - SPLIT_OPS.append(_MOE_SPLIT_OP) -except ImportError: - pass - - _BASE_DYNAMIC_ARG_DIMS: dict[str, int] = { "input_ids": 0, "positions": 0, diff --git a/python/sglang/srt/models/utils.py b/python/sglang/srt/models/utils.py index e8a03a2f5a19..605d9e26dfe5 100644 --- a/python/sglang/srt/models/utils.py +++ b/python/sglang/srt/models/utils.py @@ -37,7 +37,6 @@ from sglang.srt.model_executor.forward_context import get_token_to_kv_pool from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import get_exec from sglang.srt.utils import get_current_device_stream_fast, is_cpu, is_cuda, is_hip from sglang.srt.utils.custom_op import register_custom_op @@ -428,25 +427,7 @@ def rot_pos_ids(h: int, w: int, spatial_merge_size: int) -> torch.Tensor: def _reshape_for_qk_norm(x: torch.Tensor, head_dim: int) -> torch.Tensor: - """Reshape a (..., H*D) tensor into (..., H, D) ahead of QK RMSNorm. - - On CUDA with the inductor piecewise-cuda-graph compiler, return a - stride-preserving view so inductor can fuse this reshape with the - subsequent RMSNorm (and any upstream/downstream FP8 quant) into a - single triton kernel -- the original motivation of #21734. - - Everywhere else (ROCm, or CUDA with the eager PCG fallback), use the - flat 2D reshape that forces a copy when the input is a non-contiguous - QKV-split stride-trick view. ROCm's RMSNorm kernels assume contiguous - inputs and fault on strided tensors (root cause of the #21734 revert - in #23159). - """ - - if ( - _is_cuda - and get_exec().graph.cuda_graph_config.prefill.tc_compiler == "inductor" - ): - return x.view(*x.shape[:-1], -1, head_dim) + """Flatten head rows, making non-contiguous QKV views safe for RMSNorm.""" return x.reshape(-1, head_dim) @@ -515,8 +496,6 @@ def apply_qk_norm( and allow_inplace # TODO(dark): this can be relaxed if needed and (q_eps == k_eps) # TODO(dark): this can also be relaxed and not envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.get() - and get_exec().graph.cuda_graph_config.prefill.tc_compiler - != "inductor" # let inductor fuse QK norm and can_use_fused_inplace_qknorm(head_dim, q.dtype) ): fused_inplace_qknorm( diff --git a/python/sglang/srt/models/xllm.py b/python/sglang/srt/models/xllm.py index 59a5e77a09b0..686c9168ec31 100644 --- a/python/sglang/srt/models/xllm.py +++ b/python/sglang/srt/models/xllm.py @@ -27,7 +27,6 @@ """Inference-only xLLM K2MoE and MoVA models compatible with HF weights.""" import math -from contextlib import nullcontext from typing import Any, Dict, Iterable, Optional, Tuple, Union import torch @@ -76,11 +75,6 @@ ParallelLMHead, VocabParallelEmbedding, ) -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.runtime_context import get_exec, get_parallel @@ -1634,11 +1628,7 @@ def forward( ) for i in range(self.start_layer, self.end_layer): - ctx = ( - nullcontext() - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else get_global_expert_distribution_recorder().with_current_layer(i) - ) + ctx = get_global_expert_distribution_recorder().with_current_layer(i) with ctx: layer = self.layers[i] hidden_states = layer(positions, hidden_states, forward_batch) diff --git a/python/sglang/srt/platforms/cpu.py b/python/sglang/srt/platforms/cpu.py index 9c28ae676e8b..e5546fb60655 100644 --- a/python/sglang/srt/platforms/cpu.py +++ b/python/sglang/srt/platforms/cpu.py @@ -123,7 +123,7 @@ def get_torch_distributed_backend_str(self) -> str: class CpuSRTPlatform(CpuDeviceMixin, SRTPlatform): """Default in-tree CPU SRT platform. - supports_fp8 / support_cuda_graph / support_piecewise_cuda_graph keep the + supports_fp8 / support_cuda_graph keep the conservative SRTPlatform defaults (all False), so they are not repeated here. is_pin_memory_available is repeated for explicitness: CPU has no GPU to pin host memory to. diff --git a/python/sglang/srt/platforms/cuda.py b/python/sglang/srt/platforms/cuda.py index 3ae863561b6c..b9e27ef72eb7 100644 --- a/python/sglang/srt/platforms/cuda.py +++ b/python/sglang/srt/platforms/cuda.py @@ -108,6 +108,3 @@ def supports_fp8(self) -> bool: def support_cuda_graph(self) -> bool: return True - - def support_piecewise_cuda_graph(self) -> bool: - return True diff --git a/python/sglang/srt/platforms/interface.py b/python/sglang/srt/platforms/interface.py index 7ec27272340a..95060ddd4229 100644 --- a/python/sglang/srt/platforms/interface.py +++ b/python/sglang/srt/platforms/interface.py @@ -93,10 +93,6 @@ def get_compile_backend(self, mode: str | None = None) -> str: """ return "inductor" - def get_piecewise_backend_cls(self) -> type: - """Return the piecewise compilation backend class for this platform.""" - raise NotImplementedError - def get_speculative_cache_locs_fn( self, ) -> Optional[Callable[..., torch.Tensor]]: @@ -136,14 +132,6 @@ def support_cuda_graph(self) -> bool: """ return False - def support_piecewise_cuda_graph(self) -> bool: - """Whether this platform supports piecewise CUDA graph. - - Controls PiecewiseCudaGraphRunner for the prefill/extend path - (torch.compile backend). - """ - return False - # ------------------------------------------------------------------ # Initialization # ------------------------------------------------------------------ diff --git a/python/sglang/srt/platforms/npu.py b/python/sglang/srt/platforms/npu.py index c1382ff768b2..ea1660b0a782 100644 --- a/python/sglang/srt/platforms/npu.py +++ b/python/sglang/srt/platforms/npu.py @@ -84,6 +84,3 @@ def supports_fp8(self) -> bool: def support_cuda_graph(self) -> bool: # NPUGraphRunner in hardware_backend/npu/graph_runner return True - - def support_piecewise_cuda_graph(self) -> bool: - return False diff --git a/python/sglang/srt/platforms/rocm.py b/python/sglang/srt/platforms/rocm.py index c250c482ec53..b4828ccfa381 100644 --- a/python/sglang/srt/platforms/rocm.py +++ b/python/sglang/srt/platforms/rocm.py @@ -23,7 +23,7 @@ class RocmDeviceMixin(CudaDeviceMixin): class RocmSRTPlatform(RocmDeviceMixin, SRTPlatform): """Default in-tree ROCm SRT platform. - Capability flags (supports_fp8, support_cuda_graph, support_piecewise_cuda_graph) + Capability flags (supports_fp8, support_cuda_graph) keep the conservative SRTPlatform defaults rather than mirroring CudaSRTPlatform. They are currently only consulted in OOT branches gated on is_out_of_tree(), so the defaults are behaviorally inert for the in-tree ROCm path. A follow-up diff --git a/python/sglang/srt/platforms/xpu.py b/python/sglang/srt/platforms/xpu.py index 2995d7250133..ace6cea77c1f 100644 --- a/python/sglang/srt/platforms/xpu.py +++ b/python/sglang/srt/platforms/xpu.py @@ -99,6 +99,3 @@ def supports_fp8(self) -> bool: def support_cuda_graph(self) -> bool: return True - - def support_piecewise_cuda_graph(self) -> bool: - return True diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 5742a0bd7fe3..a2d9123e19d1 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -2001,7 +2001,6 @@ def cutedsl_moe_max_num_tokens() -> int: resolution pipeline uses. Max over the prefill bound, the piecewise-prefill capture, and the decode/verify bound. """ - from sglang.srt.model_executor.cuda_graph_config import Backend spec = get_spec() num_tokens_per_req = ( @@ -2009,8 +2008,6 @@ def cutedsl_moe_max_num_tokens() -> int: ) prefill_tokens = get_schedule().max_prefill_tokens cg_config = get_exec().graph.cuda_graph_config - if cg_config is not None and cg_config.prefill.backend == Backend.TC_PIECEWISE: - prefill_tokens = max(prefill_tokens, cg_config.prefill.max_bs or 0) decode_max_bs = (cg_config.decode.max_bs if cg_config is not None else 0) or 0 return max(prefill_tokens, decode_max_bs * num_tokens_per_req) diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index 59572210b982..76d3069b8d84 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -3923,11 +3923,8 @@ def dispose_tensor(x: torch.Tensor): from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( is_in_breakable_cuda_graph, ) - from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - is_in_tc_piecewise_cuda_graph, - ) - if is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph(): + if is_in_breakable_cuda_graph(): return if get_flags().capture.disable_dispose_tensor: diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py index 0bff8187972f..65ef0c76a5e4 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py @@ -314,7 +314,7 @@ def __init__( max_context_len: int, head_dim: int, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, ): pool_batch_size = runner_batch_size or case.batch_size @@ -361,8 +361,8 @@ def __init__( prefill=PhaseConfig( backend=( Backend.DISABLED - if (disable_cuda_graph or disable_piecewise_cuda_graph) - else Backend.TC_PIECEWISE + if (disable_cuda_graph or disable_prefill_cuda_graph) + else Backend.BREAKABLE ), ), ), @@ -971,7 +971,7 @@ def build_dense_attention_fixture( dtype: torch.dtype = DEFAULT_DTYPE, device: str = DEFAULT_DEVICE, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, loc_layout: str = "shuffled_pages", ) -> DenseAttentionFixture: @@ -995,7 +995,7 @@ def build_dense_attention_fixture( max_context_len=max_context_len, head_dim=head_dim, disable_cuda_graph=disable_cuda_graph, - disable_piecewise_cuda_graph=disable_piecewise_cuda_graph, + disable_prefill_cuda_graph=disable_prefill_cuda_graph, runner_batch_size=runner_batch_size, ) try: diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py index 440e497e87d1..ea9ec25bdcfa 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py @@ -278,7 +278,7 @@ def __init__( max_context_len: int, head_dim: int, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, dsa_prefill_backend: str = "flashmla_auto", dsa_decode_backend: str = "flashmla_kv", @@ -331,8 +331,8 @@ def __init__( prefill=PhaseConfig( backend=( Backend.DISABLED - if (disable_cuda_graph or disable_piecewise_cuda_graph) - else Backend.TC_PIECEWISE + if (disable_cuda_graph or disable_prefill_cuda_graph) + else Backend.BREAKABLE ), ), ), @@ -612,7 +612,7 @@ def build_dsa_attention_fixture( dtype: torch.dtype = torch.bfloat16, device: str = DEFAULT_DEVICE, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, dsa_prefill_backend: str = "flashmla_auto", dsa_decode_backend: str = "flashmla_kv", @@ -642,7 +642,7 @@ def build_dsa_attention_fixture( max_context_len=max_context_len, head_dim=head_dim, disable_cuda_graph=disable_cuda_graph, - disable_piecewise_cuda_graph=disable_piecewise_cuda_graph, + disable_prefill_cuda_graph=disable_prefill_cuda_graph, runner_batch_size=runner_batch_size, dsa_prefill_backend=dsa_prefill_backend, dsa_decode_backend=dsa_decode_backend, @@ -813,7 +813,7 @@ def build_dsa_sparse_attention_fixture( dtype: torch.dtype = torch.bfloat16, device: str = DEFAULT_DEVICE, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, dsa_prefill_backend: str = "flashmla_auto", dsa_decode_backend: str = "flashmla_kv", @@ -852,7 +852,7 @@ def build_dsa_sparse_attention_fixture( max_context_len=max_context_len, head_dim=head_dim, disable_cuda_graph=disable_cuda_graph, - disable_piecewise_cuda_graph=disable_piecewise_cuda_graph, + disable_prefill_cuda_graph=disable_prefill_cuda_graph, runner_batch_size=runner_batch_size, dsa_prefill_backend=dsa_prefill_backend, dsa_decode_backend=dsa_decode_backend, diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py index 83b4161c9348..da1cceb08446 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py @@ -312,7 +312,7 @@ def __init__( max_context_len: int, swa_size: int, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, compression_ratios: list[int] = None, ): @@ -358,8 +358,8 @@ def __init__( prefill=PhaseConfig( backend=( Backend.DISABLED - if (disable_cuda_graph or disable_piecewise_cuda_graph) - else Backend.TC_PIECEWISE + if (disable_cuda_graph or disable_prefill_cuda_graph) + else Backend.BREAKABLE ), ), ), @@ -694,7 +694,7 @@ def build_dsv4_attention_fixture( dtype: torch.dtype = torch.bfloat16, device: str = "cuda", disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, compression_ratios: list[int] = None, ) -> DSV4AttentionFixture: @@ -731,7 +731,7 @@ def build_dsv4_attention_fixture( max_context_len=max_context_len, swa_size=swa_size, disable_cuda_graph=disable_cuda_graph, - disable_piecewise_cuda_graph=disable_piecewise_cuda_graph, + disable_prefill_cuda_graph=disable_prefill_cuda_graph, runner_batch_size=runner_batch_size, compression_ratios=compression_ratios, ) diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py index e7a3d4b61c41..c4ce5e924c57 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py @@ -213,7 +213,7 @@ def __init__( head_k_dim: int, head_v_dim: int, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, ): pool_batch_size = runner_batch_size or case.batch_size @@ -247,8 +247,8 @@ def __init__( prefill=PhaseConfig( backend=( Backend.DISABLED - if (disable_cuda_graph or disable_piecewise_cuda_graph) - else Backend.TC_PIECEWISE + if (disable_cuda_graph or disable_prefill_cuda_graph) + else Backend.BREAKABLE ), ), ), @@ -590,7 +590,7 @@ def build_gdn_attention_fixture( dtype: torch.dtype = DEFAULT_DTYPE, device: str = DEFAULT_DEVICE, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, loc_layout: str = "shuffled_pages", ) -> GDNAttentionFixture: @@ -613,7 +613,7 @@ def build_gdn_attention_fixture( head_k_dim=head_k_dim, head_v_dim=head_v_dim, disable_cuda_graph=disable_cuda_graph, - disable_piecewise_cuda_graph=disable_piecewise_cuda_graph, + disable_prefill_cuda_graph=disable_prefill_cuda_graph, runner_batch_size=runner_batch_size, ) try: diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py index 4254e7cd228a..a9b04da771dc 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py @@ -216,7 +216,7 @@ def __init__( head_k_dim: int, head_v_dim: int, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, ): pool_batch_size = runner_batch_size or case.batch_size @@ -250,8 +250,8 @@ def __init__( prefill=PhaseConfig( backend=( Backend.DISABLED - if (disable_cuda_graph or disable_piecewise_cuda_graph) - else Backend.TC_PIECEWISE + if (disable_cuda_graph or disable_prefill_cuda_graph) + else Backend.BREAKABLE ), ), ), @@ -580,7 +580,7 @@ def build_kda_attention_fixture( dtype: torch.dtype = DEFAULT_DTYPE, device: str = DEFAULT_DEVICE, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, loc_layout: str = "shuffled_pages", ) -> KDAAttentionFixture: @@ -603,7 +603,7 @@ def build_kda_attention_fixture( head_k_dim=head_k_dim, head_v_dim=head_v_dim, disable_cuda_graph=disable_cuda_graph, - disable_piecewise_cuda_graph=disable_piecewise_cuda_graph, + disable_prefill_cuda_graph=disable_prefill_cuda_graph, runner_batch_size=runner_batch_size, ) try: diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py index 42c285f2d00b..8bc6fb98fc0f 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py @@ -224,7 +224,7 @@ def __init__( max_context_len: int, head_dim: int, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, ): pool_batch_size = runner_batch_size or case.batch_size @@ -258,8 +258,8 @@ def __init__( prefill=PhaseConfig( backend=( Backend.DISABLED - if (disable_cuda_graph or disable_piecewise_cuda_graph) - else Backend.TC_PIECEWISE + if (disable_cuda_graph or disable_prefill_cuda_graph) + else Backend.BREAKABLE ), ), ), @@ -551,7 +551,7 @@ def build_lightning_attention_fixture( dtype: torch.dtype = DEFAULT_DTYPE, device: str = DEFAULT_DEVICE, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, loc_layout: str = "shuffled_pages", ) -> LightningAttentionFixture: @@ -573,7 +573,7 @@ def build_lightning_attention_fixture( max_context_len=max_context_len, head_dim=head_dim, disable_cuda_graph=disable_cuda_graph, - disable_piecewise_cuda_graph=disable_piecewise_cuda_graph, + disable_prefill_cuda_graph=disable_prefill_cuda_graph, runner_batch_size=runner_batch_size, ) try: diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py index 7b8c3675f0e3..384c9599c30a 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py @@ -308,7 +308,7 @@ def __init__( device: str, max_context_len: int, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, ): pool_batch_size = runner_batch_size or case.batch_size @@ -351,8 +351,8 @@ def __init__( prefill=PhaseConfig( backend=( Backend.DISABLED - if (disable_cuda_graph or disable_piecewise_cuda_graph) - else Backend.TC_PIECEWISE + if (disable_cuda_graph or disable_prefill_cuda_graph) + else Backend.BREAKABLE ), ), ), @@ -679,7 +679,7 @@ def build_mamba2_attention_fixture( dtype: torch.dtype = DEFAULT_DTYPE, device: str = DEFAULT_DEVICE, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, loc_layout: str = "shuffled_pages", ) -> Mamba2AttentionFixture: @@ -695,7 +695,7 @@ def build_mamba2_attention_fixture( device=device, max_context_len=max_context_len, disable_cuda_graph=disable_cuda_graph, - disable_piecewise_cuda_graph=disable_piecewise_cuda_graph, + disable_prefill_cuda_graph=disable_prefill_cuda_graph, runner_batch_size=runner_batch_size, ) try: diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py index 0a381ca0ee3f..b8eaf8aeca97 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py @@ -220,7 +220,7 @@ def __init__( kv_lora_rank: int, qk_rope_head_dim: int, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, fp8_kv_cache: bool = False, ): @@ -263,8 +263,8 @@ def __init__( prefill=PhaseConfig( backend=( Backend.DISABLED - if (disable_cuda_graph or disable_piecewise_cuda_graph) - else Backend.TC_PIECEWISE + if (disable_cuda_graph or disable_prefill_cuda_graph) + else Backend.BREAKABLE ), ), ), @@ -876,7 +876,7 @@ def build_mla_attention_fixture( dtype: torch.dtype = DEFAULT_DTYPE, device: str = DEFAULT_DEVICE, disable_cuda_graph: bool = True, - disable_piecewise_cuda_graph: bool = True, + disable_prefill_cuda_graph: bool = True, runner_batch_size: int | None = None, fp8_kv_cache: bool = False, loc_layout: str = "shuffled_pages", @@ -901,7 +901,7 @@ def build_mla_attention_fixture( kv_lora_rank=kv_lora_rank, qk_rope_head_dim=qk_rope_head_dim, disable_cuda_graph=disable_cuda_graph, - disable_piecewise_cuda_graph=disable_piecewise_cuda_graph, + disable_prefill_cuda_graph=disable_prefill_cuda_graph, runner_batch_size=runner_batch_size, fp8_kv_cache=fp8_kv_cache, ) diff --git a/python/sglang/test/kits/attention_unittest/runner_modes/split_op_runner.py b/python/sglang/test/kits/attention_unittest/runner_modes/split_op_runner.py index 9542b8c96ff4..7913bb681151 100644 --- a/python/sglang/test/kits/attention_unittest/runner_modes/split_op_runner.py +++ b/python/sglang/test/kits/attention_unittest/runner_modes/split_op_runner.py @@ -1,3 +1,4 @@ +from contextlib import nullcontext from dataclasses import dataclass, replace from typing import Any, Callable @@ -7,12 +8,6 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( enable_breakable_cuda_graph, ) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph.context_manager import ( - enable_tc_piecewise_cuda_graph as enable_piecewise_cuda_graph, -) -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph.context_manager import ( - set_tc_piecewise_forward_context as piecewise_forward_context, -) from ..attention_methods.dense_attention import DEFAULT_DEVICE as DENSE_DEFAULT_DEVICE from ..attention_methods.dense_attention import DEFAULT_DTYPE as DENSE_DEFAULT_DTYPE @@ -178,13 +173,13 @@ class SplitOpAdapter: def _check_extend_split_op_case(case) -> None: if not case.forward_mode.is_extend_without_speculative(): - raise ValueError("PCG/BCG split-op coverage expects non-spec extend cases.") + raise ValueError("Full/BCG split-op coverage expects non-spec extend cases.") def _split_op_context(*, breakable: bool): if breakable: return enable_breakable_cuda_graph() - return enable_piecewise_cuda_graph() + return nullcontext() def _make_static_forward_batch(raw_batch, static_num_tokens: int, device: str): @@ -274,7 +269,7 @@ def _run_split_op_extend_case( testcase, case, **build_kwargs, - disable_piecewise_cuda_graph=False, + disable_prefill_cuda_graph=False, ) split_inputs = adapter.fixture_inputs(split_fixture) split_initial_state = adapter.clone_state(split_fixture) @@ -307,13 +302,12 @@ def _run_split_op_extend_case( with ( torch.no_grad(), _split_op_context(breakable=breakable), - forward_context(ForwardContext(attn_backend=split_fixture.backend)), - piecewise_forward_context( - static_batch, - adapter.attention_layers(split_fixture), - None, - [], - [], + forward_context( + ForwardContext( + attn_backend=split_fixture.backend, + full_graph=not breakable, + raw_num_tokens=raw_num_tokens, + ) ), ): split_fixture.backend.init_forward_metadata(raw_batch) @@ -475,7 +469,7 @@ def run_kda_split_op_extend_case( dtype: torch.dtype = KDA_DEFAULT_DTYPE, device: str = KDA_DEFAULT_DEVICE, ): - """KDA PCG/BCG split-op extend. Verifies the live-token slicing contract + """KDA Full/BCG split-op extend. Verifies the live-token slicing contract with a larger static token buffer, mirroring GDN's split_op coverage.""" adapter = SplitOpAdapter( build_fixture=build_kda_attention_fixture, @@ -521,7 +515,7 @@ def run_lightning_split_op_extend_case( dtype: torch.dtype = LIGHTNING_DEFAULT_DTYPE, device: str = LIGHTNING_DEFAULT_DEVICE, ): - """Lightning PCG/BCG split-op extend. Same pattern as KDA/GDN.""" + """Lightning Full/BCG split-op extend. Same pattern as KDA/GDN.""" adapter = SplitOpAdapter( build_fixture=build_lightning_attention_fixture, fixture_inputs=lightning_fixture_inputs, @@ -564,7 +558,7 @@ def run_mamba2_split_op_extend_case( dtype: torch.dtype = MAMBA2_DEFAULT_DTYPE, device: str = MAMBA2_DEFAULT_DEVICE, ): - """Mamba2 PCG/BCG split-op extend. Same pattern as KDA. Mamba2's + """Mamba2 Full/BCG split-op extend. Same pattern as KDA. Mamba2's forward writes through an `empty_like(hidden_states)` buffer that short-circuits the RadixAttention dispatch path, so the per-head-vs-flat shape mismatch that blocks Lightning split-op doesn't apply.""" diff --git a/python/sglang/test/kits/dsa_metadata_kit.py b/python/sglang/test/kits/dsa_metadata_kit.py index 45de95799453..cf3e19280b86 100644 --- a/python/sglang/test/kits/dsa_metadata_kit.py +++ b/python/sglang/test/kits/dsa_metadata_kit.py @@ -43,7 +43,6 @@ def make_backend(mode, seq, req, *, fusion=True): backend.token_to_kv_pool = SimpleNamespace(slots_per_page=64) # Only attention-dispatch state is synthetic; every metadata kernel is real. backend._is_in_breakable_cuda_graph = lambda: False - backend._is_in_tc_piecewise_cuda_graph = lambda: False backend._get_device_sm = lambda: backend.device_sm_major * 10 backend._is_blackwell = lambda: backend.device_sm_major == 10 backend.dsa_topk_backend = DSATopKBackend.SGL_KERNEL diff --git a/reference/tc_piecewise/README.md b/reference/tc_piecewise/README.md new file mode 100644 index 000000000000..03e13cc8a12a --- /dev/null +++ b/reference/tc_piecewise/README.md @@ -0,0 +1,12 @@ +# TCPCG reference only + +Snapshot from Oasis-Git/sglang main at c421d16563, before TCPCG removal. + +These files preserve the old compiler adapters, FX partitioning, and per-piece +capture/replay implementation for future design reference. They are outside the installed SGLang +package, are not imported by the serving runtime, and are not a supported or +runnable backend. Imports intentionally retain their historical names; the +runtime interfaces and model split ops they referenced have been removed. + +Active torch.compile support remains under python/sglang/srt/compilation. +Shared graph tensor-lifetime helpers live under runner_backend_utils. diff --git a/python/sglang/srt/compilation/backend.py b/reference/tc_piecewise/compilation/backend.py similarity index 100% rename from python/sglang/srt/compilation/backend.py rename to reference/tc_piecewise/compilation/backend.py diff --git a/python/sglang/srt/compilation/compilation_config.py b/reference/tc_piecewise/compilation/compilation_config.py similarity index 100% rename from python/sglang/srt/compilation/compilation_config.py rename to reference/tc_piecewise/compilation/compilation_config.py diff --git a/python/sglang/srt/compilation/compilation_counter.py b/reference/tc_piecewise/compilation/compilation_counter.py similarity index 100% rename from python/sglang/srt/compilation/compilation_counter.py rename to reference/tc_piecewise/compilation/compilation_counter.py diff --git a/python/sglang/srt/compilation/compile.py b/reference/tc_piecewise/compilation/compile.py similarity index 100% rename from python/sglang/srt/compilation/compile.py rename to reference/tc_piecewise/compilation/compile.py diff --git a/python/sglang/srt/compilation/compile_phase.py b/reference/tc_piecewise/compilation/compile_phase.py similarity index 100% rename from python/sglang/srt/compilation/compile_phase.py rename to reference/tc_piecewise/compilation/compile_phase.py diff --git a/python/sglang/srt/compilation/compiler_interface.py b/reference/tc_piecewise/compilation/compiler_interface.py similarity index 100% rename from python/sglang/srt/compilation/compiler_interface.py rename to reference/tc_piecewise/compilation/compiler_interface.py diff --git a/python/sglang/srt/compilation/cuda_piecewise_backend.py b/reference/tc_piecewise/compilation/cuda_piecewise_backend.py similarity index 100% rename from python/sglang/srt/compilation/cuda_piecewise_backend.py rename to reference/tc_piecewise/compilation/cuda_piecewise_backend.py diff --git a/python/sglang/srt/compilation/inductor_pass.py b/reference/tc_piecewise/compilation/inductor_pass.py similarity index 100% rename from python/sglang/srt/compilation/inductor_pass.py rename to reference/tc_piecewise/compilation/inductor_pass.py diff --git a/python/sglang/srt/compilation/npu_piecewise_backend.py b/reference/tc_piecewise/compilation/npu_piecewise_backend.py similarity index 100% rename from python/sglang/srt/compilation/npu_piecewise_backend.py rename to reference/tc_piecewise/compilation/npu_piecewise_backend.py diff --git a/python/sglang/srt/compilation/pass_manager.py b/reference/tc_piecewise/compilation/pass_manager.py similarity index 100% rename from python/sglang/srt/compilation/pass_manager.py rename to reference/tc_piecewise/compilation/pass_manager.py diff --git a/python/sglang/srt/compilation/xpu_piecewise_backend.py b/reference/tc_piecewise/compilation/xpu_piecewise_backend.py similarity index 100% rename from python/sglang/srt/compilation/xpu_piecewise_backend.py rename to reference/tc_piecewise/compilation/xpu_piecewise_backend.py diff --git a/python/sglang/srt/model_executor/runner_backend_utils/tc_piecewise_cuda_graph/__init__.py b/reference/tc_piecewise/context/__init__.py similarity index 100% rename from python/sglang/srt/model_executor/runner_backend_utils/tc_piecewise_cuda_graph/__init__.py rename to reference/tc_piecewise/context/__init__.py diff --git a/python/sglang/srt/model_executor/runner_backend_utils/tc_piecewise_cuda_graph/context_manager.py b/reference/tc_piecewise/context/context_manager.py similarity index 100% rename from python/sglang/srt/model_executor/runner_backend_utils/tc_piecewise_cuda_graph/context_manager.py rename to reference/tc_piecewise/context/context_manager.py diff --git a/python/sglang/srt/model_executor/runner_backend/tc_piecewise_cuda_graph_backend.py b/reference/tc_piecewise/tc_piecewise_cuda_graph_backend.py similarity index 100% rename from python/sglang/srt/model_executor/runner_backend/tc_piecewise_cuda_graph_backend.py rename to reference/tc_piecewise/tc_piecewise_cuda_graph_backend.py diff --git a/test/manual/chunked_prefill/test_scripted_hybrid_swa.py b/test/manual/chunked_prefill/test_scripted_hybrid_swa.py index 1dd80b8a986e..654d99d7ec16 100644 --- a/test/manual/chunked_prefill/test_scripted_hybrid_swa.py +++ b/test/manual/chunked_prefill/test_scripted_hybrid_swa.py @@ -19,7 +19,7 @@ class TestSWABasic(ScriptedTestCase): model_path=_SWA_MODEL, chunked_prefill_size=DEFAULT_CHUNK_SIZE, mem_fraction_static=0.70, - disable_piecewise_cuda_graph=True, + cuda_graph_backend_prefill="disabled", ) def test_naive_swa_chunked(self): @@ -112,7 +112,7 @@ class TestSWAHalfWindowChunk(ScriptedTestCase): model_path=_SWA_MODEL, chunked_prefill_size=_SWA_WINDOW // 2, mem_fraction_static=0.70, - disable_piecewise_cuda_graph=True, + cuda_graph_backend_prefill="disabled", ) def test_swa_prompt_2x_window_half_chunks(self): @@ -134,7 +134,7 @@ class TestSWAChunkSizeExceedsWindow(ScriptedTestCase): model_path=_SWA_MODEL, chunked_prefill_size=_SWA_WINDOW * 2, mem_fraction_static=0.70, - disable_piecewise_cuda_graph=True, + cuda_graph_backend_prefill="disabled", ) def test_swa_chunk_size_exceeds_window(self): @@ -155,7 +155,7 @@ class TestSWARadix(ScriptedTestCase): chunked_prefill_size=DEFAULT_CHUNK_SIZE, mem_fraction_static=0.70, disable_radix_cache=False, - disable_piecewise_cuda_graph=True, + cuda_graph_backend_prefill="disabled", ) def test_swa_radix_partial_hit_straddles_window(self): diff --git a/test/manual/chunked_prefill/test_scripted_regression.py b/test/manual/chunked_prefill/test_scripted_regression.py index 4a91e0a01702..eecfdb8d6674 100644 --- a/test/manual/chunked_prefill/test_scripted_regression.py +++ b/test/manual/chunked_prefill/test_scripted_regression.py @@ -496,7 +496,7 @@ class TestRegressionGptOss(ScriptedTestCase): chunked_prefill_size=DEFAULT_CHUNK_SIZE, model_path="openai/gpt-oss-20b", mem_fraction_static=0.70, - disable_piecewise_cuda_graph=True, + cuda_graph_backend_prefill="disabled", ) def test_chunked_stash_bounded_by_kv_committed_len(self): diff --git a/test/manual/nightly/test_vlms_piecewise_cuda_graph.py b/test/manual/nightly/test_vlms_piecewise_cuda_graph.py index 83f911f59f83..a14551e2082c 100644 --- a/test/manual/nightly/test_vlms_piecewise_cuda_graph.py +++ b/test/manual/nightly/test_vlms_piecewise_cuda_graph.py @@ -137,7 +137,7 @@ def _run_vlm_mmmu_test( "--trust-remote-code", "--cuda-graph-max-bs-prefill", "8192", - "--cuda-graph-backend-prefill=tc_piecewise", + "--cuda-graph-backend-prefill=breakable", "--tp=8", "--cuda-graph-tc-compiler=eager", "--disable-radix-cache", diff --git a/test/manual/nightly/test_vlms_vit_cuda_graph.py b/test/manual/nightly/test_vlms_vit_cuda_graph.py index e86dc06c2264..06de10282719 100644 --- a/test/manual/nightly/test_vlms_vit_cuda_graph.py +++ b/test/manual/nightly/test_vlms_vit_cuda_graph.py @@ -140,7 +140,7 @@ def _run_vlm_mmmu_test( other_args=[ "--mm-attention-backend", "fa3", - "--cuda-graph-backend-prefill=tc_piecewise", + "--cuda-graph-backend-prefill=breakable", "--cuda-graph-max-bs-prefill", "8192", "--chunked-prefill-size", diff --git a/test/manual/piecewise_cuda_graph/test_piecewise_cuda_graph_support_1_gpu_archived.py b/test/manual/piecewise_cuda_graph/test_piecewise_cuda_graph_support_1_gpu_archived.py deleted file mode 100644 index 66e8ae784d5b..000000000000 --- a/test/manual/piecewise_cuda_graph/test_piecewise_cuda_graph_support_1_gpu_archived.py +++ /dev/null @@ -1,63 +0,0 @@ -"""Archived test classes split out of test/registered/piecewise_cuda_graph/test_piecewise_cuda_graph_support_1_gpu.py. - -Originally registered with `register_cuda_ci(...)`. Moved here as part of -the per-commit pruning effort to keep the code reachable manually. -Run with `python3 test/manual/piecewise_cuda_graph/test_piecewise_cuda_graph_support_1_gpu_archived.py`. -""" - -import unittest - -from sglang.srt.utils import kill_process_tree -from sglang.test.run_eval import run_eval -from sglang.test.test_utils import ( - DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - DEFAULT_URL_FOR_TEST, - CustomTestCase, - SimpleNamespace, - popen_launch_server, -) - - -# CI Registration -class TestPiecewiseCudaGraphInternVL25(CustomTestCase): - """Test piecewise CUDA graph with InternVL2.5-8B model""" - - @classmethod - def setUpClass(cls): - cls.model = "OpenGVLab/InternVL2_5-8B" - cls.base_url = DEFAULT_URL_FOR_TEST - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - other_args=[ - "--cuda-graph-backend-prefill=tc_piecewise", - "--disable-radix-cache", - ], - ) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - def test_gsm8k_accuracy(self): - args = SimpleNamespace( - base_url=self.base_url, - model=self.model, - eval_name="gsm8k", - num_examples=None, - num_threads=1024, - ) - - metrics = run_eval(args) - print(f"GSM8K Accuracy: {metrics['score']:.3f}") - - # Baseline (no piecewise CUDA graph): 0.571 — this eval uses 5-shot - # concatenated text via chat API, which scores lower than reported - # benchmarks (~77.8%) that use proper CoT chat format. The threshold - # is set 5% below observed to catch catastrophic regressions. - self.assertGreaterEqual(metrics["score"], 0.54) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/manual/piecewise_cudagraph/test_disaggregation_piecewise_cuda_graph.py b/test/manual/piecewise_cudagraph/test_disaggregation_piecewise_cuda_graph.py deleted file mode 100644 index 9b7b2ca24573..000000000000 --- a/test/manual/piecewise_cudagraph/test_disaggregation_piecewise_cuda_graph.py +++ /dev/null @@ -1,87 +0,0 @@ -import unittest -from types import SimpleNamespace - -from sglang.test.server_fixtures.disaggregation_fixture import ( - PDDisaggregationServerBase, -) -from sglang.test.sgl_eval_utils import run_sgl_eval -from sglang.test.test_utils import ( - DEFAULT_MODEL_NAME_FOR_TEST, - DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - popen_launch_pd_server, -) - - -class TestDisaggregationPiecewiseCudaGraph(PDDisaggregationServerBase): - """Test piecewise CUDA graph support in disaggregation prefill server""" - - @classmethod - def setUpClass(cls): - super().setUpClass() - cls.model = DEFAULT_MODEL_NAME_FOR_TEST - - # Start servers - cls.start_prefill() - cls.start_decode() - - # Wait for both to be ready - cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill) - cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode) - - cls.launch_lb() - - @classmethod - def start_prefill(cls): - prefill_args = [ - "--trust-remote-code", - "--disaggregation-mode", - "prefill", - "--tp", - "1", - "--cuda-graph-backend-prefill=tc_piecewise", - ] - prefill_args += cls.transfer_backend + cls.rdma_devices - cls.process_prefill = popen_launch_pd_server( - cls.model, - cls.prefill_url, - timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - other_args=prefill_args, - ) - - @classmethod - def start_decode(cls): - decode_args = [ - "--trust-remote-code", - "--disaggregation-mode", - "decode", - "--tp", - "1", - "--base-gpu-id", - "1", - ] - decode_args += cls.transfer_backend + cls.rdma_devices - cls.process_decode = popen_launch_pd_server( - cls.model, - cls.decode_url, - timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - other_args=decode_args, - ) - - def test_gsm8k_accuracy(self): - """Verify that piecewise cuda graph works correctly in prefill server""" - args = SimpleNamespace( - base_url=self.base_url, - model=self.model, - eval_name="gsm8k", - max_tokens=512, - num_examples=200, - num_threads=128, - ) - metrics = run_sgl_eval(args) - print(f"GSM8K accuracy with piecewise cuda graph: {metrics['score']:.3f}") - - self.assertGreater(metrics["score"], 0.62) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/amd/test_deepseek_r1_mxfp4_8gpu.py b/test/registered/amd/test_deepseek_r1_mxfp4_8gpu.py deleted file mode 100644 index 7a4532be9a01..000000000000 --- a/test/registered/amd/test_deepseek_r1_mxfp4_8gpu.py +++ /dev/null @@ -1,182 +0,0 @@ -import unittest -from types import SimpleNamespace - -import requests - -from sglang.srt.environ import envs -from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_amd_ci -from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k -from sglang.test.send_one import BenchArgs, send_one_prompt -from sglang.test.test_utils import ( - DEFAULT_URL_FOR_TEST, - CustomTestCase, - is_in_ci, - popen_launch_server, - write_github_step_summary, -) - -register_amd_ci(est_time=1300, suite="stage-c-test-large-8-gpu-amd-mi35x") - -DEEPSEEK_R1_MODEL_PATH = "amd/DeepSeek-R1-MXFP4-Preview" -SERVER_LAUNCH_TIMEOUT = 1800 - - -class TestDeepseekR1MXFP4(CustomTestCase): - @classmethod - def setUpClass(cls): - cls.model = DEEPSEEK_R1_MODEL_PATH - cls.base_url = DEFAULT_URL_FOR_TEST - - # Workaround: AITER custom all-gather corrupts CUDA-graph IPC buffer - # registration and triggers a decode-time "Memory access fault" on - # MI35x TP=8. Disable until the AITER-side fix lands (see PR body). - envs.SGLANG_USE_AITER_AG.set(False) - - other_args = [ - "--tp", - "8", - "--chunked-prefill-size", - "131072", - "--model-loader-extra-config", - '{"enable_multithread_load": true}', - "--cuda-graph-backend-prefill=tc_piecewise", - "--cuda-graph-tc-compiler", - "eager", - "--cuda-graph-max-bs-prefill", - "8192", - ] - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=SERVER_LAUNCH_TIMEOUT, - other_args=other_args, - ) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - def test_a_gsm8k( - self, - ): # Append an "a" to make this test run first (alphabetically) to warm up the server - requests.get(self.base_url + "/flush_cache") - - args = SimpleNamespace( - num_shots=8, - data_path=None, - num_questions=1319, - parallel=1319, - max_new_tokens=512, - host="127.0.0.1", - port=int(self.base_url.split(":")[-1]), - ) - metrics = run_eval_few_shot_gsm8k(args) - print(f"{metrics=}") - - if is_in_ci(): - write_github_step_summary( - f'### test_gsm8k (deepseek-r1-mxfp4)\n{metrics["accuracy"]=:.3f}\n' - ) - self.assertGreater(metrics["accuracy"], 0.94) - - def test_bs_1_speed(self): - args = BenchArgs(port=int(self.base_url.split(":")[-1]), max_new_tokens=2048) - _, speed = send_one_prompt(args) - - print(f"{speed=:.2f}") - - if is_in_ci(): - write_github_step_summary( - f"### test_bs_1_speed (deepseek-r1-mxfp4)\n{speed=:.2f} token/s\n" - ) - self.assertGreater(speed, 75) - - -class TestDeepseekR1MXFP4MTP(CustomTestCase): - @classmethod - def setUpClass(cls): - cls.model = DEEPSEEK_R1_MODEL_PATH - cls.base_url = DEFAULT_URL_FOR_TEST - - envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.set(True) - # Same AITER custom all-gather workaround as TestDeepseekR1MXFP4 above. - envs.SGLANG_USE_AITER_AG.set(False) - - other_args = [ - "--tp", - "8", - "--chunked-prefill-size", - "131072", - "--speculative-algorithm", - "EAGLE", - "--speculative-num-steps", - "3", - "--speculative-eagle-topk", - "1", - "--speculative-num-draft-tokens", - "4", - "--model-loader-extra-config", - '{"enable_multithread_load": true}', - ] - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=SERVER_LAUNCH_TIMEOUT, - other_args=other_args, - ) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - def test_a_gsm8k( - self, - ): # Append an "a" to make this test run first (alphabetically) to warm up the server - requests.get(self.base_url + "/flush_cache") - - args = SimpleNamespace( - num_shots=5, - data_path=None, - num_questions=200, - max_new_tokens=512, - parallel=128, - host="127.0.0.1", - port=int(self.base_url.split(":")[-1]), - ) - metrics = run_eval_few_shot_gsm8k(args) - print(f"{metrics=}") - - server_info = requests.get(self.base_url + "/server_info") - avg_spec_accept_length = server_info.json()["internal_states"][0][ - "avg_spec_accept_length" - ] - print(f"{avg_spec_accept_length=}") - - if is_in_ci(): - write_github_step_summary( - f"### test_gsm8k (deepseek-r1-mxfp4 mtp)\n" - f'{metrics["accuracy"]=:.3f}\n' - f"{avg_spec_accept_length=:.2f}\n" - ) - self.assertGreater(metrics["accuracy"], 0.94) - self.assertGreater(avg_spec_accept_length, 2.04) - - def test_bs_1_speed(self): - args = BenchArgs(port=int(self.base_url.split(":")[-1]), max_new_tokens=2048) - acc_length, speed = send_one_prompt(args) - - print(f"{acc_length=:.2f} {speed=:.2f}") - - if is_in_ci(): - write_github_step_summary( - f"### test_bs_1_speed (deepseek-r1-mxfp4 mtp)\n" - f"{acc_length=:.2f}\n" - f"{speed=:.2f} token/s\n" - ) - self.assertGreater(acc_length, 2.04) - self.assertGreater(speed, 150) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/attention/unittests/dense/test_triton.py b/test/registered/attention/unittests/dense/test_triton.py index 64559579caa9..75b21a5690f3 100644 --- a/test/registered/attention/unittests/dense/test_triton.py +++ b/test/registered/attention/unittests/dense/test_triton.py @@ -366,7 +366,7 @@ def test_runner_mode_cuda_graph_decode_cases(self): @unittest.skipIf( is_hip(), "split-op extend runner exercises the piecewise-CUDA-graph path " - "(TcPiecewiseForwardContext.num_tokens), which is not wired on ROCm.", + "(prefill runner token buffers), which is not wired on ROCm.", ) def test_runner_mode_split_op_extend_cases(self): for case, static_num_tokens in self.SPLIT_OP_CASES: diff --git a/test/registered/attention/unittests/dsa/README.md b/test/registered/attention/unittests/dsa/README.md index d629a5578fd5..70bad7759918 100644 --- a/test/registered/attention/unittests/dsa/README.md +++ b/test/registered/attention/unittests/dsa/README.md @@ -92,7 +92,7 @@ hardware/SDK. The variant tests live in `test_dsa.py` as enabling default tests. - Runner-mode integration is now plumbed at the fixture level: `DSAMockModelRunner` accepts `disable_cuda_graph`, - `disable_piecewise_cuda_graph`, and `runner_batch_size` kwargs; + `disable_prefill_cuda_graph`, and `runner_batch_size` kwargs; `build_dsa_attention_fixture` passes them through; and `dsa_attention.py` exposes the standard adapter callbacks (`make_dsa_case_with_prefix_lens`, `dsa_fixture_inputs`, diff --git a/test/registered/attention/unittests/gdn/test_torch_native.py b/test/registered/attention/unittests/gdn/test_torch_native.py index ef516ab5b7ee..92a4d2cc895a 100644 --- a/test/registered/attention/unittests/gdn/test_torch_native.py +++ b/test/registered/attention/unittests/gdn/test_torch_native.py @@ -78,7 +78,7 @@ def test_layout_robustness_cases(self): @unittest.skipIf( is_hip(), "split-op extend runner exercises the piecewise-CUDA-graph path " - "(TcPiecewiseForwardContext.num_tokens), which is not wired on ROCm.", + "(prefill runner token buffers), which is not wired on ROCm.", ) def test_runner_mode_split_op_extend_cases(self): for case, static_num_tokens in self.SPLIT_OP_CASES: diff --git a/test/registered/attention/unittests/gdn/test_triton.py b/test/registered/attention/unittests/gdn/test_triton.py index 129e76a5c806..8b90643e95d1 100644 --- a/test/registered/attention/unittests/gdn/test_triton.py +++ b/test/registered/attention/unittests/gdn/test_triton.py @@ -289,7 +289,7 @@ def test_runner_mode_cuda_graph_decode_cases(self): @unittest.skipIf( is_hip(), "split-op extend runner exercises the piecewise-CUDA-graph path " - "(TcPiecewiseForwardContext.num_tokens), which is not wired on ROCm.", + "(prefill runner token buffers), which is not wired on ROCm.", ) def test_runner_mode_split_op_extend_cases(self): for case, static_num_tokens in self.SPLIT_OP_CASES: diff --git a/test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py b/test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py index 9a078c277227..890ac15492ed 100644 --- a/test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py +++ b/test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py @@ -84,7 +84,7 @@ def _build_hybrid_backend(testcase, case: MLAAttentionCase): kv_lora_rank=_KV_LORA_RANK, qk_rope_head_dim=_QK_ROPE_HEAD_DIM, disable_cuda_graph=True, - disable_piecewise_cuda_graph=True, + disable_prefill_cuda_graph=True, runner_batch_size=None, fp8_kv_cache=False, ) diff --git a/test/registered/attention/unittests/kda/test_triton.py b/test/registered/attention/unittests/kda/test_triton.py index d9f24bbabd60..82dc20a7adc9 100644 --- a/test/registered/attention/unittests/kda/test_triton.py +++ b/test/registered/attention/unittests/kda/test_triton.py @@ -235,7 +235,7 @@ def test_runner_mode_eagle_verify_cuda_graph_cases(self): @unittest.skipIf( is_hip(), "split-op extend runner exercises the piecewise-CUDA-graph path " - "(TcPiecewiseForwardContext.num_tokens), which is not wired on ROCm.", + "(prefill runner token buffers), which is not wired on ROCm.", ) def test_runner_mode_split_op_extend_cases(self): for case, static_num_tokens in self.SPLIT_OP_CASES: diff --git a/test/registered/attention/unittests/swa/test_triton.py b/test/registered/attention/unittests/swa/test_triton.py index aeeb84c08527..00335cdc0ad0 100644 --- a/test/registered/attention/unittests/swa/test_triton.py +++ b/test/registered/attention/unittests/swa/test_triton.py @@ -297,7 +297,7 @@ def test_runner_mode_cuda_graph_decode_cases(self): @unittest.skipIf( is_hip(), "split-op extend runner exercises the piecewise-CUDA-graph path " - "(TcPiecewiseForwardContext.num_tokens), which is not wired on ROCm.", + "(prefill runner token buffers), which is not wired on ROCm.", ) def test_runner_mode_split_op_extend_cases(self): for case, static_num_tokens in self.SPLIT_OP_CASES: diff --git a/test/registered/cp/test_cp_strategy_unit.py b/test/registered/cp/test_cp_strategy_unit.py index 70da8f1e58ad..d3d934593152 100644 --- a/test/registered/cp/test_cp_strategy_unit.py +++ b/test/registered/cp/test_cp_strategy_unit.py @@ -131,7 +131,7 @@ def _make_runner(self): runner._is_full_backend = False runner.enable_lora = False runner._capture_chunked_prefix = False - runner.prefill_backend_name = Backend.TC_PIECEWISE + runner.prefill_backend_name = Backend.BREAKABLE runner.has_mha_companion_layers = False runner.capture_hidden_mode = CaptureHiddenMode.NULL runner.capture_num_tokens = [2048, 2304] diff --git a/test/registered/cuda_graph/breakable/test_bcg_with_speculative_decoding.py b/test/registered/cuda_graph/breakable/test_bcg_with_speculative_decoding.py index b67bc159ade9..fcde0e93a8cc 100644 --- a/test/registered/cuda_graph/breakable/test_bcg_with_speculative_decoding.py +++ b/test/registered/cuda_graph/breakable/test_bcg_with_speculative_decoding.py @@ -1,7 +1,5 @@ """Test breakable CUDA graph (BCG) coexisting with EAGLE3 speculative -decoding. Sibling of test_pcg_with_speculative_decoding.py — same -target/draft pair, only flips the prefill backend from tc_piecewise to -breakable. Verifies the draft-side BCG plumbing in PrefillCudaGraphRunner +decoding. Verifies the draft-side BCG plumbing in PrefillCudaGraphRunner stays wired (capture_hidden_mode for EAGLE, static_draft_hidden_states buffer sized from the draft's fc input, EagleDraftInput at capture, and the load_batch refresh). diff --git a/test/registered/cuda_graph/breakable/test_breakable_cuda_graph.py b/test/registered/cuda_graph/breakable/test_breakable_cuda_graph.py index 9c9faa1eb00d..399079477349 100644 --- a/test/registered/cuda_graph/breakable/test_breakable_cuda_graph.py +++ b/test/registered/cuda_graph/breakable/test_breakable_cuda_graph.py @@ -8,7 +8,6 @@ """ import unittest -from unittest.mock import patch import torch @@ -69,7 +68,7 @@ def test_single_break(self): intermediate = torch.zeros(4, device=self.device) y = torch.zeros(4, device=self.device) - @self.eager_on_graph(enable=True) + @self.eager_on_graph def eager_op(src): return src * 2.0 @@ -92,11 +91,11 @@ def test_multiple_breaks(self): x = torch.zeros(4, device=self.device) y = torch.zeros(4, device=self.device) - @self.eager_on_graph(enable=True) + @self.eager_on_graph def add_one(src): return src + 1.0 - @self.eager_on_graph(enable=True) + @self.eager_on_graph def double(src): return src * 2.0 @@ -115,24 +114,10 @@ def double(src): torch.cuda.synchronize() self.assertTrue(torch.allclose(y, torch.full((4,), 16.0, device=self.device))) - def test_eager_on_graph_disabled(self): - """@eager_on_graph(enable=False) should be a no-op passthrough.""" - - @self.eager_on_graph(enable=False) - def my_fn(x): - return x + 1.0 - - # Should just be the original function - t = torch.tensor([1.0, 2.0], device=self.device) - result = my_fn(t) - self.assertTrue( - torch.allclose(result, torch.tensor([2.0, 3.0], device=self.device)) - ) - def test_eager_on_graph_outside_capture(self): """@eager_on_graph called outside capture should run the function directly.""" - @self.eager_on_graph(enable=True) + @self.eager_on_graph def my_fn(x): return x + 1.0 @@ -147,7 +132,7 @@ def test_replay_updates_output(self): x = torch.zeros(4, device=self.device) y = torch.zeros(4, device=self.device) - @self.eager_on_graph(enable=True) + @self.eager_on_graph def scale(src): return src * 3.0 @@ -175,7 +160,7 @@ def test_side_stream_join_across_break(self): y = torch.zeros(4, device=self.device) stream = torch.cuda.Stream(self.device) - @self.eager_on_graph(enable=True) + @self.eager_on_graph def identity(src): return src @@ -199,7 +184,7 @@ def test_eager_output_is_held_strongly_for_replay_bridge(self): x = torch.zeros(4, device=self.device) y = torch.zeros(4, device=self.device) - @self.eager_on_graph(enable=True) + @self.eager_on_graph def scale(src): return src * 3.0 @@ -216,56 +201,6 @@ def scale(src): "eager output bridge buffer must be strongly captured", ) - def test_attention_narrows_padded_positions(self): - from sglang.srt.layers.radix_attention import unified_attention_with_output - - num_tokens = 3 - padded_num_tokens = 5 - forward_batch = SimpleNamespace( - global_num_token_non_padded_cpu=num_tokens, - out_cache_loc=torch.arange(padded_num_tokens, device=self.device), - positions=torch.arange(padded_num_tokens, device=self.device), - ) - context = SimpleNamespace( - forward_batch=forward_batch, - attention_layers=[object()], - mha_companion_layers=None, - num_tokens=padded_num_tokens, - raw_num_tokens=num_tokens, - ) - observed = {} - - def attention_forward(query, key, value, layer, batch, save_kv_cache): - observed["positions"] = batch.positions.clone() - observed["out_cache_loc"] = batch.out_cache_loc.clone() - return torch.ones_like(query) - - output = torch.full((padded_num_tokens, 2), float("nan"), device=self.device) - with ( - patch( - "sglang.srt.layers.radix_attention.get_tc_piecewise_forward_context", - return_value=context, - ), - patch( - "sglang.srt.layers.radix_attention.get_attn_backend", - return_value=SimpleNamespace(forward=attention_forward), - ), - ): - unified_attention_with_output( - torch.zeros((padded_num_tokens, 2), device=self.device), - torch.zeros((padded_num_tokens, 1, 2), device=self.device), - torch.zeros((padded_num_tokens, 1, 2), device=self.device), - output, - True, - 0, - ) - - expected = torch.arange(num_tokens, device=self.device) - torch.testing.assert_close(observed["positions"], expected) - torch.testing.assert_close(observed["out_cache_loc"], expected) - self.assertEqual(forward_batch.positions.shape[0], padded_num_tokens) - self.assertEqual(forward_batch.out_cache_loc.shape[0], padded_num_tokens) - class TestCopyOutput(CustomTestCase): """Test the _copy_output helper for structured output writeback.""" diff --git a/test/registered/cuda_graph/piecewise/test_pcg_with_speculative_decoding.py b/test/registered/cuda_graph/piecewise/test_pcg_with_speculative_decoding.py deleted file mode 100644 index ebb57687030d..000000000000 --- a/test/registered/cuda_graph/piecewise/test_pcg_with_speculative_decoding.py +++ /dev/null @@ -1,45 +0,0 @@ -"""Test piecewise CUDA graph coexisting with speculative decoding (EAGLE3). - -PCG handles prefill/extend path while speculative decoding (EAGLE3) uses -decode CUDA graphs. This test verifies they don't interfere with each -other. MTP / STANDALONE / NGRAM variants moved to the sibling file -test_pcg_with_speculative_decoding_extra.py. -""" - -import unittest - -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.server_fixtures.pcg_spec_fixture import PCGSpecBase - -register_cuda_ci(est_time=130, stage="weekly", runner_config="4-gpu-h100") - - -class TestPCGWithEAGLE3(PCGSpecBase, unittest.TestCase): - """PCG + EAGLE3 on Qwen3-30B-A3B-Instruct-2507.""" - - model = "Qwen/Qwen3-30B-A3B-Instruct-2507" - server_args = [ - "--tp", - "2", - "--trust-remote-code", - "--cuda-graph-backend-prefill=tc_piecewise", - "--mem-fraction-static", - "0.6", - "--speculative-algorithm", - "EAGLE3", - "--speculative-draft-model-path", - "lmsys/SGLang-EAGLE3-Qwen3-30B-A3B-Instruct-2507-SpecForge-Nex", - "--speculative-num-steps", - "5", - "--speculative-eagle-topk", - "4", - "--speculative-num-draft-tokens", - "8", - ] - timeout_mult = 3 - server_env = {"SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN": "1"} - accuracy_threshold = 0.75 - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/cuda_graph/piecewise/test_pcg_with_speculative_decoding_dflash.py b/test/registered/cuda_graph/piecewise/test_pcg_with_speculative_decoding_dflash.py deleted file mode 100644 index 1a1a6412a7e5..000000000000 --- a/test/registered/cuda_graph/piecewise/test_pcg_with_speculative_decoding_dflash.py +++ /dev/null @@ -1,52 +0,0 @@ -"""Test piecewise CUDA graph coexisting with speculative decoding (DFLASH). - -PCG handles prefill/extend path while DFlash needs target aux hidden states -from prefill to materialize draft KV cache. This verifies PCG captures that -path with the DFlash hidden-state variant enabled. -""" - -import unittest - -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.server_fixtures.pcg_spec_fixture import PCGSpecBase -from sglang.test.test_utils import ( - DEFAULT_DRAFT_MODEL_DFLASH, - DEFAULT_TARGET_MODEL_DFLASH, - CustomTestCase, -) - -register_cuda_ci(est_time=110, stage="weekly", runner_config="1-gpu-large") - - -class TestPCGWithDFlash(PCGSpecBase, CustomTestCase): - """PCG + DFLASH on Llama-3.1-8B-Instruct.""" - - model = DEFAULT_TARGET_MODEL_DFLASH - server_args = [ - "--trust-remote-code", - "--attention-backend", - "flashinfer", - "--cuda-graph-backend-prefill", - "tc_piecewise", - "--speculative-algorithm", - "DFLASH", - "--speculative-draft-model-path", - DEFAULT_DRAFT_MODEL_DFLASH, - "--page-size", - "1", - "--max-running-requests", - "64", - # Keep headroom for the draft KV pool + piecewise cuda graph - # private pools on 32GB CI cards. - "--mem-fraction-static", - "0.7", - "--cuda-graph-bs-decode", - *[str(i) for i in range(1, 65)], - ] - server_env = {"SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN": "1"} - accuracy_threshold = 0.75 - speedup_threshold = 2.8 - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/cuda_graph/piecewise/test_pcg_with_speculative_decoding_extra.py b/test/registered/cuda_graph/piecewise/test_pcg_with_speculative_decoding_extra.py deleted file mode 100644 index 4b16fef8c8b3..000000000000 --- a/test/registered/cuda_graph/piecewise/test_pcg_with_speculative_decoding_extra.py +++ /dev/null @@ -1,79 +0,0 @@ -"""Extra: PCG coexistence with non-EAGLE3 speculative decoding variants. - -EAGLE3 lives in the sibling file test_pcg_with_speculative_decoding.py. -""" - -import unittest - -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.server_fixtures.pcg_spec_fixture import PCGSpecBase - -register_cuda_ci(est_time=450, stage="weekly", runner_config="4-gpu-h100") - - -class TestPCGWithMTP(PCGSpecBase, unittest.TestCase): - """PCG + MTP (NEXTN) on Qwen3.5-35B-A3B with FP8.""" - - model = "Qwen/Qwen3.5-35B-A3B" - server_args = [ - "--tp", - "2", - "--trust-remote-code", - "--quantization", - "fp8", - "--mamba-radix-cache-strategy", - "extra_buffer", - "--speculative-algorithm", - "NEXTN", - "--reasoning-parser", - "qwen3", - ] - timeout_mult = 3 - max_tokens = 8192 - thinking_mode = "qwen3" - accuracy_threshold = 0.75 - - -class TestPCGWithSTANDALONE(PCGSpecBase, unittest.TestCase): - """PCG + STANDALONE on Llama-3.1-8B-Instruct + Llama-3.2-1B-Instruct.""" - - model = "meta-llama/Llama-3.1-8B-Instruct" - server_args = [ - "--trust-remote-code", - "--cuda-graph-backend-prefill=tc_piecewise", - "--mem-fraction-static", - "0.5", - "--speculative-algorithm", - "STANDALONE", - "--speculative-draft-model-path", - "meta-llama/Llama-3.2-1B-Instruct", - "--speculative-num-steps", - "3", - "--speculative-eagle-topk", - "1", - "--speculative-num-draft-tokens", - "4", - ] - accuracy_threshold = 0.50 - - -class TestPCGWithNGRAM(PCGSpecBase, unittest.TestCase): - """PCG + NGRAM on Qwen2.5-Coder-7B-Instruct.""" - - model = "Qwen/Qwen2.5-Coder-7B-Instruct" - server_args = [ - "--trust-remote-code", - "--cuda-graph-backend-prefill=tc_piecewise", - "--speculative-algorithm", - "NGRAM", - "--speculative-num-draft-tokens", - "16", - "--cuda-graph-max-bs-decode", - "8", - "--mem-fraction-static", - "0.8", - ] - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/cuda_graph/piecewise/test_piecewise_cuda_graph_support_1_gpu.py b/test/registered/cuda_graph/piecewise/test_piecewise_cuda_graph_support_1_gpu.py deleted file mode 100644 index b9d334251da8..000000000000 --- a/test/registered/cuda_graph/piecewise/test_piecewise_cuda_graph_support_1_gpu.py +++ /dev/null @@ -1,131 +0,0 @@ -import unittest - -import torch -from transformers import AutoProcessor - -from sglang import Engine -from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci -from sglang.test.run_eval import run_eval -from sglang.test.test_utils import ( - DEFAULT_IMAGE_URL, - DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - DEFAULT_URL_FOR_TEST, - CustomTestCase, - SimpleNamespace, - build_vlm_image_prompt, - is_in_amd_ci, - popen_launch_server, -) - -# CI Registration -register_cuda_ci(est_time=250, stage="nightly", runner_config="1-gpu-large") -register_amd_ci(est_time=180, suite="stage-b-test-1-gpu-large-amd") - - -# The 192GB mi300x runners have less headroom than the 256GB mi325x ones they -# replaced: the auto-derived fraction left too little room for the ViT -# activations plus the piecewise graph private pools, and the server died under -# the 1024-thread gsm8k load. -AMD_MEM_FRACTION_STATIC = 0.6 - - -class TestPiecewiseCudaGraphQwen25VL(CustomTestCase): - """Test piecewise CUDA graph with Qwen2.5-VL-7B-Instruct model""" - - @classmethod - def setUpClass(cls): - cls.model = "Qwen/Qwen2.5-VL-7B-Instruct" - cls.base_url = DEFAULT_URL_FOR_TEST - other_args = [ - "--cuda-graph-backend-prefill=tc_piecewise", - "--disable-radix-cache", - ] - if is_in_amd_ci(): - other_args += ["--mem-fraction-static", str(AMD_MEM_FRACTION_STATIC)] - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - other_args=other_args, - ) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - def test_gsm8k_accuracy(self): - args = SimpleNamespace( - base_url=self.base_url, - model=self.model, - eval_name="gsm8k", - num_examples=None, - num_threads=1024, - ) - - metrics = run_eval(args) - print(f"GSM8K Accuracy: {metrics['score']:.3f}") - - self.assertGreaterEqual(metrics["score"], 0.80) - - -class TestPiecewiseCudaGraphQwen25VLEmbedding(CustomTestCase): - """Test piecewise CUDA graph with Qwen2.5-VL-3B-Instruct embedding model""" - - def test_embedding(self): - model_path = "Qwen/Qwen2.5-VL-3B-Instruct" - text = build_vlm_image_prompt( - AutoProcessor.from_pretrained(model_path), "What is in this picture?" - ) - extra_args = ( - {"mem_fraction_static": AMD_MEM_FRACTION_STATIC} if is_in_amd_ci() else {} - ) - - engine = Engine( - model_path=model_path, - enable_multimodal=True, - is_embedding=True, - cuda_graph_backend_prefill="tc_piecewise", - **extra_args, - ) - out = engine.encode([text], image_data=[DEFAULT_IMAGE_URL])[0]["embedding"] - engine.shutdown() - self.assertGreater(len(out), 0) - - engine = Engine( - model_path=model_path, - enable_multimodal=True, - is_embedding=True, - cuda_graph_backend_prefill="disabled", - **extra_args, - ) - out_without_pcg = engine.encode([text], image_data=[DEFAULT_IMAGE_URL])[0][ - "embedding" - ] - engine.shutdown() - self.assertGreater(len(out_without_pcg), 0) - - t_out = torch.tensor(out) - t_out_without_pcg = torch.tensor(out_without_pcg) - max_abs_diff = (t_out - t_out_without_pcg).abs().max().item() - max_rel_diff = ( - ((t_out - t_out_without_pcg).abs() / (t_out_without_pcg.abs() + 1e-8)) - .max() - .item() - ) - print( - f"PCG embedding diff: max_abs={max_abs_diff:.6f}, max_rel={max_rel_diff:.6f}" - ) - self.assertTrue( - torch.allclose( - t_out, - t_out_without_pcg, - atol=1e-2, - rtol=1e-2, - ), - f"Piecewise CUDA graph embedding mismatch: max_abs_diff={max_abs_diff}, max_rel_diff={max_rel_diff}", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/cuda_graph/test_cuda_piecewise_backend.py b/test/registered/cuda_graph/test_cuda_piecewise_backend.py deleted file mode 100644 index 7efb5f78de92..000000000000 --- a/test/registered/cuda_graph/test_cuda_piecewise_backend.py +++ /dev/null @@ -1,56 +0,0 @@ -from unittest.mock import MagicMock - -import pytest -import torch - -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=11, stage="base-b", runner_config="1-gpu-small") - -if not torch.cuda.is_available(): - pytest.skip("CUDA piecewise backend requires CUDA", allow_module_level=True) - -import sglang.srt.compilation.cuda_piecewise_backend as cuda_backend -from sglang.srt.compilation.cuda_piecewise_backend import ( - ConcreteSizeEntry, - CUDAPiecewiseBackend, -) - - -def test_runtime_recompile_without_capture_stream_falls_back(monkeypatch): - """A Dynamo replacement backend cannot capture outside a PCG session.""" - compile_config = MagicMock() - compile_config.get_capture_sizes.return_value = [] - backend = CUDAPiecewiseBackend( - graph=MagicMock(), - compile_config=compile_config, - inductor_config={}, - graph_pool=None, - piecewise_compile_index=0, - total_piecewise_compiles=1, - sym_shape_indices=[0], - compiled_graph_for_general_shape=MagicMock(), - sglang_backend=MagicMock(), - ) - backend.first_run_finished = True - fallback = MagicMock(return_value="fallback-result") - backend.concrete_size_entries = { - 4: ConcreteSizeEntry( - runtime_shape=4, - need_to_compile=False, - use_cudagraph=True, - runnable=fallback, - num_finished_warmup=1, - ) - } - - monkeypatch.setattr(cuda_backend, "get_pcg_capture_stream", lambda: None) - monkeypatch.setattr(cuda_backend, "is_in_torch_compile_warmup", lambda: False) - monkeypatch.setattr(cuda_backend, "print_warning_once", lambda _message: None) - - assert backend(4) == "fallback-result" - fallback.assert_called_once_with(4) - - -if __name__ == "__main__": - raise SystemExit(pytest.main([__file__, "-v"])) diff --git a/test/registered/npu/basic_function/optimization_debug/test_npu_piecewise_graph_prefill.py b/test/registered/npu/basic_function/optimization_debug/test_npu_piecewise_graph_prefill.py deleted file mode 100644 index 1db72e5fb438..000000000000 --- a/test/registered/npu/basic_function/optimization_debug/test_npu_piecewise_graph_prefill.py +++ /dev/null @@ -1,77 +0,0 @@ -import subprocess -import unittest - -from sglang.test.ascend.gsm8k_ascend_mixin import GSM8KAscendMixin -from sglang.test.ascend.test_ascend_utils import ( - QWEN2_5_7B_INSTRUCT_WEIGHTS_PATH, - write_results_to_github_step_summary, -) -from sglang.test.ci.ci_register import register_npu_ci -from sglang.test.test_utils import ( - CustomTestCase, - run_bench_one_batch, -) - -register_npu_ci(est_time=400, suite="base-b-test-1-npu-a3") -register_npu_ci(est_time=400, suite="nightly-1-npu-a3", nightly=True) - - -TOKENS_TO_CAPTURE = [i for i in range(128, 4096, 128)] - - -class TestPiecewiseGraphPrefillCorrectness(GSM8KAscendMixin, CustomTestCase): - model = QWEN2_5_7B_INSTRUCT_WEIGHTS_PATH - other_args = [ - "--trust-remote-code", - "--mem-fraction-static", - 0.8, - "--attention-backend", - "ascend", - "--cuda-graph-bs-decode", - 128, - "--cuda-graph-backend-prefill=tc_piecewise", - "--cuda-graph-bs-prefill", - *TOKENS_TO_CAPTURE, - ] - accuracy = 0.84 - num_questions = 1319 - - -class TestPiecewiseGraphPrefillBenchmark(CustomTestCase): - model = QWEN2_5_7B_INSTRUCT_WEIGHTS_PATH - other_args = [ - "--trust-remote-code", - "--mem-fraction-static", - 0.8, - "--attention-backend", - "ascend", - "--cuda-graph-backend-prefill=tc_piecewise", - "--cuda-graph-bs-prefill", - ] + TOKENS_TO_CAPTURE - - latency = 0.045 - - def test_latency(self): - print(f"##=== Testing prefill latency: {self.model} ===##") - model_metrics = { - "server": subprocess.list2cmdline(map(str, self.other_args)), - "client": "bench_one_batch", - "latency_threshold": self.latency, - } - try: - prefill_latency, _, _ = run_bench_one_batch( - self.model, - other_args=self.other_args, - ) - model_metrics["latency"] = float(prefill_latency) - self.assertLess(prefill_latency, self.latency) - except Exception as e: - model_metrics["error"] = e - print(f"Error testing {self.model}: {e}") - self.fail(f"Test failed for {self.model}: {e}") - finally: - write_results_to_github_step_summary({self.model: model_metrics}) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/compilation/test_torch_compile_decoration.py b/test/registered/unit/compilation/test_torch_compile_decoration.py index 7d94466660c7..ce54327d821e 100644 --- a/test/registered/unit/compilation/test_torch_compile_decoration.py +++ b/test/registered/unit/compilation/test_torch_compile_decoration.py @@ -1,78 +1,52 @@ +"""The retained torch.compile entry point is independent of TCPCG.""" + from types import SimpleNamespace from unittest.mock import patch import pytest import torch +from sglang.srt.compilation import torch_compile_decoration as decoration from sglang.test.ci.ci_register import register_cpu_ci -register_cpu_ci(est_time=8, suite="base-a-test-cpu") - -from sglang.srt.compilation.compile import ( - _infer_dynamic_arg_dims_from_annotations, - _mark_dynamic_forward_batch, - _runtime_dynamic_dim_for_argument, -) - - -class _MropeModel: - def forward( - self, - input_ids: torch.Tensor, - positions: torch.Tensor, - forward_batch, - ): - return input_ids, positions, forward_batch - - -class _StringAnnotatedMropeModel: - def forward( - self, - input_ids: "torch.Tensor", - positions: "torch.Tensor", - forward_batch, - ): - return input_ids, positions, forward_batch +register_cpu_ci(est_time=3, suite="base-a-test-cpu") -def test_positions_marks_the_token_axis_dynamic_for_mrope_and_1d_rope(): - dynamic_dims = _infer_dynamic_arg_dims_from_annotations(_MropeModel.forward) - string_dynamic_dims = _infer_dynamic_arg_dims_from_annotations( - _StringAnnotatedMropeModel.forward - ) +class Model(torch.nn.Module): + def forward(self, x): + return x + 2 - assert dynamic_dims["input_ids"] == 0 - assert dynamic_dims["positions"] == -1 - assert string_dynamic_dims["positions"] == -1 +def test_disabled_compile_keeps_raw_forward(): + model = Model() + with patch.object(torch, "compile") as compile_model: + with decoration.patch_model( + model, False, 4, SimpleNamespace(ca_comm=None) + ) as forward: + assert forward(torch.tensor(3)).item() == 5 + compile_model.assert_not_called() -def test_runtime_dynamic_dim_uses_the_token_axis_for_mrope_metadata(): - assert _runtime_dynamic_dim_for_argument("positions") == -1 - assert _runtime_dynamic_dim_for_argument("position_ids") == -1 - assert _runtime_dynamic_dim_for_argument("mrope_positions") == -1 - assert _runtime_dynamic_dim_for_argument("input_ids") == 0 - -def test_forward_batch_marks_token_and_batch_metadata_dynamic(): - batch = SimpleNamespace( - input_embeds=torch.empty(8, 16), - seq_lens=torch.empty(2, dtype=torch.int64), - mrope_positions=torch.empty(3, 8, dtype=torch.int64), - scalar=torch.tensor(1), - ) - marked = [] - - def record_mark_dynamic(value, dims): - marked.append((id(value), tuple(dims))) - - with patch("torch._dynamo.maybe_mark_dynamic", side_effect=record_mark_dynamic): - _mark_dynamic_forward_batch(batch) - - assert (id(batch.input_embeds), (0,)) in marked - assert (id(batch.seq_lens), (0,)) in marked - assert (id(batch.mrope_positions), (1,)) in marked - assert all(value_id != id(batch.scalar) for value_id, _ in marked) - - -if __name__ == "__main__": - raise SystemExit(pytest.main([__file__, "-v"])) +@pytest.mark.parametrize("raises", [False, True]) +def test_compile_restores_module_policy_and_communicator(raises): + model = Model() + communicator = object() + group = SimpleNamespace(ca_comm=communicator) + with ( + patch.object(decoration, "_to_torch") as convert, + patch.object( + torch, "compile", side_effect=lambda forward, **kwargs: forward + ) as compile_model, + ): + try: + with decoration.patch_model(model, True, 4, group) as forward: + assert forward(torch.tensor(3)).item() == 5 + group.ca_comm = None + if raises: + raise RuntimeError("forward failed") + except RuntimeError: + assert raises + assert group.ca_comm is communicator + assert convert.call_args_list[0].kwargs == {"reverse": False, "num_tokens": 4} + assert convert.call_args_list[1].kwargs == {"reverse": True, "num_tokens": 4} + compile_model.assert_called_once() diff --git a/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py b/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py index b76e82bb7022..a16f3830c367 100644 --- a/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py +++ b/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py @@ -15,7 +15,6 @@ from sglang.srt.model_executor.cuda_graph_config import ( Backend, CudaGraphConfig, - Phase, PhaseConfig, ) from sglang.srt.model_executor.forward_batch_info import ( @@ -75,34 +74,9 @@ def test_unknown_multimodal_arch_is_not_opted_in(self): ) ) - def test_supported_multimodal_model_upgrades_default_to_tc_piecewise(self): - args = ServerArgs(model_path="dummy") - args._model_config = SimpleNamespace( - is_multimodal_piecewise_cuda_graph_supported=True, - is_multimodal_breakable_cuda_graph_supported=False, - ) - args.cuda_graph_config = CudaGraphConfig( - prefill=PhaseConfig(backend=Backend.BREAKABLE) - ) - args._cuda_graph_config_locked = set() - - args.attention_backend = "fa3" - - with patch( - "sglang.srt.arg_groups.cuda_graph_hook" - ".disable_tc_piecewise_cudagraph_if_incompatible" - ) as disable_if_incompatible: - apply_cuda_graph_compatibility(args) - - self.assertEqual( - resolution_result(args, "cuda_graph_config").prefill.backend, - Backend.TC_PIECEWISE, - ) - disable_if_incompatible.assert_called_once() - def test_trtllm_mla_stays_on_breakable(self): args = ServerArgs(model_path="dummy") - # trtllm_mla skips the tc_piecewise upgrade and keeps breakable, which + # trtllm_mla keeps breakable, which # now serves MLA by falling back to the flashinfer MLA impl for extend. args._model_config = SimpleNamespace( is_multimodal_piecewise_cuda_graph_supported=True, @@ -125,27 +99,6 @@ def test_trtllm_mla_stays_on_breakable(self): Backend.BREAKABLE, ) - def test_explicit_tc_piecewise_overrides_trtllm_mla_default(self): - args = ServerArgs(model_path="dummy") - args.cuda_graph_config = CudaGraphConfig( - prefill=PhaseConfig(backend=Backend.TC_PIECEWISE) - ) - args._cuda_graph_config_locked = {(Phase.PREFILL, "backend")} - - args.attention_backend = "trtllm_mla" - - apply_cuda_graph_compatibility(args) - - self.assertEqual( - resolution_result(args, "cuda_graph_config").prefill.backend, - Backend.TC_PIECEWISE, - ) - - def test_multimodal_inputs_keep_tc_piecewise_prefill_enabled(self): - runner = self._make_prefill_runner(Backend.TC_PIECEWISE) - - self.assertTrue(runner.can_run_graph(self._make_multimodal_forward_batch())) - def test_multimodal_inputs_keep_breakable_prefill_enabled(self): runner = self._make_prefill_runner(Backend.BREAKABLE) @@ -172,7 +125,7 @@ def test_embedding_gemma_forces_breakable_prefill(self): ) args.cuda_graph_config = CudaGraphConfig( decode=PhaseConfig(backend=Backend.FULL), - prefill=PhaseConfig(backend=Backend.TC_PIECEWISE), + prefill=PhaseConfig(backend=Backend.FULL), ) args.disable_radix_cache = False args.chunked_prefill_size = 2048 diff --git a/test/registered/unit/layer_boundary/test_cutedsl_ar_fusion.py b/test/registered/unit/layer_boundary/test_cutedsl_ar_fusion.py index 6abfb5ae291f..87fe87a0d930 100644 --- a/test/registered/unit/layer_boundary/test_cutedsl_ar_fusion.py +++ b/test/registered/unit/layer_boundary/test_cutedsl_ar_fusion.py @@ -458,11 +458,7 @@ def test_early_shared_load_touches_only_the_finalize_routes(): def test_dual_stream_op_pins_the_deferral_off_under_a_deferring_caller(): - """The op's Tensor schema cannot carry a handoff; the dispatcher would raise - "Unable to cast ... to Tensor".""" - from sglang.srt.models.deepseek_v2 import ( # noqa: F401 (registers the op) - dsv2_flashinfer_moe_dual_stream_graph, - ) + """The captured dual-stream path returns tensors without deferred handoffs.""" class _DeferRecordingMoE: def forward_normal_dual_stream(self, hidden_states): @@ -471,17 +467,12 @@ def forward_normal_dual_stream(self, hidden_states): reset_context() fusion = _DeferRecordingMoE() - op = torch.ops.sglang.dsv2_flashinfer_moe_dual_stream_graph.default - # The CUDA key runs the real schema while the stub keeps tensors on CPU. - cuda_key = torch._C.DispatchKeySet(torch._C.DispatchKey.CUDA) - with ( - get_forward().scoped(defer_moe_finalize=True), - patch( - "sglang.srt.models.deepseek_v2.get_tc_piecewise_forward_context", - return_value=SimpleNamespace(moe_fusions={0: fusion}), - ), - ): - out = op.redispatch(cuda_key, torch.zeros(4, 8), 0, True, False) + from sglang.srt.models.deepseek_v2 import DeepseekV2MoE + + with get_forward().scoped(defer_moe_finalize=True): + out = DeepseekV2MoE._forward_moe_dual_stream_graph( + fusion, torch.zeros(4, 8), True, False + ) assert get_forward().defer_moe_finalize is True assert fusion.seen_defer is False diff --git a/test/registered/unit/layers/attention/test_dsa_head_gate_guard.py b/test/registered/unit/layers/attention/test_dsa_head_gate_guard.py index 246b24c4a6a6..7da0c6b8845c 100644 --- a/test/registered/unit/layers/attention/test_dsa_head_gate_guard.py +++ b/test/registered/unit/layers/attention/test_dsa_head_gate_guard.py @@ -19,23 +19,23 @@ class TestDsaHeadGateGuard(CustomTestCase): """ def test_definition_and_import_agree(self): - from sglang.srt.layers.attention.dsa import dsa_indexer, dsa_prefill_cuda_graph + from sglang.srt.layers.attention.dsa import dsa_indexer, head_gate for name in HELPERS: with self.subTest(helper=name): self.assertEqual( - hasattr(dsa_prefill_cuda_graph, name), - hasattr(dsa_indexer, name), + name in vars(head_gate), + name in vars(dsa_indexer), f"{name} is defined on one side of the guard but not the other", ) def test_helpers_exist_on_cuda_and_hip(self): - from sglang.srt.layers.attention.dsa import dsa_prefill_cuda_graph + from sglang.srt.layers.attention.dsa import head_gate expected = is_cuda() or is_hip() for name in HELPERS: with self.subTest(helper=name): - self.assertEqual(hasattr(dsa_prefill_cuda_graph, name), expected) + self.assertEqual(name in vars(head_gate), expected) if __name__ == "__main__": diff --git a/test/registered/unit/layers/quantization/test_fp8_utils.py b/test/registered/unit/layers/quantization/test_fp8_utils.py index d0ed39ca3a1c..9d6a012f2f56 100644 --- a/test/registered/unit/layers/quantization/test_fp8_utils.py +++ b/test/registered/unit/layers/quantization/test_fp8_utils.py @@ -195,7 +195,7 @@ def test_native_scalar_a_static_prequant_and_dynamic_scale_shapes(self): exec_config = SimpleNamespace( graph=SimpleNamespace( cuda_graph_config=SimpleNamespace( - prefill=SimpleNamespace(tc_compiler="none") + prefill=SimpleNamespace(backend="disabled") ) ) ) diff --git a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py index d9399f2195c8..a310cb40ea71 100644 --- a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py +++ b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py @@ -325,10 +325,6 @@ def _is_eligible(self, **overrides): patch(f"{_INDEXER}.is_cuda", return_value=True), patch(f"{_INDEXER}.is_hip", return_value=False), get_parallel().override(attn_cp_size=1), - patch( - f"{_INDEXER}.is_in_tc_piecewise_cuda_graph", - return_value=overrides.get("piecewise_graph", False), - ), patch(f"{_INDEXER}.is_in_breakable_cuda_graph", return_value=False), patch("torch.cuda.is_current_stream_capturing", return_value=False), ): @@ -353,7 +349,6 @@ def test_eligibility_is_fail_closed(self): {"batch_size": 20_000}, {"tbo": (1, 2)}, {"prefill_graph": True}, - {"piecewise_graph": True}, {"fp4": True}, ): with self.subTest(case=case): @@ -766,7 +761,6 @@ def test_capture_warmup_skips_dynamic_budget_but_eager_forward_uses_it(self): patch(f"{_METADATA}.get_is_capture_mode", return_value=capture_mode), patch("torch.cuda.is_current_stream_capturing", return_value=False), patch(f"{_METADATA}.is_in_breakable_cuda_graph", return_value=False), - patch(f"{_METADATA}.is_in_tc_piecewise_cuda_graph", return_value=False), patch( f"{_METADATA}.mqa_logits_budget_bytes", return_value=4096 ) as budget, diff --git a/test/registered/unit/layers/test_radix_attention.py b/test/registered/unit/layers/test_radix_attention.py index e2d7dc5a7459..c9eeef4683e2 100644 --- a/test/registered/unit/layers/test_radix_attention.py +++ b/test/registered/unit/layers/test_radix_attention.py @@ -1,583 +1,222 @@ -"""CPU unit tests for the graph-safe ``RadixAttention`` interface.""" +"""CPU contracts shared by the full and breakable prefill attention paths.""" -import unittest -from contextlib import ExitStack from types import SimpleNamespace from unittest.mock import patch +import pytest import torch -import sglang.srt.layers.radix_attention as radix_attention_module from sglang.srt.layers.radix_attention import RadixAttention -from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_context import ForwardContext, forward_context -from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, - set_tc_piecewise_forward_context, -) from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase -register_cpu_ci(est_time=11, suite="base-a-test-cpu") - - -class _RecordingAttentionBackend: - def __init__(self, *, return_lse=True): - self.calls = [] - self.return_lse = return_lse - - def forward( - self, - query, - key, - value, - attention_layer, - forward_batch, - save_kv_cache, - **kwargs, - ): - self.calls.append( - SimpleNamespace( - query=query, - key=key, - value=value, - attention_layer=attention_layer, - forward_batch=forward_batch, - output=forward_batch._attn_output, - out_cache_loc=forward_batch.out_cache_loc.clone(), - save_kv_cache=save_kv_cache, - kwargs=kwargs, +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +def make_batch(tokens=2, rows=4, return_lse=False): + return SimpleNamespace( + forward_mode=SimpleNamespace(is_extend=lambda: True), + global_num_token_non_padded_cpu=tokens, + out_cache_loc=torch.arange(rows), + positions=torch.arange(rows), + _attn_output=None, + mha_return_lse=return_lse, + ) + + +@pytest.mark.parametrize("breakable", [False, True]) +@pytest.mark.parametrize("return_lse", [False, True]) +@pytest.mark.parametrize("tokens", [0, 2, 4]) +def test_dense_outputs_padding_and_lse(breakable, return_lse, tokens): + layer = RadixAttention(2, 3, 1.0, 2, 7) + batch = make_batch(tokens=tokens, return_lse=return_lse) + original = batch.out_cache_loc, batch.positions + calls = [] + + def attention(q, k, v, actual_layer, live_batch, save_kv_cache, **kwargs): + assert actual_layer is layer + assert live_batch is batch + assert q.shape[0] == tokens + assert k.shape[0] == v.shape[0] == 3 + assert batch.positions.shape[0] == batch.out_cache_loc.shape[0] == tokens + calls.append(q) + out = torch.full_like(q, 3) + if return_lse: + return out, torch.full((tokens, 2), 7.0) + return out + + q = torch.zeros(4, 2, 3) + with ( + forward_context( + ForwardContext( + SimpleNamespace(forward=attention), + full_graph=not breakable, + raw_num_tokens=tokens, ) - ) - output = torch.full_like(query, 3) - lse = torch.full((query.shape[0], query.shape[1]), 7, dtype=torch.float32) - return (output, lse) if self.return_lse else output - - -class TestRadixAttentionGraphInterface(CustomTestCase): - @staticmethod - def _new_layer() -> RadixAttention: - layer = RadixAttention( - num_heads=2, - head_dim=3, - scaling=1.0, - num_kv_heads=2, - layer_id=0, - ) - return layer - - @staticmethod - def _new_impl_context( - attention_layers, - *, - mha_companion_layers=None, - num_tokens=4, - real_num_tokens=2, + ), + patch( + "sglang.srt.layers.radix_attention.is_in_breakable_cuda_graph", + return_value=breakable, + ), ): - forward_batch = SimpleNamespace( - global_num_token_non_padded_cpu=real_num_tokens, - out_cache_loc=torch.arange(num_tokens, dtype=torch.int64), - positions=torch.arange(num_tokens, dtype=torch.int64), - _attn_output=None, - mha_return_lse=False, - ) - return SimpleNamespace( - forward_batch=forward_batch, - attention_layers=attention_layers, - mha_companion_layers=mha_companion_layers, - num_tokens=None, - raw_num_tokens=None, + result = layer(q, q, q, batch, key_value_num_tokens=3) + output, lse = result if return_lse else (result, None) + torch.testing.assert_close(output[:tokens], torch.full_like(output[:tokens], 3)) + assert torch.count_nonzero(output[tokens:]) == 0 + assert len(calls) == bool(tokens) + assert batch.out_cache_loc is original[0] and batch.positions is original[1] + assert batch._attn_output is None + if return_lse: + assert lse.shape == (4, 2) + assert torch.all(lse[:tokens] == 7) + assert torch.count_nonzero(lse[tokens:]) == 0 + + +def test_extra_kwargs_and_exception_restore(): + layer = RadixAttention(2, 3, 1.0, 2, 0) + batch = make_batch() + sentinel = torch.empty(1) + batch._attn_output = sentinel + original = batch.out_cache_loc, batch.positions + modifier = lambda score: score + q = torch.zeros(4, 2, 3) + aux = [torch.arange(4)] + + def attention(q, k, v, actual_layer, live_batch, save_kv_cache, **kwargs): + assert kwargs["score_mod"] is modifier + assert kwargs["aux_tensors"][0].shape == (2,) + assert kwargs["q_rope"].shape[0] == 2 + assert kwargs["k_rope"].shape[0] == 3 + assert kwargs["q_descale"].shape[0] == 2 + raise RuntimeError("backend failed") + + with forward_context( + ForwardContext( + SimpleNamespace(forward=attention), full_graph=True, raw_num_tokens=2 ) - - def test_forward_dispatches_all_graph_and_lse_variants(self): - layer = self._new_layer() - query = torch.zeros((4, 2, 3)) - key = torch.zeros_like(query) - value = torch.zeros_like(query) - - op_names = { - (False, False): "unified_attention_with_output", - (False, True): "unified_attention_with_output_and_lse", - (True, False): "breakable_unified_attention_with_output", - (True, True): "breakable_unified_attention_with_output_and_lse", - } - - for breakable in (False, True): - for return_lse in (False, True): - with self.subTest(breakable=breakable, return_lse=return_lse): - forward_batch = SimpleNamespace( - forward_mode=ForwardMode.EXTEND, - mha_return_lse=return_lse, - ) - calls = [] - - def output_only(*args, **kwargs): - args[3].fill_(5) - calls.append(kwargs) - - def output_and_lse(*args, **kwargs): - args[3].fill_(5) - calls.append(kwargs) - return torch.full((4, 2), 11, dtype=torch.float32) - - with ExitStack() as stack: - stack.enter_context( - patch.object( - radix_attention_module, - "get_tc_piecewise_forward_context", - return_value=SimpleNamespace( - mha_companion_layers=[layer] - ), - ) - ) - stack.enter_context( - patch.object( - radix_attention_module, - "is_in_breakable_cuda_graph", - return_value=breakable, - ) - ) - mocks = { - name: stack.enter_context( - patch.object( - radix_attention_module, - name, - side_effect=( - output_and_lse - if name.endswith("and_lse") - else output_only - ), - ) - ) - for name in op_names.values() - } - result = layer( - query, - key, - value, - forward_batch, - key_value_num_tokens=3, - ) - - selected_name = op_names[(breakable, return_lse)] - for name, mock in mocks.items(): - self.assertEqual(mock.call_count, int(name == selected_name)) - - self.assertEqual( - calls, - [ - { - "use_mha_companion": True, - "key_value_num_tokens": 3, - } - ], - ) - if return_lse: - output, lse = result - self.assertEqual(lse.shape, (4, 2)) - self.assertTrue(torch.all(lse == 11)) - else: - output = result - self.assertEqual(output.shape, query.shape) - self.assertTrue(torch.all(output == 5)) - - def test_prefill_wrapper_opt_out_preserves_expanded_rows_and_batch(self): - """Expanded attention rows must not be sliced using the runner's token count.""" - layer = RadixAttention( - num_heads=2, - head_dim=3, - scaling=1.0, - num_kv_heads=2, - layer_id=0, - use_prefill_attention_wrapper=False, - ) - runner_batch = self._new_impl_context([layer], num_tokens=2).forward_batch - expanded_batch = SimpleNamespace( - forward_mode=ForwardMode.EXTEND, - out_cache_loc=runner_batch.out_cache_loc, - _attn_output=None, - mha_return_lse=False, - ) - query = torch.arange(48, dtype=torch.float32).view(8, 2, 3) - key = query.flatten(1).clone() - value = -key - backend = _RecordingAttentionBackend(return_lse=False) - - with ( - forward_context(ForwardContext(backend)), - set_tc_piecewise_forward_context( - runner_batch, [layer], None, [], [], num_tokens=8, full_graph=True - ), - ): - output = layer( - query, key, value, expanded_batch, save_kv_cache=False, causal=False - ) - self.assertIs( - get_tc_piecewise_forward_context().forward_batch, runner_batch - ) - - self.assertTrue(torch.equal(output, torch.full_like(query, 3))) - self.assertEqual(len(backend.calls), 1) - call = backend.calls[0] - self.assertIs(call.query, query) - for actual, original in ((call.key, key), (call.value, value)): - self.assertEqual(actual.shape, (8, 2, 3)) - self.assertEqual(actual.data_ptr(), original.data_ptr()) - self.assertTrue(torch.equal(actual.flatten(), original.flatten())) - self.assertIs(call.attention_layer, layer) - self.assertIs(call.forward_batch, expanded_batch) - self.assertIs(expanded_batch.out_cache_loc, runner_batch.out_cache_loc) - self.assertFalse(call.save_kv_cache) - self.assertEqual(call.kwargs, {"causal": False}) - self.assertEqual(runner_batch.global_num_token_non_padded_cpu, 2) - - def test_deferred_norm_rope_operands_follow_real_tokens_on_each_call(self): - layer = self._new_layer() - query = torch.zeros((4, 2, 3)) - positions = torch.arange(4) - temp_scale = torch.arange(4, dtype=torch.float32).reshape(4, 1) - norm_weight = torch.ones(3) - operands = { - "mxfp8_norm_rope_positions": positions, - "mxfp8_norm_rope_temp_scale": temp_scale, - "norm_weight": norm_weight, - } - for breakable in (False, True): - with self.subTest(breakable=breakable): - context = self._new_impl_context([layer]) - forward_batch = context.forward_batch - forward_batch.forward_mode = ForwardMode.EXTEND - original_cache_loc = forward_batch.out_cache_loc - backend = _RecordingAttentionBackend(return_lse=False) - with ( - patch.object( - radix_attention_module, - "get_tc_piecewise_forward_context", - return_value=context, - ), - patch.object( - radix_attention_module, - "get_attn_backend", - return_value=backend, - ), - patch.object( - radix_attention_module, - "is_in_breakable_cuda_graph", - return_value=breakable, - ), - patch.object( - radix_attention_module, - "breakable_attention_with_output_extra_kwargs", - side_effect=radix_attention_module.attention_with_output_extra_kwargs, - ) as graph_break, - ): - for real_tokens in (2, 4): - forward_batch.global_num_token_non_padded_cpu = real_tokens - positions.add_(10) - result = layer(query, query, query, forward_batch, **operands) - call = backend.calls[-1] - self.assertEqual(call.query.shape[0], real_tokens) - self.assertTrue( - torch.equal( - call.kwargs["mxfp8_norm_rope_positions"], - positions[:real_tokens], - ) - ) - self.assertTrue( - torch.equal( - call.kwargs["mxfp8_norm_rope_temp_scale"], - temp_scale[:real_tokens], - ) - ) - self.assertIs(call.kwargs["norm_weight"], norm_weight) - self.assertIs(forward_batch.out_cache_loc, original_cache_loc) - self.assertTrue(torch.all(result[:real_tokens] == 3)) - self.assertEqual(positions.shape[0], 4) - self.assertEqual(temp_scale.shape[0], 4) - self.assertEqual(graph_break.call_count, 2 if breakable else 0) - - def test_impl_preserves_attention_identity_and_lse(self): - mqa = SimpleNamespace() - mha = SimpleNamespace() - context = self._new_impl_context([mqa], mha_companion_layers=[mha]) - forward_batch = context.forward_batch - original_out_cache_loc = forward_batch.out_cache_loc - backend = _RecordingAttentionBackend() - query = torch.zeros((4, 2, 3)) - - with ( - patch.object( - radix_attention_module, - "get_tc_piecewise_forward_context", - return_value=context, - ), - patch.object( - radix_attention_module, "get_attn_backend", return_value=backend - ), - ): - for use_mha_companion, expected_layer in ((False, mqa), (True, mha)): - with self.subTest(use_mha_companion=use_mha_companion): - output = torch.empty_like(query) - lse = radix_attention_module._unified_attention_with_output_impl( - query, - query, - query, - output, - False, - 0, - use_mha_companion, - True, - ) - - call_record = backend.calls[-1] - self.assertIs(call_record.attention_layer, expected_layer) - self.assertEqual(call_record.query.shape, (2, 2, 3)) - self.assertEqual(call_record.key.shape, (2, 2, 3)) - self.assertEqual(call_record.value.shape, (2, 2, 3)) - self.assertEqual(call_record.output.shape, (2, 2, 3)) - self.assertEqual(call_record.out_cache_loc.tolist(), [0, 1]) - self.assertFalse(call_record.save_kv_cache) - self.assertTrue(torch.all(output[:2] == 3)) - self.assertEqual(lse.shape, (4, 2)) - self.assertTrue(torch.all(lse[:2] == 7)) - self.assertTrue(torch.all(lse[2:] == 0)) - self.assertIs(forward_batch.out_cache_loc, original_out_cache_loc) - - def test_extra_kwargs_path_returns_bucket_shaped_lse(self): - attention_layer = SimpleNamespace() - context = self._new_impl_context([attention_layer]) - backend = _RecordingAttentionBackend() - query = torch.zeros((4, 2, 3)) - - with ( - patch.object( - radix_attention_module, - "get_tc_piecewise_forward_context", - return_value=context, - ), - patch.object( - radix_attention_module, "get_attn_backend", return_value=backend - ), - ): - lse = radix_attention_module.attention_with_output_extra_kwargs( - query, - query, - query, - torch.empty_like(query), - False, - 0, - {"return_lse": True}, - ) - - self.assertEqual(lse.tolist(), [[7, 7], [7, 7], [0, 0], [0, 0]]) - - def test_impl_uses_independent_query_and_key_value_extents(self): - attention_layer = SimpleNamespace() - context = self._new_impl_context([attention_layer]) - forward_batch = context.forward_batch - original_out_cache_loc = forward_batch.out_cache_loc - backend = _RecordingAttentionBackend() - query = torch.zeros((4, 2, 3)) - key = torch.zeros((6, 2, 3)) - value = torch.zeros((6, 2, 3)) - k_rope = torch.zeros((6, 2, 1)) - output = torch.empty_like(query) - - with ( - patch.object( - radix_attention_module, - "get_tc_piecewise_forward_context", - return_value=context, - ), - patch.object( - radix_attention_module, "get_attn_backend", return_value=backend - ), - ): - lse = radix_attention_module._unified_attention_with_output_impl( - query, - key, - value, - output, - False, - 0, - False, - True, - key_value_num_tokens=5, - k_rope=k_rope, - ) - - call_record = backend.calls[-1] - self.assertEqual(call_record.query.shape, (2, 2, 3)) - self.assertEqual(call_record.key.shape, (5, 2, 3)) - self.assertEqual(call_record.value.shape, (5, 2, 3)) - self.assertEqual(call_record.kwargs["k_rope"].shape, (5, 2, 1)) - self.assertEqual(call_record.output.shape, (2, 2, 3)) - self.assertEqual(lse.shape, (4, 2)) - self.assertIs(forward_batch.out_cache_loc, original_out_cache_loc) - - def test_impl_preserves_output_only_contract(self): - attention_layer = SimpleNamespace() - context = self._new_impl_context([attention_layer]) - forward_batch = context.forward_batch - original_out_cache_loc = forward_batch.out_cache_loc - backend = _RecordingAttentionBackend(return_lse=False) - query = torch.zeros((4, 2, 3)) - output = torch.empty_like(query) - - with ( - patch.object( - radix_attention_module, - "get_tc_piecewise_forward_context", - return_value=context, - ), - patch.object( - radix_attention_module, "get_attn_backend", return_value=backend - ), - ): - lse = radix_attention_module._unified_attention_with_output_impl( - query, - query, - query, - output, - False, - 0, - False, - False, - ) - - self.assertIsNone(lse) - self.assertIs(backend.calls[-1].attention_layer, attention_layer) - self.assertTrue(torch.all(output[:2] == 3)) - self.assertIs(forward_batch.out_cache_loc, original_out_cache_loc) - - def test_impl_zero_real_tokens_returns_zeroed_lse(self): - # Regression: an idle DP rank whose fabricated EXTEND batch is masked to - # 0 real tokens skips attention entirely. The skip must still honor the - # LSE return mode -- unified_attention_with_output_and_lse asserts a - # tensor comes back, so returning a bare None raised AssertionError as - # soon as any 0-real-token call needed LSE (chunked-prefix MHA merge). - attention_layer = SimpleNamespace() - context = self._new_impl_context([attention_layer], real_num_tokens=0) - backend = _RecordingAttentionBackend() - query = torch.zeros((4, 2, 3)) - output = torch.full_like(query, float("nan")) - - with ( - patch.object( - radix_attention_module, - "get_tc_piecewise_forward_context", - return_value=context, - ), - patch.object( - radix_attention_module, "get_attn_backend", return_value=backend - ), - ): - lse = radix_attention_module._unified_attention_with_output_impl( - query, - query, - query, - output, - False, - 0, - False, - True, - ) - - self.assertEqual(backend.calls, []) - self.assertTrue(torch.all(output == 0)) - # Same shape/dtype the registered fake impl declares, so - # unified_attention_with_output_and_lse's `assert lse is not None` holds. - self.assertEqual(lse.shape, (4, 2)) - self.assertEqual(lse.dtype, torch.float32) - self.assertTrue(torch.all(lse == 0)) - - def test_impl_zero_real_tokens_output_only_returns_none(self): - # The 0-token skip must not start returning a tensor on the non-LSE - # path: unified_attention_with_output is registered with an inplace - # (None-returning) schema. - attention_layer = SimpleNamespace() - context = self._new_impl_context([attention_layer], real_num_tokens=0) - backend = _RecordingAttentionBackend(return_lse=False) - query = torch.zeros((4, 2, 3)) - output = torch.full_like(query, float("nan")) - - with ( - patch.object( - radix_attention_module, - "get_tc_piecewise_forward_context", - return_value=context, - ), - patch.object( - radix_attention_module, "get_attn_backend", return_value=backend - ), - ): - lse = radix_attention_module._unified_attention_with_output_impl( - query, - query, - query, - output, - False, - 0, - False, - False, - ) - - self.assertIsNone(lse) - self.assertEqual(backend.calls, []) - self.assertTrue(torch.all(output == 0)) - - def test_extra_kwargs_zero_real_tokens_zeroes_output(self): - # Regression: attention_with_output_extra_kwargs (Inkling score_mod / - # aux_tensors) narrowed to query[:0] and copied output[:0], so with 0 - # real tokens the preallocated torch.empty output was never written and - # its garbage (NaN/Inf) flowed into residuals and MoE routing. Only ROCm - # zeroed the padded tail, so every other platform leaked it. - attention_layer = SimpleNamespace() - context = self._new_impl_context([attention_layer], real_num_tokens=0) - backend = _RecordingAttentionBackend(return_lse=False) - query = torch.zeros((4, 2, 3)) - output = torch.full_like(query, float("nan")) - - with ( - patch.object( - radix_attention_module, - "get_tc_piecewise_forward_context", - return_value=context, - ), - patch.object( - radix_attention_module, "get_attn_backend", return_value=backend - ), - ): - radix_attention_module.attention_with_output_extra_kwargs( - query, - query, - query, - output, - False, - 0, - {}, + ): + with pytest.raises(RuntimeError, match="backend failed"): + layer( + q, + q, + q, + batch, + key_value_num_tokens=3, + score_mod=modifier, + aux_tensors=aux, + q_rope=q, + k_rope=q, + q_descale=q, ) - - self.assertEqual(backend.calls, []) - self.assertTrue(torch.all(output == 0)) - - def test_lse_fake_impl_declares_shape_and_dtype(self): - query = torch.empty((5, 3, 7), dtype=torch.float16) - output = torch.empty_like(query) - - lse = radix_attention_module._unified_attention_with_output_and_lse_fake( - query, - None, - None, - output, - False, - 0, + assert batch.out_cache_loc is original[0] and batch.positions is original[1] + assert batch._attn_output is sentinel + assert aux[0].shape == (4,) + + +@pytest.mark.parametrize("tokens", [0, 2]) +@pytest.mark.parametrize("has_index_value", [False, True]) +def test_sparse_full_graph_two_outputs(tokens, has_index_value): + layer = RadixAttention(2, 3, 1.0, 2, 0) + batch = make_batch(tokens=tokens) + q = torch.zeros(4, 2, 3) + idx = torch.zeros(4, 1, 2) + + def attention(q, k, v, layer, batch, save_kv_cache, **kwargs): + assert kwargs["idx_q"].shape == (tokens, 1, 2) + assert kwargs["idx_k"].shape == (tokens, 1, 2) + return ( + torch.full((tokens, 1, 2), 5.0) if has_index_value else None, + torch.ones_like(q), ) - self.assertEqual(lse.shape, (5, 3)) - self.assertEqual(lse.dtype, torch.float32) - self.assertEqual(lse.device, query.device) - - -if __name__ == "__main__": - unittest.main() + with forward_context( + ForwardContext( + SimpleNamespace(forward=attention), full_graph=True, raw_num_tokens=tokens + ) + ): + index_output, output = layer(q, q, q, batch, idx_q=idx, idx_k=idx) + assert output.shape == (4, 6) and index_output.shape == (4, 2) + assert torch.all(output[:tokens] == 1) + assert torch.count_nonzero(output[tokens:]) == 0 + if has_index_value: + assert torch.all(index_output[:tokens] == 5) + assert torch.count_nonzero(index_output[tokens:]) == 0 + + +@pytest.mark.parametrize("sparse", [False, True]) +def test_direct_backend_dispatch(sparse): + layer = RadixAttention(2, 3, 1.0, 2, 0) + batch = make_batch() + q = torch.zeros(4, 2, 3) + result = object() + backend = SimpleNamespace(forward=lambda *args, **kwargs: result) + with ( + forward_context(ForwardContext(backend)), + patch( + "sglang.srt.layers.radix_attention.is_in_breakable_cuda_graph", + return_value=sparse, + ), + ): + kwargs = {"idx_q": q, "idx_k": q} if sparse else {} + assert layer(q, q, q, batch, **kwargs) is result + + +def test_output_dtype_follows_values(): + layer = RadixAttention(2, 3, 1.0, 2, 0) + batch = make_batch() + q = torch.zeros(4, 2, 3, dtype=torch.float32) + v = q.to(torch.bfloat16) + backend = SimpleNamespace( + forward=lambda q, k, v, *args, **kwargs: torch.ones_like(v) + ) + with forward_context(ForwardContext(backend, full_graph=True, raw_num_tokens=2)): + assert layer(q, q, v, batch).dtype == torch.bfloat16 + + +@pytest.mark.parametrize("return_lse", [False, True]) +def test_attention_replay_uses_live_batch_and_static_output_buffers(return_lse): + from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( + breakable_cuda_graph as bcg, + ) + + layer = RadixAttention(2, 3, 1.0, 2, 0) + q = torch.zeros(4, 2, 3) + captured_batch = make_batch(return_lse=return_lse) + live_batch = make_batch(tokens=1, return_lse=return_lse) + seen = [] + + def attention(query, key, value, actual_layer, batch, save_kv_cache, **kwargs): + seen.append(batch) + out = torch.full_like(query, float(len(seen))) + lse = torch.full((query.shape[0], 2), float(len(seen))) + return (out, lse) if return_lse else out + + graph = bcg.BreakableCUDAGraph() + capture = SimpleNamespace( + cuda_graph=graph, + _barrier_fn=None, + _end_current_segment=lambda: None, + _begin_new_segment=lambda: None, + ) + backend = SimpleNamespace(forward=attention) + with ( + forward_context(ForwardContext(backend, raw_num_tokens=2)), + patch( + "sglang.srt.layers.radix_attention.is_in_breakable_cuda_graph", + return_value=True, + ), + ): + token = bcg._current_capture_var.set(capture) + try: + result = layer(q, q, q, captured_batch) + finally: + bcg._current_capture_var.reset(token) + output, lse = result if return_lse else (result, None) + pointer = output.data_ptr() + with forward_context(ForwardContext(backend, raw_num_tokens=1)): + graph._break_fns[0](live_batch) + assert seen == [captured_batch, live_batch] + assert output.data_ptr() == pointer + assert torch.all(output[:1] == 2) and torch.count_nonzero(output[1:]) == 0 + if return_lse: + assert torch.all(lse[:1] == 2) and torch.count_nonzero(lse[1:]) == 0 diff --git a/test/registered/unit/layers/test_radix_linear_attention.py b/test/registered/unit/layers/test_radix_linear_attention.py index a11230a8e7a1..a921e8dbdb23 100644 --- a/test/registered/unit/layers/test_radix_linear_attention.py +++ b/test/registered/unit/layers/test_radix_linear_attention.py @@ -66,6 +66,57 @@ def forward(self, *, layer, forward_batch, mixed_qkv, a, b): class TestRadixLinearAttentionPadding(CustomTestCase): + def test_graph_dispatch_distinguishes_prefill_decode_and_verify(self): + layer = radix_linear_attention.RadixLinearAttention( + layer_id=0, + num_q_heads=1, + num_k_heads=1, + num_v_heads=2, + head_q_dim=4, + head_k_dim=4, + head_v_dim=4, + ) + for mode, extend, verify in ( + ("prefill", True, False), + ("decode", False, False), + ("verify", True, True), + ): + for breakable, full in ((False, False), (True, False), (False, True)): + batch = SimpleNamespace( + forward_mode=SimpleNamespace( + is_extend=lambda: extend, + is_target_verify=lambda: verify, + is_extend_without_speculative=lambda: extend and not verify, + ), + global_num_token_non_padded_cpu=3, + out_cache_loc=torch.arange(3), + ) + with ( + self.subTest(mode=mode, breakable=breakable, full=full), + patch.object( + radix_linear_attention, + "is_in_breakable_cuda_graph", + return_value=breakable, + ), + patch.object( + radix_linear_attention, + "is_in_full_prefill_graph", + return_value=full, + ), + patch.object( + radix_linear_attention, + "get_attn_backend", + return_value=_FakeAttentionBackend(), + ), + patch.object(layer, "_eager_linear_attention") as eager, + ): + layer.forward( + batch, torch.zeros(3, 8), torch.zeros(3, 2), torch.zeros(3, 2) + ) + self.assertEqual( + eager.called, extend and (full or (breakable and not verify)) + ) + def test_eager_padded_input_is_sliced_and_output_shape_is_restored(self): layer = radix_linear_attention.RadixLinearAttention( layer_id=0, @@ -86,8 +137,8 @@ def test_eager_padded_input_is_sliced_and_output_shape_is_restored(self): with ( patch.object( radix_linear_attention, - "get_tc_piecewise_forward_context", - return_value=None, + "is_in_full_prefill_graph", + return_value=False, ), patch.object( radix_linear_attention, @@ -126,8 +177,8 @@ def test_target_verify_keeps_physical_rows_matching_its_metadata(self): with ( patch.object( radix_linear_attention, - "get_tc_piecewise_forward_context", - return_value=None, + "is_in_full_prefill_graph", + return_value=False, ), patch.object( radix_linear_attention, @@ -165,8 +216,8 @@ def test_eager_backend_failure_restores_out_cache_loc(self): with ( patch.object( radix_linear_attention, - "get_tc_piecewise_forward_context", - return_value=None, + "is_in_full_prefill_graph", + return_value=False, ), patch.object( radix_linear_attention, @@ -201,8 +252,8 @@ def test_padded_output_tail_is_initialized(self): with ( patch.object( radix_linear_attention, - "get_tc_piecewise_forward_context", - return_value=context, + "is_in_full_prefill_graph", + return_value=True, ), patch.object( radix_linear_attention, @@ -210,12 +261,13 @@ def test_padded_output_tail_is_initialized(self): return_value=_FakeAttentionBackend(), ), ): - radix_linear_attention._unified_linear_attention_with_output_impl( + radix_linear_attention._linear_attention_with_output_impl( mixed_qkv=torch.zeros((padded_num_tokens, 8)), a=torch.zeros((padded_num_tokens, 2)), b=torch.zeros((padded_num_tokens, 2)), output=output, - layer_id=0, + attention_layer=context.attention_layers[0], + forward_batch=forward_batch, ) torch.testing.assert_close(output[:, :3], torch.full((1, 3, 2, 4), 5.0)) diff --git a/test/registered/unit/model_executor/model_runner_components/test_cuda_graph_setup.py b/test/registered/unit/model_executor/model_runner_components/test_cuda_graph_setup.py index ac9aad7da824..a77bc8f4fd82 100644 --- a/test/registered/unit/model_executor/model_runner_components/test_cuda_graph_setup.py +++ b/test/registered/unit/model_executor/model_runner_components/test_cuda_graph_setup.py @@ -5,10 +5,8 @@ from sglang.srt.model_executor.model_runner_components import cuda_graph_setup from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import ( - _align_pipeline_layers, capture_decode_graph, has_standard_gqa_for_all_local_layers, - index_attention_layers_by_global_id, ) from sglang.test.ci.ci_register import register_cpu_ci @@ -31,42 +29,6 @@ def test_standard_gqa_gate_is_unchanged_without_pipeline_parallelism(): ) -def test_pipeline_attention_metadata_is_indexed_by_global_layer_id(): - layer23 = SimpleNamespace(layer_id=23) - layer24 = SimpleNamespace(layer_id=24) - companion24 = object() - - attention, companions = index_attention_layers_by_global_id( - [layer23, layer24], [None, companion24] - ) - - assert len(attention) == 25 - assert all(layer is None for layer in attention[:23]) - assert attention[23] is layer23 - assert attention[24] is layer24 - assert companions[23] is None - assert companions[24] is companion24 - - -def test_reuse_tables_pass_through_but_distinct_duplicates_raise(): - looped = SimpleNamespace(layer_id=1) - companion = object() - attention_in = [SimpleNamespace(layer_id=0), looped, looped] - companions_in = [None, companion, companion] - - attention, companions = index_attention_layers_by_global_id( - attention_in, companions_in - ) - - assert attention is attention_in - assert companions is companions_in - - with pytest.raises(ValueError, match="duplicate attention layer_id: 2"): - index_attention_layers_by_global_id( - [SimpleNamespace(layer_id=2), SimpleNamespace(layer_id=2)], [None, None] - ) - - def test_model_runner_can_override_decode_graph_runner(monkeypatch): from sglang.srt.runtime_context import get_context @@ -111,30 +73,5 @@ def _decode_cuda_graph_runner_cls(self): override.restore() -def test_align_pipeline_layers_uses_absolute_indices(): - class PipelineStage: - start_layer = 3 - end_layer = 5 - layers = [object()] * 8 - - local_layers = ["layer-3", "layer-4"] - assert _align_pipeline_layers(local_layers, PipelineStage()) == [ - None, - None, - None, - "layer-3", - "layer-4", - None, - None, - None, - ] - full_model = SimpleNamespace(layers=local_layers) - assert _align_pipeline_layers(local_layers, full_model) == local_layers - with pytest.raises(AssertionError, match="together"): - _align_pipeline_layers( - local_layers, SimpleNamespace(start_layer=0, layers=local_layers) - ) - - if __name__ == "__main__": sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/model_executor/model_runner_components/test_layer_setup.py b/test/registered/unit/model_executor/model_runner_components/test_layer_setup.py index bf9a5b16ef67..32c86dce4908 100644 --- a/test/registered/unit/model_executor/model_runner_components/test_layer_setup.py +++ b/test/registered/unit/model_executor/model_runner_components/test_layer_setup.py @@ -5,9 +5,11 @@ from types import SimpleNamespace from unittest.mock import patch +from torch import nn + from sglang.srt.distributed.utils import get_pp_indices from sglang.srt.model_executor.model_runner_components.layer_setup import ( - compute_attention_and_moe_layers, + compute_attention_layer_info, resolve_layer_indices, ) from sglang.test.ci.ci_register import register_cpu_ci @@ -16,8 +18,8 @@ register_cpu_ci(est_time=6, suite="base-a-test-cpu") -class TestComputeAttentionAndMoeLayers(CustomTestCase): - def test_deepseek_mla_registers_mha_companion(self): +class TestComputeAttentionLayerInfo(CustomTestCase): + def test_deepseek_mla_detects_mha_companion(self): attn_mqa = SimpleNamespace() attn_mha = SimpleNamespace() layer_model = SimpleNamespace( @@ -28,27 +30,49 @@ def test_deepseek_mla_registers_mha_companion(self): ] ) - attention_layers, _, _, _, mha_companion_layers = ( - compute_attention_and_moe_layers(layer_model) + attention_layer_count, has_mha_companion_layers = compute_attention_layer_info( + layer_model ) - self.assertEqual(attention_layers, [attn_mqa]) - self.assertEqual(mha_companion_layers, [attn_mha]) + self.assertEqual(attention_layer_count, 1) + self.assertTrue(has_mha_companion_layers) self.assertNotIn("_pcg_mha_companion", vars(attn_mqa)) - def test_pipeline_placeholders_preserve_global_layer_ids(self): + def test_pipeline_placeholders_do_not_count_as_attention(self): local_attention = SimpleNamespace() layer_model = SimpleNamespace( layers=[SimpleNamespace(), SimpleNamespace()] + [SimpleNamespace(self_attn=SimpleNamespace(attn=local_attention))] ) - attention_layers, _, _, _, mha_companion_layers = ( - compute_attention_and_moe_layers(layer_model) + attention_layer_count, has_mha_companion_layers = compute_attention_layer_info( + layer_model + ) + + self.assertEqual(attention_layer_count, 1) + self.assertFalse(has_mha_companion_layers) + + def test_loop_attention_counts_executions(self): + attention = nn.Identity() + layer_model = SimpleNamespace( + layers=[ + SimpleNamespace( + self_attn=SimpleNamespace( + attn=nn.ModuleList([attention, attention]) + ) + ) + ] ) + self.assertEqual(compute_attention_layer_info(layer_model), (2, False)) - self.assertEqual(attention_layers, [None, None, local_attention]) - self.assertEqual(mha_companion_layers, [None, None, None]) + def test_module_dict_skips_non_attention_layers(self): + supported = nn.Module() + supported.self_attn = nn.Module() + supported.self_attn.attn = nn.Identity() + layer_model = SimpleNamespace( + layers=nn.ModuleDict({"local": supported, "placeholder": nn.Identity()}) + ) + self.assertEqual(compute_attention_layer_info(layer_model), (1, False)) NUM_LAYERS = 36 diff --git a/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py b/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py index fa383ff034b9..e0eb494912b3 100644 --- a/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py +++ b/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py @@ -231,15 +231,6 @@ def test_unsupported_overlap_is_rejected_instead_of_falling_back(self): dict(options=_make_options(cuda_graph_enabled=False)), "CUDA graph capture is disabled", ), - ( - "tc_piecewise_prefill", - dict( - options=_make_options( - prefill_cuda_graph_backend=Backend.TC_PIECEWISE - ) - ), - "tc_piecewise prefill CUDA graphs are not supported", - ), ( "pt_checkpoint", dict(load_config=LoadConfig(load_format=LoadFormat.PT)), diff --git a/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py b/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py index b6d4187bdbc2..b16380b8b7a7 100644 --- a/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py +++ b/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py @@ -30,7 +30,7 @@ def _make_runner(self): runner._is_full_backend = False runner.enable_lora = False runner._capture_chunked_prefix = False - runner.prefill_backend_name = Backend.TC_PIECEWISE + runner.prefill_backend_name = Backend.BREAKABLE runner.has_mha_companion_layers = False runner.prefer_eager_mixed_prefill = False runner.capture_hidden_mode = CaptureHiddenMode.NULL diff --git a/test/registered/unit/model_executor/test_cuda_graph_config_backends.py b/test/registered/unit/model_executor/test_cuda_graph_config_backends.py new file mode 100644 index 000000000000..6ce0432c359e --- /dev/null +++ b/test/registered/unit/model_executor/test_cuda_graph_config_backends.py @@ -0,0 +1,43 @@ +"""Validate backend retirement without importing accelerator backends.""" + +import argparse +from unittest.mock import patch + +import pytest + +from sglang.srt.model_executor.cuda_graph_config import ( + Backend, + CudaGraphConfig, + default_prefill_backend, + parse_cuda_graph_config_arg, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +@pytest.mark.parametrize("backend", [Backend.FULL, Backend.BREAKABLE, Backend.DISABLED]) +def test_supported_prefill_backends_round_trip(backend): + raw = parse_cuda_graph_config_arg('{"prefill":{"backend":"' + backend + '"}}') + config = CudaGraphConfig.from_dict(raw) + assert config.prefill.backend == backend + + +def test_removed_backend_is_rejected(): + with pytest.raises(ValueError, match="tc_piecewise was removed"): + CudaGraphConfig.from_dict({"prefill": {"backend": "tc_piecewise"}}) + + +def test_removed_compiler_option_is_rejected(): + with pytest.raises(argparse.ArgumentTypeError, match="tc_compiler"): + parse_cuda_graph_config_arg('{"prefill":{"tc_compiler":"eager"}}') + with pytest.raises(ValueError, match="tc_compiler was removed"): + CudaGraphConfig.from_dict({"prefill": {"tc_compiler": "inductor"}}) + + +@pytest.mark.parametrize( + "cuda, expected", [(True, Backend.BREAKABLE), (False, Backend.DISABLED)] +) +def test_prefill_defaults_do_not_enable_unvalidated_platforms(cuda, expected): + with patch("sglang.srt.utils.is_cuda", return_value=cuda): + assert default_prefill_backend() == expected diff --git a/test/registered/unit/model_executor/test_eager_on_graph.py b/test/registered/unit/model_executor/test_eager_on_graph.py new file mode 100644 index 000000000000..47c6980f4529 --- /dev/null +++ b/test/registered/unit/model_executor/test_eager_on_graph.py @@ -0,0 +1,118 @@ +"""CPU tests of eager replay binding without requiring CUDA graph capture.""" + +from contextlib import contextmanager +from types import SimpleNamespace + +import pytest +import torch + +from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( + breakable_cuda_graph as bcg, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") + + +@contextmanager +def recording_capture(): + events = [] + graph = bcg.BreakableCUDAGraph() + capture = SimpleNamespace( + cuda_graph=graph, + _barrier_fn=None, + _end_current_segment=lambda: events.append("end"), + _begin_new_segment=lambda: events.append("begin"), + ) + token = bcg._current_capture_var.set(capture) + try: + yield graph, events + finally: + bcg._current_capture_var.reset(token) + + +@pytest.mark.parametrize("keyword_batch", [False, True]) +def test_method_replay_rebinds_batch_and_preserves_output_address(keyword_batch): + class Layer: + @bcg.eager_on_graph + def _eager_scale(self, x, forward_batch): + return x * forward_batch.scale + + layer = Layer() + x = torch.tensor([2.0]) + captured_batch = SimpleNamespace(scale=3) + with recording_capture() as (graph, events): + if keyword_batch: + output = layer._eager_scale(x, forward_batch=captured_batch) + else: + output = layer._eager_scale(x, captured_batch) + assert events == ["end", "begin"] + pointer = output.data_ptr() + x.fill_(4) + graph._break_fns[0](SimpleNamespace(scale=5)) + assert output.item() == 20 and output.data_ptr() == pointer + with pytest.raises(ValueError, match="ForwardBatch"): + graph._break_fns[0](None) + + +def test_method_capture_stub_and_real_replay(): + calls = [] + + class Layer: + def _capture_stub(self, x, forward_batch): + calls.append("stub") + return torch.zeros_like(x) + + @bcg.eager_on_graph(capture_stub=_capture_stub) + def _eager_scale(self, x, forward_batch): + calls.append("real") + return x * forward_batch.scale + + layer = Layer() + with recording_capture() as (graph, _): + output = layer._eager_scale(torch.ones(2), SimpleNamespace(scale=3)) + assert calls == ["stub"] + graph._break_fns[0](SimpleNamespace(scale=7)) + assert calls == ["stub", "real"] + torch.testing.assert_close(output, torch.full((2,), 7.0)) + + +def test_plain_callable_and_eager_passthrough(): + @bcg.eager_on_graph + def increment(x): + return x + 1 + + assert increment(torch.tensor(2)).item() == 3 + x = torch.tensor(3) + with recording_capture() as (graph, _): + output = increment(x) + x.fill_(9) + graph._break_fns[0](None) + assert output.item() == 10 + + +def test_replay_does_not_retain_capture_or_serving_batches(): + import gc + import weakref + + class Batch: + scale = 3 + + @bcg.eager_on_graph + def scale(x, forward_batch): + return x * forward_batch.scale + + batch = Batch() + capture_ref = weakref.ref(batch) + with recording_capture() as (graph, _): + scale(torch.ones(2), batch) + del batch + gc.collect() + assert capture_ref() is None + + batch = Batch() + live_ref = weakref.ref(batch) + graph._break_fns[0](batch) + del batch + gc.collect() + assert live_ref() is None diff --git a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py index 29805fb3f2e6..77c7422eecbe 100644 --- a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py +++ b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py @@ -242,15 +242,10 @@ def test_low_free_memory_still_captures_prefill_graph(self): model_config=SimpleNamespace(context_len=8192, num_hidden_layers=1), layer_info=SimpleNamespace(start_layer=0, end_layer=1), req_to_token_pool=SimpleNamespace(size=1), - get_cuda_graph_layers=lambda _layer_model: ( - [object()], - [], - [], - [], - [None], - ), ) - language_model = SimpleNamespace(layers=[object()]) + language_model = SimpleNamespace( + layers=[SimpleNamespace(self_attn=SimpleNamespace(attn=object()))] + ) with ( patch.object(graph_setup, "check_cuda_graph_backend", return_value=False), @@ -275,30 +270,6 @@ def test_low_free_memory_still_captures_prefill_graph(self): self.assertIs(capture.runner, prefill_runner) - def test_eagle_target_tc_piecewise_skips_last_mode_capture(self): - eager_runner = object() - # The server-side hidden-state ceiling and graph config are bag leaves. - override = get_context().override_server_args( - enable_return_hidden_states=True, - return_hidden_states_mode="last", - cuda_graph_config=SimpleNamespace( - prefill=SimpleNamespace(backend=Backend.TC_PIECEWISE) - ), - ) - override.install() - self.addCleanup(override.restore) - model_runner = SimpleNamespace( - is_draft_worker=False, - spec_algorithm=SimpleNamespace(is_eagle=lambda: True), - ) - - capture = capture_prefill_graph( - model_runner=model_runner, - eager_runner=eager_runner, - ) - - self.assertIs(capture.runner, eager_runner) - def test_pp_proxy_output_is_trimmed_to_raw_prefill_tokens(self): runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner) runner.raw_num_tokens = 3 diff --git a/test/registered/unit/models/test_inkling_graph_fusion_gates.py b/test/registered/unit/models/test_inkling_graph_fusion_gates.py new file mode 100644 index 000000000000..e1d0f659209d --- /dev/null +++ b/test/registered/unit/models/test_inkling_graph_fusion_gates.py @@ -0,0 +1,66 @@ +"""BCG prefill must not change decode or verification fusion eligibility.""" + +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import torch + +from sglang.srt.model_executor.forward_context import ForwardContext, forward_context +from sglang.srt.models.inkling_common.kernels import comm +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +@pytest.mark.parametrize("scope", ["eager", "breakable", "full"]) +@pytest.mark.parametrize("mode", ["extend", "decode", "verify"]) +@pytest.mark.parametrize("scattered", [False, True]) +def test_fusion_gate_preserves_non_prefill_behavior(scope, mode, scattered): + batch = SimpleNamespace( + forward_mode=SimpleNamespace( + is_draft_extend_v2=lambda: False, + is_extend=lambda: mode != "decode", + is_decode=lambda: mode == "decode", + is_target_verify=lambda: mode == "verify", + is_extend_without_speculative=lambda: mode == "extend", + ) + ) + group = SimpleNamespace( + world_size=4, + torch_symm_mem_comm=SimpleNamespace(disabled=False, dtype=torch.bfloat16), + ) + gate = ( + comm.scattered_ar_sconv_fusable + if scattered + else comm.fullwidth_ar_sconv_fusable + ) + with ( + forward_context(ForwardContext(attn_backend=None, full_graph=scope == "full")), + patch.object(comm, "is_cuda", return_value=True), + patch.object( + comm, "is_in_breakable_cuda_graph", return_value=scope == "breakable" + ), + patch.object( + comm, + "get_exec", + return_value=SimpleNamespace( + comm=SimpleNamespace(enable_scattered_sconv=scattered) + ), + ), + patch.object( + comm.envs.SGLANG_OPT_USE_INKLING_CUSTOM_AR, "get", return_value=True + ), + patch.object( + comm.envs.SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV, "get", return_value=True + ), + patch.object( + comm, + "_get_inkling_ar_resources", + return_value=SimpleNamespace(ssconv_out=4096 * 128), + ), + ): + expected = (scattered or mode == "extend") and not ( + scope == "breakable" and mode == "extend" + ) + assert gate(group, batch, 4096, 128, torch.bfloat16) == expected diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 67010b13066f..f4487c95f914 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -25,7 +25,6 @@ ) from sglang.srt.arg_groups.cuda_graph_hook import ( apply_cuda_graph_compatibility, - disable_tc_piecewise_cudagraph_if_incompatible, finalize_cuda_graph_prefill_max_context, handle_cuda_graph_config, ) @@ -2482,44 +2481,6 @@ def test_rejects_fp4_kv_cache(self): self._validate_prefill_only_args(kv_cache_dtype=kv_cache_dtype) -class TestCudaGraphConfigDataclassAccess(CustomTestCase): - @patch( - "sglang.srt.model_executor.runner_backend." - "tc_piecewise_cuda_graph_backend.get_moe_a2a_backend" - ) - def test_tc_piecewise_build_config_reads_phase_config_dataclass( - self, mock_get_moe_a2a_backend - ): - from sglang.srt.model_executor.runner_backend.tc_piecewise_cuda_graph_backend import ( - TcPiecewiseCudaGraphBackend, - ) - - mock_backend = mock_get_moe_a2a_backend.return_value - mock_backend.is_deepep.return_value = False - mock_backend.is_mooncake.return_value = False - from sglang.srt.runtime_context import get_context - - # The graph configuration is a bag leaf; the debug switch is raw input - # and stays on the argument. - override = get_context().override_server_args( - cuda_graph_config=CudaGraphConfig( - prefill=PhaseConfig( - backend=Backend.TC_PIECEWISE, - bs=[32, 64], - tc_compiler="eager", - ) - ) - ) - override.install() - self.addCleanup(override.restore) - server_args = SimpleNamespace(enable_torch_compile_debug_mode=False) - - config = TcPiecewiseCudaGraphBackend.build_compilation_config(server_args) - - self.assertEqual(config.get_capture_sizes(), [32, 64]) - self.assertEqual(config.compiler, "eager") - - class TestPipelineParallelCompat(CustomTestCase): """Features supported with `pipeline-parallel-size > 1`.""" @@ -2802,33 +2763,6 @@ def test_lora_paths_keep_breakable_prefill_graph(self): Backend.BREAKABLE, ) - def test_lora_still_disables_tc_piecewise_prefill_graph(self): - # Pin the tc_piecewise LoRA rule itself, with the hardware rule - # neutralized so this runs on CPU-only CI. - args = ServerArgs(model_path="dummy", enable_lora=True) - args._model_config = SimpleNamespace( - hf_config=SimpleNamespace(architectures=["LlamaForCausalLM"]), - is_piecewise_cuda_graph_disabled_model=False, - is_multimodal=False, - is_multimodal_piecewise_cuda_graph_supported=False, - ) - args.cuda_graph_config = CudaGraphConfig( - prefill=PhaseConfig(backend=Backend.TC_PIECEWISE) - ) - with ( - override_platform(is_hip=False), - override_platform(is_npu=False), - patch("sglang.srt.arg_groups.cuda_graph_hook.is_cpu", return_value=False), - patch("sglang.srt.arg_groups.cuda_graph_hook.is_mps", return_value=False), - override_platform(is_xpu=False), - ): - disable_tc_piecewise_cudagraph_if_incompatible(args) - - self.assertEqual( - resolution_result(args, "cuda_graph_config").prefill.backend, - Backend.DISABLED, - ) - class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase): """The BCG "multimodal model" rule exempts archs on the BCG multimodal @@ -2894,38 +2828,30 @@ def test_allowlist_membership(self): class TestCutedslMoeMaxNumTokens(CustomTestCase): """The shared CuteDSL MoE per-forward token bound. Fields are set directly to exercise the math independently of __post_init__ resolution. - - cg-refactor: the legacy disable_piecewise_cuda_graph / - piecewise_cuda_graph_max_tokens / cuda_graph_max_bs fields were - consolidated into cuda_graph_config; the helper accepts the legacy - kwarg names for test readability and translates them to the per-phase - dataclasses. """ - def _args(self, **overrides): + def _args( + self, + *, + prefill_backend=Backend.BREAKABLE, + prefill_graph_max_tokens=2048, + decode_graph_max_bs=512, + **overrides, + ): server_args = ServerArgs(model_path="dummy") fields = dict( speculative_algorithm=None, speculative_num_draft_tokens=None, max_prefill_tokens=16384, - disable_piecewise_cuda_graph=False, - piecewise_cuda_graph_max_tokens=2048, - cuda_graph_max_bs=512, ) fields.update(overrides) - disable_piecewise = fields.pop("disable_piecewise_cuda_graph") - piecewise_max = fields.pop("piecewise_cuda_graph_max_tokens") - cg_max_bs = fields.pop("cuda_graph_max_bs") for key, value in fields.items(): setattr(server_args, key, value) server_args.cuda_graph_config = CudaGraphConfig( - decode=PhaseConfig(backend=Backend.FULL, max_bs=cg_max_bs), + decode=PhaseConfig(backend=Backend.FULL, max_bs=decode_graph_max_bs), prefill=PhaseConfig( - backend=( - Backend.DISABLED if disable_piecewise else Backend.TC_PIECEWISE - ), - max_bs=piecewise_max, - tc_compiler="eager", + backend=prefill_backend, + max_bs=prefill_graph_max_tokens, ), ) return server_args @@ -2934,20 +2860,20 @@ def test_prefill_dominates_in_default_config(self): self.assertEqual(cutedsl_moe_max_num_tokens(self._args()), 16384) def test_speculative_decoding_scales_decode_bound(self): - # decode bound 512 * 8 dominates the small prefill/piecewise bounds + # decode bound 512 * 8 dominates the small prefill bounds args = self._args( max_prefill_tokens=512, - piecewise_cuda_graph_max_tokens=512, + prefill_graph_max_tokens=512, speculative_algorithm="EAGLE", speculative_num_draft_tokens=8, ) self.assertEqual(cutedsl_moe_max_num_tokens(args), 4096) - def test_piecewise_bound_excluded_when_disabled(self): + def test_prefill_graph_bound_excluded_when_disabled(self): args = self._args( max_prefill_tokens=512, - disable_piecewise_cuda_graph=True, - cuda_graph_max_bs=64, + prefill_backend=Backend.DISABLED, + decode_graph_max_bs=64, ) self.assertEqual(cutedsl_moe_max_num_tokens(args), 512) diff --git a/test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py b/test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py index ac79697f078a..49f07720ceb3 100644 --- a/test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py +++ b/test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py @@ -16,7 +16,7 @@ Capture is wired: the captured batch reads `out_cache_loc` out of the registry slot, refilled before each replay from the already-rebound kernel-facing loc, and the read tables are refilled out-of-graph from the live v2p. So BREAKABLE -(the CUDA default) and TC_PIECEWISE must be left alone -- an earlier gate +(the CUDA default) must be left alone -- an earlier gate disabled every prefill backend outright, which cost every unified run its prefill graph. @@ -79,7 +79,7 @@ class TestUnifiedPrefillCudaGraphGate(unittest.TestCase): def test_non_full_prefill_backends_are_left_enabled(self): """BUG REGRESSION. Unified used to disable prefill capture outright, so the default BREAKABLE graph silently never ran.""" - for backend in (Backend.BREAKABLE, Backend.TC_PIECEWISE): + for backend in (Backend.BREAKABLE,): for attn in (("fa4", "fa4"), ("triton", "triton")): with self.subTest(prefill=backend, attn=attn): cg = _run_handler(prefill_backend=backend, attention_backends=attn) diff --git a/test/registered/unit/test_server_args_migration.py b/test/registered/unit/test_server_args_migration.py index 0371f11a0385..ae1cf8a8450e 100644 --- a/test/registered/unit/test_server_args_migration.py +++ b/test/registered/unit/test_server_args_migration.py @@ -168,7 +168,7 @@ def backends(argv): return sa.disable_cuda_graph, config.decode.backend, config.prefill.backend # Not a literal: the prefill default is BREAKABLE on CUDA and - # TC_PIECEWISE elsewhere, and this file runs on the CPU runner. + # disabled elsewhere, and this file runs on the CPU runner. self.assertEqual(backends([]), (False, Backend.FULL, default_prefill_backend())) self.assertEqual( backends(["--disable-cuda-graph"]), diff --git a/test/registered/xpu/test_breakable_cuda_graph.py b/test/registered/xpu/test_breakable_cuda_graph.py index 680e2a2595ec..dd95d70084cc 100644 --- a/test/registered/xpu/test_breakable_cuda_graph.py +++ b/test/registered/xpu/test_breakable_cuda_graph.py @@ -79,7 +79,7 @@ def test_single_break(self): intermediate = torch.zeros(4, device=self.device) y = torch.zeros(4, device=self.device) - @self.eager_on_graph(enable=True) + @self.eager_on_graph def eager_op(src): return src * 2.0 @@ -102,11 +102,11 @@ def test_multiple_breaks(self): x = torch.zeros(4, device=self.device) y = torch.zeros(4, device=self.device) - @self.eager_on_graph(enable=True) + @self.eager_on_graph def add_one(src): return src + 1.0 - @self.eager_on_graph(enable=True) + @self.eager_on_graph def double(src): return src * 2.0 @@ -125,24 +125,10 @@ def double(src): get_device_module().synchronize() self.assertTrue(torch.allclose(y, torch.full((4,), 16.0, device=self.device))) - def test_eager_on_graph_disabled(self): - """@eager_on_graph(enable=False) should be a no-op passthrough.""" - - @self.eager_on_graph(enable=False) - def my_fn(x): - return x + 1.0 - - # Should just be the original function - t = torch.tensor([1.0, 2.0], device=self.device) - result = my_fn(t) - self.assertTrue( - torch.allclose(result, torch.tensor([2.0, 3.0], device=self.device)) - ) - def test_eager_on_graph_outside_capture(self): """@eager_on_graph called outside capture should run the function directly.""" - @self.eager_on_graph(enable=True) + @self.eager_on_graph def my_fn(x): return x + 1.0 @@ -157,7 +143,7 @@ def test_replay_updates_output(self): x = torch.zeros(4, device=self.device) y = torch.zeros(4, device=self.device) - @self.eager_on_graph(enable=True) + @self.eager_on_graph def scale(src): return src * 3.0 @@ -184,7 +170,7 @@ def test_eager_output_is_held_strongly_for_replay_bridge(self): x = torch.zeros(4, device=self.device) y = torch.zeros(4, device=self.device) - @self.eager_on_graph(enable=True) + @self.eager_on_graph def scale(src): return src * 3.0 diff --git a/test/registered/xpu/test_xpu_graph.py b/test/registered/xpu/test_xpu_graph.py index d5123d60aa85..39eb73a4c041 100644 --- a/test/registered/xpu/test_xpu_graph.py +++ b/test/registered/xpu/test_xpu_graph.py @@ -1,8 +1,8 @@ """ -XPU graph tests: verifies decode full-graph and prefill tc_piecewise graph +XPU graph tests: verifies decode full-graph and prefill breakable graph on Intel XPU produce valid outputs. - - TestXPUGraph : decode full-graph and prefill tc_piecewise graph enabled + - TestXPUGraph : decode full-graph and prefill breakable graph enabled together in a single bench_one_batch invocation. Usage: @@ -38,13 +38,13 @@ class TestXPUGraph(CustomTestCase): - """Decode full-graph + prefill tc_piecewise together.""" + """Decode full-graph + prefill breakable together.""" def test_full_graph_runs(self): args = [ *_COMMON_ARGS, "--cuda-graph-config", - '{"decode":{"backend":"full"},"prefill":{"backend":"tc_piecewise","tc_compiler":"eager"}}', + '{"decode":{"backend":"full"},"prefill":{"backend":"breakable"}}', "--cuda-graph-bs-prefill", "64", "128", @@ -60,7 +60,7 @@ def test_full_graph_runs(self): self.assertGreater( prefill_latency, 0, - "prefill latency must be > 0 with tc_piecewise XPU graph", + "prefill latency must be > 0 with breakable XPU graph", ) self.assertGreater( decode_throughput, 0, "decode throughput must be > 0 with full XPU graph"