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
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:
+49
@@ -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)
|
||||
+54
@@ -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,
|
||||
)
|
||||
Reference in New Issue
Block a user