From efebcab43ed851405bd41cfdf206245945dbb8ee Mon Sep 17 00:00:00 2001 From: Shu Wang Date: Sat, 28 Mar 2026 17:55:48 -0500 Subject: [PATCH] Support skip-softmax attention (#19089) --- docs/references/environment_variables.md | 2 + python/sglang/bench_serving.py | 1 + python/sglang/benchmark/datasets/__init__.py | 2 + .../sglang/benchmark/datasets/longbench_v2.py | 104 ++++++++++++++++++ python/sglang/srt/environ.py | 4 + .../srt/layers/attention/nsa_backend.py | 2 + .../layers/attention/trtllm_mha_backend.py | 3 + .../layers/attention/trtllm_mla_backend.py | 4 + 8 files changed, 122 insertions(+) create mode 100644 python/sglang/benchmark/datasets/longbench_v2.py diff --git a/docs/references/environment_variables.md b/docs/references/environment_variables.md index 29d6b6962..1bfb7ce8f 100644 --- a/docs/references/environment_variables.md +++ b/docs/references/environment_variables.md @@ -45,6 +45,8 @@ SGLang supports various environment variables that can be used to configure its | `SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH` | Enable NCCL for gathering when preparing mlp sync batch under overlap scheduler (without this flag gloo is used for gathering) | `false` | | `SGLANG_SYMM_MEM_PREALLOC_GB_SIZE` | Size of preallocated GPU buffer (in GB) for NCCL symmetric memory pool to limit memory fragmentation. Only have an effect when server arg `--enable-symm-mem` is set. | `-1` | | `SGLANG_CUSTOM_ALLREDUCE_ALGO` | The algorithm of custom all-reduce. Set to `oneshot` or `1stage` to force use one-shot. Set to `twoshot` or `2stage` to force use two-shot. | `` | +| `SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR` | Skip-softmax threshold scale factor for TRT-LLM prefill attention in flashinfer. `None` means standard attention. See https://arxiv.org/abs/2512.12087 | `None` | +| `SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR` | Skip-softmax threshold scale factor for TRT-LLM decode attention in flashinfer. `None` means standard attention. See https://arxiv.org/abs/2512.12087 | `None` | ## DeepGEMM Configuration (Advanced Optimization) diff --git a/python/sglang/bench_serving.py b/python/sglang/bench_serving.py index 4e79ece7e..c3129aa39 100644 --- a/python/sglang/bench_serving.py +++ b/python/sglang/bench_serving.py @@ -1937,6 +1937,7 @@ if __name__ == "__main__": "mmmu", "image", "mooncake", + "longbench_v2", ], help="Name of the dataset to benchmark on.", ) diff --git a/python/sglang/benchmark/datasets/__init__.py b/python/sglang/benchmark/datasets/__init__.py index 63612d52e..615e3a241 100644 --- a/python/sglang/benchmark/datasets/__init__.py +++ b/python/sglang/benchmark/datasets/__init__.py @@ -6,6 +6,7 @@ from sglang.benchmark.datasets.generated_shared_prefix import ( GeneratedSharedPrefixDataset, ) from sglang.benchmark.datasets.image import ImageDataset +from sglang.benchmark.datasets.longbench_v2 import LongBenchV2Dataset from sglang.benchmark.datasets.mmmu import MMMUDataset from sglang.benchmark.datasets.mooncake import MooncakeDataset from sglang.benchmark.datasets.openai_dataset import OpenAIDataset @@ -24,6 +25,7 @@ DATASET_MAPPING: Dict[str, Type[BaseDataset]] = { "mmmu": MMMUDataset, "image": ImageDataset, "mooncake": MooncakeDataset, + "longbench_v2": LongBenchV2Dataset, } diff --git a/python/sglang/benchmark/datasets/longbench_v2.py b/python/sglang/benchmark/datasets/longbench_v2.py new file mode 100644 index 000000000..e8a642957 --- /dev/null +++ b/python/sglang/benchmark/datasets/longbench_v2.py @@ -0,0 +1,104 @@ +import random +from argparse import Namespace +from dataclasses import dataclass +from typing import List, Optional + +from transformers import PreTrainedTokenizerBase + +from sglang.benchmark.datasets.common import BaseDataset, DatasetRow + +LONGBENCH_V2_REPO_ID = "THUDM/LongBench-v2" +LONGBENCH_V2_DEFAULT_OUTPUT_LEN = 10 # answer letter + short explanation + + +def _format_prompt(example: dict) -> str: + return ( + f"{example['context']}\n\n" + f"Question: {example['question']}\n" + f"A. {example['choice_A']}\n" + f"B. {example['choice_B']}\n" + f"C. {example['choice_C']}\n" + f"D. {example['choice_D']}\n" + f"Answer:" + ) + + +@dataclass +class LongBenchV2Dataset(BaseDataset): + dataset_path: str + num_requests: int + fixed_output_len: Optional[int] + context_len: Optional[int] + + @classmethod + def from_args(cls, args: Namespace) -> "LongBenchV2Dataset": + return cls( + dataset_path=args.dataset_path, + num_requests=args.num_prompts, + fixed_output_len=args.sharegpt_output_len, + context_len=args.sharegpt_context_len, + ) + + def load( + self, tokenizer: PreTrainedTokenizerBase, model_id=None + ) -> List[DatasetRow]: + return sample_longbench_v2_requests( + dataset_path=self.dataset_path, + num_requests=self.num_requests, + tokenizer=tokenizer, + fixed_output_len=self.fixed_output_len, + context_len=self.context_len, + ) + + +def sample_longbench_v2_requests( + dataset_path: str, + num_requests: int, + tokenizer: PreTrainedTokenizerBase, + fixed_output_len: Optional[int] = None, + context_len: Optional[int] = None, +) -> List[DatasetRow]: + output_len = ( + fixed_output_len + if fixed_output_len is not None + else LONGBENCH_V2_DEFAULT_OUTPUT_LEN + ) + + # Load dataset + if dataset_path: + # Local file (parquet or JSON lines) + import pandas as pd + + if dataset_path.endswith(".parquet"): + df = pd.read_parquet(dataset_path) + examples = df.to_dict(orient="records") + else: + import json + + with open(dataset_path) as f: + examples = [json.loads(line) for line in f if line.strip()] + else: + from datasets import load_dataset + + ds = load_dataset(LONGBENCH_V2_REPO_ID, split="train") + examples = list(ds) + + random.shuffle(examples) + + rows: List[DatasetRow] = [] + for example in examples: + if len(rows) >= num_requests: + break + + prompt = _format_prompt(example) + prompt_ids = tokenizer(prompt).input_ids + prompt_len = len(prompt_ids) + + if context_len is not None and prompt_len + output_len > context_len: + continue + + rows.append( + DatasetRow(prompt=prompt, prompt_len=prompt_len, output_len=output_len) + ) + + return rows diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index e96abfcb1..f35c829a5 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -339,6 +339,10 @@ class Envs: SGLANG_IS_FLASHINFER_AVAILABLE = EnvBool(True) # Default to the pick from flashinfer SGLANG_FLASHINFER_WORKSPACE_SIZE = EnvInt(384 * 1024 * 1024) + # Skip-softmax threshold scale factor for TRT-LLM attention (prefill and decode separately). + # None = standard attention. See https://arxiv.org/abs/2512.12087 + SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR = EnvFloat(None) + SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR = EnvFloat(None) # TODO(mmangkad): Remove this once the FlashInfer unified allreduce-fusion # transport issue on GB200/GB300 platforms is fixed and verified resolved. SGLANG_FLASHINFER_FORCE_POSIX_FD_TRANSPORT = EnvBool(None) diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index dea5c5348..dd3754f1b 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -1790,6 +1790,7 @@ class NativeSparseAttnBackend( enable_pdl=False, is_causal=causal, return_lse=False, + skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get(), ) # Use FA3 for SM90 (Hopper/H200) @@ -2025,6 +2026,7 @@ class NativeSparseAttnBackend( sparse_mla_top_k=self.nsa_index_topk, bmm1_scale=bmm1_scale, backend="trtllm-gen", + skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), ) # Output: [batch, q_len=1, heads, v_dim] -> [batch, heads, v_dim] return out.squeeze(1) diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index 09f3f409a..205b5fadf 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -773,6 +773,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): bmm2_scale=bmm2_scale, window_left=layer.sliding_window_size, sinks=attention_sink, + skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), out_dtype=self.q_data_type, # model_runner.dtype ) @@ -855,6 +856,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): bmm2_scale=bmm2_scale, window_left=layer.sliding_window_size, sinks=attention_sink, + skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), out_dtype=self.q_data_type, # model_runner.dtype q_len_per_req=self.forward_metadata.max_seq_len_q, ) @@ -874,6 +876,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): cum_seq_lens_kv=self.forward_metadata.cu_seqlens_k, window_left=layer.sliding_window_size, sinks=attention_sink, + skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get(), out_dtype=self.q_data_type, # model_runner.dtype ) diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index b74124da2..54eb273ed 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -14,6 +14,7 @@ import triton import triton.language as tl from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph +from sglang.srt.environ import envs from sglang.srt.layers.attention.flashinfer_mla_backend import ( FlashInferMLAAttnBackend, FlashInferMLAMultiStepDraftBackend, @@ -875,6 +876,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): seq_lens=forward_batch.seq_lens.to(torch.int32), max_seq_len=metadata.max_seq_len_k, bmm1_scale=bmm1_scale, + skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), ) # Reshape output directly without slicing @@ -1062,6 +1064,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): seq_lens=metadata.seq_lens_k, max_seq_len=max_seq_len, bmm1_scale=bmm1_scale, + skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), ) if needs_unpad: @@ -1099,6 +1102,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): "bmm1_scale": q_scale * k_scale * layer.scaling, "bmm2_scale": v_scale, "cum_seq_lens_q": self.forward_prefill_metadata.cum_seq_lens, + "skip_softmax_threshold_scale_factor": envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get(), } # When chunked prefix cache is enabled, dispatch to different path for ragged attention.