Support skip-softmax attention (#19089)

This commit is contained in:
Shu Wang
2026-03-28 15:55:48 -07:00
committed by GitHub
parent 40ff652862
commit efebcab43e
8 changed files with 122 additions and 0 deletions
+1
View File
@@ -1937,6 +1937,7 @@ if __name__ == "__main__":
"mmmu",
"image",
"mooncake",
"longbench_v2",
],
help="Name of the dataset to benchmark on.",
)
@@ -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,
}
@@ -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
+4
View File
@@ -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)
@@ -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)
@@ -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
)
@@ -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.