chore: import upstream snapshot with attribution
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled

This commit is contained in:
wehub-resource-sync
2026-07-13 12:38:16 +08:00
commit 94057c3d3e
7152 changed files with 2120455 additions and 0 deletions
@@ -0,0 +1,49 @@
# Copyright 2024-2025 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.
# ==============================================================================
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from sglang.srt.configs.model_config import is_deepseek_dsa
from sglang.srt.speculative.eagle_draft_extend_cuda_graph_runner import (
EAGLEDraftExtendCudaGraphRunner,
)
if TYPE_CHECKING:
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker
class EAGLEDraftExtendNpuGraphRunner(EAGLEDraftExtendCudaGraphRunner):
def __init__(self, eagle_worker: EagleDraftWorker):
super().__init__(eagle_worker)
def _cache_loc_dtype(self):
return torch.int32
def _replay_graph(self, shape_key, forward_batch):
if not is_deepseek_dsa(self.model_runner.model_config.hf_config):
seq_lens = forward_batch.seq_lens_cpu.tolist() + [0] * (
self.bs - self.raw_bs
)
return self.backend.replay_with_input_update(
shape_key,
seq_lens=seq_lens,
attr_name="actual_seq_lengths_kv",
attr_type=[],
)
else:
return self.backend.replay(shape_key, forward_batch)
@@ -0,0 +1,69 @@
# Copyright 2025 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.
# ==============================================================================
from __future__ import annotations
from typing import TYPE_CHECKING, Dict, Union
import torch
from sglang.srt.configs.model_config import AttentionArch, is_deepseek_dsa
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
EAGLEDraftCudaGraphRunner,
)
if TYPE_CHECKING:
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker
class EAGLEDraftNpuGraphRunner(EAGLEDraftCudaGraphRunner):
def __init__(self, eagle_worker: EagleDraftWorker):
self._init_arch_map()
super().__init__(eagle_worker)
def _init_arch_map(self):
self.attr_name: Dict[str, str] = {
AttentionArch.MLA: "actual_seq_lengths_kv",
AttentionArch.MHA: "context_lens",
}
self.attr_type: Dict[str, Union[list, torch.Tensor]] = {
AttentionArch.MLA: [],
AttentionArch.MHA: torch.Tensor(),
}
def _cache_loc_dtype(self):
return torch.int32
def _get_update_attr_name(self):
return self.attr_name[AttentionArch.MLA]
def _get_update_attr_type(self):
return self.attr_type[AttentionArch.MLA]
def _replay_graph(self, shape_key, forward_batch):
if not is_deepseek_dsa(self.model_runner.model_config.hf_config):
seq_lens_for_each_draft_step = []
for speculative_step_id in range(self.speculative_num_steps - 1):
seq_lens_cpu = (
forward_batch.seq_lens_cpu[: self.raw_bs] + speculative_step_id + 1
)
seq_lens = seq_lens_cpu.tolist() + [0] * (self.bs - self.raw_bs)
seq_lens_for_each_draft_step.append(seq_lens)
attr_name = self._get_update_attr_name()
cpu_update_input = [{attr_name: sl} for sl in seq_lens_for_each_draft_step]
return self.backend.replay_with_input_update(
shape_key, seq_lens=None, cpu_update_input=cpu_update_input
)
else:
return self.backend.replay(shape_key, forward_batch)
@@ -0,0 +1,54 @@
# Copyright 2024-2025 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.
# ==============================================================================
from __future__ import annotations
import logging
from typing import TYPE_CHECKING
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import (
MultiLayerEagleDraftExtendCudaGraphRunner,
MultiLayerEagleMultiStepDraftExtendCudaGraphRunner,
)
logger = logging.getLogger(__name__)
if TYPE_CHECKING:
pass
class MultiLayerEagleDraftExtendNpuGraphRunner(
MultiLayerEagleDraftExtendCudaGraphRunner
):
def _replay_graph(self, shape_key, forward_batch):
seq_lens = self.buffers.seq_lens_cpu[: self.raw_bs].tolist() + [0] * (
self.bs - self.raw_bs
)
return self.backend.replay_with_input_update(
shape_key,
seq_lens=seq_lens,
attr_name="actual_seq_kvlen",
attr_type=[],
)
class MultiLayerEagleMultiStepDraftExtendNpuGraphRunner(
MultiLayerEagleMultiStepDraftExtendCudaGraphRunner
):
def _create_runner(self, step: int) -> MultiLayerEagleDraftExtendNpuGraphRunner:
return MultiLayerEagleDraftExtendNpuGraphRunner(self.eagle_worker, step)
def _cuda_graph_disabled(self) -> bool:
return cuda_graph_fully_disabled()
@@ -0,0 +1,177 @@
"""NPUCudaGraphBackend — Ascend NPU full-graph capture (torch.npu.NPUGraph).
Mirrors FullCudaGraphBackend with two differences:
- Captures via torch.npu.graph(...) into torch.npu.NPUGraph.
- replay_with_input_update(shape_key, seq_lens, attr_name) rebinds
the recorded graph's input bindings for variable seq_lens at replay
time (NPU's NPUGraph.update(...) API).
torch.npu is imported lazily inside methods so the module loads on
non-NPU hosts.
"""
from __future__ import annotations
import threading
from contextlib import AbstractContextManager, contextmanager
from functools import partial
from typing import TYPE_CHECKING, Any, Callable, Dict, Optional
import numpy as np
import torch
from sglang.srt.constants import GPU_MEMORY_TYPE_CUDA_GRAPH
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
set_graph_pool_id,
)
from sglang.srt.model_executor.runner.shape_key import ShapeKey
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
BaseCudaGraphBackend,
)
from sglang.srt.utils import empty_context, get_bool_env_var
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.runner.base_cuda_graph_runner import (
BaseCudaGraphRunner,
)
class NPUCudaGraphBackend(BaseCudaGraphBackend):
"""One torch.npu.NPUGraph per shape; attention metadata captured
inside the graph. replay_with_input_update substitutes fresh
seq_lens without re-recording."""
def __init__(
self,
cuda_graph_runner: BaseCudaGraphRunner,
*,
enable_memory_saver: bool = False,
) -> None:
self._graphs: Dict[Any, Any] = {}
self._outputs: Dict[Any, Any] = {}
self._pool = None
self._device_module = cuda_graph_runner.device_module
self._tp_group = cuda_graph_runner.model_runner.tp_group
self._capture_stream = None
self._memory_saver_adapter: Optional[Any] = TorchMemorySaverAdapter.create(
enable=enable_memory_saver
and get_bool_env_var("SGLANG_MEMORY_SAVER_CUDA_GRAPH")
)
self._enable_torch_compile = getattr(
cuda_graph_runner, "enable_torch_compile", False
)
@contextmanager
def capture_session(self, stream):
if self._pool is None:
self._pool = self._device_module.graph_pool_handle()
set_graph_pool_id(self._pool)
self._capture_stream = stream
try:
yield
finally:
self._capture_stream = None
def capture_one(
self,
shape_key: ShapeKey,
forward_fn: Callable[[], Any],
dummies: Optional[Any] = None,
post_warmup_hook: Optional[Callable[[], None]] = None,
) -> None:
import torch_npu # noqa: F401 (verifies NPU availability)
# Two warmups so kernels are loaded and one-time setup is paid before capture.
# post_warmup_hook lets the attention backend reset state that warmup mutated.
for _ in range(2):
self._device_module.synchronize()
self._tp_group.barrier()
forward_fn()
if post_warmup_hook is not None:
post_warmup_hook()
graph = torch.npu.NPUGraph()
if self._enable_torch_compile:
skip_guard_context = torch.compiler.set_stance(skip_guard_eval_unsafe=True)
else:
skip_guard_context = empty_context()
graph_ctx: Callable[..., AbstractContextManager]
if (
self._memory_saver_adapter is not None
and self._memory_saver_adapter.enabled
):
graph_ctx = partial(
self._memory_saver_adapter.cuda_graph,
tag=GPU_MEMORY_TYPE_CUDA_GRAPH,
)
else:
graph_ctx = torch.npu.graph
with skip_guard_context, graph_ctx(
graph,
pool=self._pool,
stream=self._capture_stream,
auto_dispatch_capture=True,
):
out = forward_fn()
self._graphs[shape_key] = graph
self._outputs[shape_key] = out
def can_run(self, forward_batch: ForwardBatch, shape_key: ShapeKey) -> bool:
return shape_key in self._graphs
@contextmanager
def replay_session(self):
yield
def replay(
self,
shape_key: ShapeKey,
static_forward_batch: ForwardBatch,
**kwargs,
) -> Any:
self._graphs[shape_key].replay()
return self._outputs[shape_key]
def replay_with_input_update(
self,
shape_key: ShapeKey,
seq_lens: Any,
attr_name: str = None,
attr_type: Any = None,
cpu_update_input: list = None,
) -> Any:
"""Rebind seq_lens on the recorded NPU graph in a background
thread, then replay. Used when the model is not deepseek-nsa.
Two calling conventions:
1. (legacy) seq_lens + attr_name + attr_type:
Constructs cpu_update_input=[{attr_name: seq_lens}] internally.
2. cpu_update_input: A list of {attr_name: seq_lens} dicts,
one per speculative step. Used by EAGLE draft runners.
"""
if cpu_update_input is None:
if isinstance(attr_type, torch.Tensor):
seq_lens = torch.from_numpy(np.array(seq_lens).astype(np.int32))
cpu_update_input = [{attr_name: seq_lens}]
graph = self._graphs[shape_key]
def _update():
graph.update(cpu_update_input=cpu_update_input)
thread = threading.Thread(target=_update)
thread.start()
graph.replay()
thread.join()
return self._outputs[shape_key]
def cleanup(self) -> None:
self._graphs.clear()
self._outputs.clear()
self._pool = None
@@ -0,0 +1,284 @@
# 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.
# ==============================================================================
"""Run the model with NPU graph and torch.compile.
NPUGraphRunner is a thin subclass of DecodeCudaGraphRunner: the
factory returns NPUCudaGraphBackend for NPU devices, so all
capture/replay mechanics live in the backend. This class adds:
- NPU-specific patch_model monkey-patch for the decode-Full +
torch.compile path.
- Profile context override (NPU profiler emits to disk, not in-mem).
- Replay override that issues an async NPUGraph.update for
seq_lens before replay (skipped for deepseek-nsa).
- Smaller cache_loc dtype (int32 instead of int64).
"""
from __future__ import annotations
import logging
import os
from contextlib import contextmanager
from pathlib import Path
from typing import TYPE_CHECKING, Dict, Optional, Union
import numpy as np
import torch
from sglang.srt.configs.model_config import (
AttentionArch,
is_deepseek_dsa,
is_deepseek_v4,
)
from sglang.srt.distributed.parallel_state import GroupCoordinator
from sglang.srt.environ import envs
from sglang.srt.model_executor.runner import DecodeCudaGraphRunner
from sglang.srt.utils import (
empty_context,
get_bool_env_var,
get_compiler_backend,
is_npu,
)
is_npu = is_npu()
if is_npu:
import torch_npu
from torch_npu.profiler import ProfilerActivity, profile
logger = logging.getLogger(__name__)
if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
@contextmanager
def patch_model_npu(
model: torch.nn.Module,
enable_compile: bool,
num_tokens: int,
tp_group: GroupCoordinator,
):
if enable_compile:
backend = get_compiler_backend("npugraph_ex")
yield torch.compile(
torch.no_grad()(model.forward),
fullgraph=True,
dynamic=False,
backend=backend,
)
else:
yield model.forward
class NPUGraphRunner(DecodeCudaGraphRunner):
"""A NPUGraphRunner runs the forward pass of a model with NPU graph and torch.compile."""
def __init__(
self,
model_runner: ModelRunner,
*,
attn_backend=None,
speculative_num_steps: Optional[int] = None,
speculative_num_draft_tokens: Optional[int] = None,
):
# NPU patch_model override: monkey-patch torch_compile_decoration's
# patch_model with the NPU-specific version.
from sglang.srt.compilation import torch_compile_decoration
torch_compile_decoration.patch_model = patch_model_npu
super().__init__(
model_runner,
attn_backend=attn_backend,
speculative_num_steps=speculative_num_steps,
speculative_num_draft_tokens=speculative_num_draft_tokens,
)
self.update_attr_name = None
self.update_attr_type = None
self.model_runner = model_runner
self._init_arch_map()
self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False")
self.if_use_v2 = any(
arch
in ("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM", "Step3p5ForCausalLM")
for arch in (model_runner.model_config.hf_config.architectures or [])
)
def _init_arch_map(self):
if self.is_dllm:
self.attr_name: Dict[str, str] = {
AttentionArch.MLA: "actual_seq_lengths_kv",
AttentionArch.MHA: "actual_seq_lengths_kv",
"TARGET_VERIFY": "actual_seq_kvlen",
}
else:
self.attr_name: Dict[str, str] = {
AttentionArch.MLA: "actual_seq_lengths_kv",
AttentionArch.MHA: "context_lens",
"TARGET_VERIFY": "actual_seq_kvlen",
}
self.attr_type: Dict[str, Union[list, torch.Tensor]] = {
AttentionArch.MLA: [],
AttentionArch.MHA: torch.Tensor(),
"TARGET_VERIFY": [],
}
def _create_device_graph(self):
return torch.npu.NPUGraph()
def _capture_graph(self, graph, pool, stream, run_once_fn):
if self.enable_torch_compile:
skip_guard_context = torch.compiler.set_stance(skip_guard_eval_unsafe=True)
else:
skip_guard_context = empty_context()
with (
skip_guard_context,
torch.npu.graph(
graph,
pool=pool,
stream=stream,
auto_dispatch_capture=True,
),
):
out = run_once_fn()
return out
def _get_update_attr_name(self):
if self.if_use_v2:
return self.attr_name["TARGET_VERIFY"]
return self.attr_name[AttentionArch.MLA]
def _get_update_attr_type(self):
if self.if_use_v2:
return self.attr_type["TARGET_VERIFY"]
return self.attr_type[AttentionArch.MLA]
def _update_inputs(self, seq_lens):
if isinstance(self.update_attr_type, torch.Tensor):
seq_lens = torch.from_numpy(np.array(seq_lens).astype(np.int32))
self.graphs[self.bs].update(
cpu_update_input=[{self.update_attr_name: seq_lens}]
)
def _cache_loc_dtype(self):
return torch.int32
def _init_profile_context_and_memory_record(self):
output_dir = os.path.join(
os.getenv("SGLANG_TORCH_PROFILER_DIR", "/tmp"), "graph_capture_profile"
)
if not Path(output_dir).exists():
Path(output_dir).mkdir(parents=True, exist_ok=True)
logger.info(
f"Profiling starts for graph capture for NPU. Traces will be saved to: {output_dir}"
)
experimental_config = torch_npu.profiler._ExperimentalConfig(
export_type=[torch_npu.profiler.ExportType.Text],
profiler_level=torch_npu.profiler.ProfilerLevel.Level1,
)
profile_context = profile(
activities=[ProfilerActivity.CPU, ProfilerActivity.NPU],
record_shapes=True,
profile_memory=True,
on_trace_ready=torch_npu.profiler.tensorboard_trace_handler(
output_dir, async_mode=True
),
experimental_config=experimental_config,
)
return profile_context
def _post_process_after_profile(self, prof_context):
# for NPU, profile data will be saved to disk for further analysis.
pass
def execute(
self,
forward_batch: ForwardBatch,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
if forward_batch.needs_forward_metadata_init():
self.load_batch(forward_batch, pp_proxy_tensors)
else:
# In speculative decoding, these two fields are still needed.
self.buffers.input_ids[: self.raw_num_token].copy_(forward_batch.input_ids)
self.buffers.positions[: self.raw_num_token].copy_(forward_batch.positions)
if (
self.model_runner.spec_algorithm.is_dflash()
and self.model_runner.is_draft_worker
and forward_batch.input_embeds is not None
):
self.buffers.input_embeds[: self.raw_num_token].copy_(
forward_batch.input_embeds
)
if (
envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.get()
and forward_batch.mrope_positions is not None
):
self.buffers.mrope_positions[:, : self.raw_num_token].copy_(
forward_batch.mrope_positions
)
graph_key = self._make_graph_key(self.bs)
if not (
is_deepseek_dsa(self.model_runner.model_config.hf_config)
or is_deepseek_v4(self.model_runner.model_config.hf_config)
):
if forward_batch.forward_mode.is_target_verify():
seq_lens_cpu = forward_batch.seq_lens.cpu() + self.num_tokens_per_bs
seq_lens = seq_lens_cpu.tolist() + [0] * (self.bs - self.raw_bs)
else:
seq_lens = forward_batch.seq_lens.cpu().tolist() + [0] * (
self.bs - self.raw_bs
)
output = self.backend.replay_with_input_update(
graph_key,
seq_lens=seq_lens,
attr_name=self._get_update_attr_name(),
attr_type=self._get_update_attr_type(),
)
else:
output = self.backend.replay(graph_key, forward_batch)
if isinstance(output, LogitsProcessorOutput):
if self.is_dllm:
next_token_logits = None
full_logits = (
output.full_logits[: self.raw_num_token]
if output.full_logits is not None
else None
)
else:
full_logits = None
next_token_logits = (
output.next_token_logits[: self.raw_num_token]
if output.next_token_logits is not None
else None
)
return LogitsProcessorOutput(
next_token_logits=next_token_logits,
full_logits=full_logits,
hidden_states=(
output.hidden_states[: self.raw_num_token]
if output.hidden_states is not None
else None
),
)
else:
assert isinstance(output, PPProxyTensors)
return PPProxyTensors({k: v[: self.bs] for k, v in output.tensors.items()})
@@ -0,0 +1,229 @@
# Copyright 2023-2025 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.
# ==============================================================================
"""ViT NPU Graph Runner class."""
from __future__ import annotations
from typing import Dict, Hashable, List, Optional, Tuple
import torch
import torch.nn as nn
import torch_npu
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
set_graph_pool_id,
)
from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
from sglang.srt.runtime_context import get_server_args
class ViTNpuGraphRunner(ViTCudaGraphRunner):
"""Generic ViT NPU Graph Runner.
This runner captures the "blocks + merger + deepstack merger (optional)" part
of a vision transformer into a NPU graph and replays it for identical shapes.
Optional for Qwen3 deepstack:
- vit.deepstack_vision_indexes: Sequence[int]
- vit.deepstack_merger_list: nn.ModuleList (same length as deepstack_vision_indexes)
"""
_graph_memory_pool = None
def __init__(
self,
vit: nn.Module,
) -> None:
super().__init__(vit)
self.device_module = torch.get_device_module(self.device)
self.cu_seq_lens: Dict[Hashable, torch.Tensor] = {}
# rotary position buffers shared across graphs
self.sin_cos_ws: Dict[Hashable, Tuple[torch.Tensor, torch.Tensor]] = {}
@property
def device(self) -> torch.device:
return self.vit.device
@property
def dtype(self) -> torch.dtype:
return self.vit.dtype
def _create_graph(
self,
graph_key: int,
):
graph = torch_npu.npu.NPUGraph()
vit = self.vit
override_backend = get_server_args().mm_attention_backend
with torch_npu.npu.graph(graph, pool=ViTNpuGraphRunner._graph_memory_pool):
y = None
deepstack_outs: List[torch.Tensor] = []
deepstack_capture_idx = 0
for layer_num, blk in enumerate(vit.blocks):
if override_backend == "ascend_attn":
cu_seq_lens = self.cu_seq_lens[graph_key]
else:
raise RuntimeError("Not supported ViT attention backend")
if layer_num == 0:
y = blk(
self.block_input[graph_key],
cu_seqlens=cu_seq_lens,
rotary_pos_emb_cos=self.sin_cos_ws[graph_key][0],
rotary_pos_emb_sin=self.sin_cos_ws[graph_key][1],
output_ws=self.block_ws[graph_key],
)
else:
y = blk(
y,
cu_seqlens=cu_seq_lens,
rotary_pos_emb_cos=self.sin_cos_ws[graph_key][0],
rotary_pos_emb_sin=self.sin_cos_ws[graph_key][1],
output_ws=self.block_ws[graph_key],
)
# Optional deepstack support (Qwen3-VL)
if (
self._deepstack_visual_indexes
and layer_num in self._deepstack_visual_indexes
):
if self._deepstack_merger_list is None:
raise RuntimeError(
"deepstack_visual_indexes exists but deepstack_merger_list is missing."
)
deepstack_out = self._deepstack_merger_list[deepstack_capture_idx](
y
)
deepstack_outs.append(deepstack_out)
deepstack_capture_idx += 1
main_out = vit.merger(y)
if deepstack_outs:
self.block_output[graph_key] = torch.cat(
[main_out] + deepstack_outs, dim=1
)
else:
self.block_output[graph_key] = main_out
self.block_graphs[graph_key] = graph
def create_graph(
self,
x_3d: torch.Tensor, # [S, 1, H]
cu_seqlens: torch.Tensor,
rotary_pos_emb_cos: Optional[torch.Tensor] = None,
rotary_pos_emb_sin: Optional[torch.Tensor] = None,
) -> int:
vit = self.vit
graph_key = self._get_graph_key(x_3d)
if graph_key in self.block_graphs:
return graph_key
if ViTNpuGraphRunner._graph_memory_pool is None:
ViTNpuGraphRunner._graph_memory_pool = (
self.device_module.graph_pool_handle()
)
# Set graph pool id globally to be able to use symmetric memory
set_graph_pool_id(ViTNpuGraphRunner._graph_memory_pool)
# pre-allocate workspace
attn_module: VisionAttention = vit.blocks[0].attn
num_heads = attn_module.num_attention_heads_per_partition
attn_head_dim = attn_module.head_size
if graph_key not in self.block_output:
self.block_output[graph_key] = x_3d
self.block_input[graph_key] = x_3d
self.block_ws[graph_key] = torch.empty(
graph_key,
num_heads,
attn_head_dim,
device=self.device,
dtype=self.dtype,
)
if rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None:
self.sin_cos_ws[graph_key] = (rotary_pos_emb_cos, rotary_pos_emb_sin)
if graph_key not in self.cu_seq_lens:
seq_lens = cu_seqlens[1:] - cu_seqlens[:-1]
self.cu_seq_lens[graph_key] = seq_lens.to("cpu").to(torch.int32)
if rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None:
self._create_graph(
graph_key=graph_key,
)
return graph_key
def replay(
self,
graph_key: int,
x_3d: torch.Tensor,
rotary_pos_emb_cos: Optional[torch.Tensor] = None,
rotary_pos_emb_sin: Optional[torch.Tensor] = None,
output_indices: Optional[torch.Tensor] = None,
) -> torch.Tensor:
if rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None:
# update rotary workspace content
self.sin_cos_ws[graph_key][0].copy_(rotary_pos_emb_cos)
self.sin_cos_ws[graph_key][1].copy_(rotary_pos_emb_sin)
# copy input
self.block_input[graph_key].copy_(x_3d)
# replay
self.block_graphs[graph_key].replay()
out = self.block_output[graph_key]
# Optional output reordering (Qwen2.5-VL window permutation inverse)
if output_indices is not None:
out = out.index_select(0, output_indices)
return out
def run(
self,
x: torch.Tensor,
cu_seqlens: torch.Tensor,
rotary_pos_emb_cos: Optional[torch.Tensor] = None,
rotary_pos_emb_sin: Optional[torch.Tensor] = None,
output_indices: Optional[torch.Tensor] = None,
) -> torch.Tensor:
# x: [seq_len, hidden] -> [S, B=1, H]
x_3d = x.unsqueeze(1)
graph_key = self._get_graph_key(x_3d)
if graph_key not in self.block_graphs:
self.create_graph(
x_3d=x_3d,
cu_seqlens=cu_seqlens,
rotary_pos_emb_cos=rotary_pos_emb_cos,
rotary_pos_emb_sin=rotary_pos_emb_sin,
)
return self.replay(
graph_key=graph_key,
x_3d=x_3d,
rotary_pos_emb_cos=rotary_pos_emb_cos,
rotary_pos_emb_sin=rotary_pos_emb_sin,
output_indices=output_indices,
)