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:
@@ -0,0 +1,213 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.vla.prefix_cache import (
|
||||
PrefixContext,
|
||||
VLADensePrefixCache,
|
||||
)
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
set_graph_pool_id,
|
||||
)
|
||||
from sglang.srt.model_executor.runner_utils.pool import (
|
||||
get_or_create_global_graph_memory_pool,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class VLADenoiseGraphSignature:
|
||||
batch_size: int
|
||||
prefix_len: int
|
||||
action_horizon: int
|
||||
action_dim: int
|
||||
dtype: str
|
||||
parallel_layout: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class _CapturedDenoiseGraph:
|
||||
graph: torch.cuda.CUDAGraph
|
||||
static_prefix_context: PrefixContext
|
||||
static_x_t: torch.Tensor
|
||||
static_timestep: torch.Tensor
|
||||
static_output: torch.Tensor
|
||||
current_context_id: int | None = None
|
||||
current_context_digest: str | None = None
|
||||
|
||||
|
||||
def _clone_past_key_values(past_key_values: Any) -> Any:
|
||||
return VLADensePrefixCache(
|
||||
tuple(
|
||||
(keys.detach().clone(), values.detach().clone(), sliding_window)
|
||||
for keys, values, sliding_window in past_key_values
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _copy_past_key_values_(dst: Any, src: Any) -> None:
|
||||
for (dst_keys, dst_values, _), (src_keys, src_values, _) in zip(
|
||||
dst, src, strict=True
|
||||
):
|
||||
dst_keys.copy_(src_keys)
|
||||
dst_values.copy_(src_values)
|
||||
|
||||
|
||||
def _clone_prefix_context(prefix_context: PrefixContext) -> PrefixContext:
|
||||
return PrefixContext(
|
||||
past_key_values=_clone_past_key_values(prefix_context.past_key_values),
|
||||
prefix_pad_masks=prefix_context.prefix_pad_masks.detach().clone(),
|
||||
prefix_len=prefix_context.prefix_len,
|
||||
layout=dict(prefix_context.layout),
|
||||
cache_key_digest=prefix_context.cache_key_digest,
|
||||
)
|
||||
|
||||
|
||||
def _copy_prefix_context_(dst: PrefixContext, src: PrefixContext) -> None:
|
||||
dst.prefix_pad_masks.copy_(src.prefix_pad_masks)
|
||||
_copy_past_key_values_(dst.past_key_values, src.past_key_values)
|
||||
dst.cache_key_digest = src.cache_key_digest
|
||||
|
||||
|
||||
class VLADenoiseGraphRunner:
|
||||
"""Full CUDA graph runner for one VLA action-denoise step.
|
||||
|
||||
Each signature owns fixed input and output buffers. This does not use
|
||||
diffusion BCG and does not capture prefix encoding or token decode.
|
||||
"""
|
||||
|
||||
def __init__(self, enabled: bool = True):
|
||||
self.enabled = enabled
|
||||
self._captured: dict[VLADenoiseGraphSignature, _CapturedDenoiseGraph] = {}
|
||||
self._disabled_signatures: set[VLADenoiseGraphSignature] = set()
|
||||
self._capture_stream: torch.cuda.Stream | None = None
|
||||
self._graph_pool: Any = None
|
||||
|
||||
def _sync_context_if_needed(
|
||||
self,
|
||||
captured: _CapturedDenoiseGraph,
|
||||
prefix_context: PrefixContext,
|
||||
) -> None:
|
||||
context_id = id(prefix_context.past_key_values)
|
||||
context_digest = prefix_context.cache_key_digest
|
||||
if (
|
||||
context_digest is not None
|
||||
and captured.current_context_digest == context_digest
|
||||
):
|
||||
captured.current_context_id = context_id
|
||||
return
|
||||
if captured.current_context_id == context_id:
|
||||
return
|
||||
_copy_prefix_context_(captured.static_prefix_context, prefix_context)
|
||||
captured.current_context_id = context_id
|
||||
captured.current_context_digest = context_digest
|
||||
|
||||
def _capture(
|
||||
self,
|
||||
signature: VLADenoiseGraphSignature,
|
||||
step_fn: Callable[..., torch.Tensor],
|
||||
prefix_context: PrefixContext,
|
||||
x_t: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
) -> _CapturedDenoiseGraph:
|
||||
static_prefix_context = _clone_prefix_context(prefix_context)
|
||||
static_x_t = x_t.detach().clone()
|
||||
static_timestep = timestep.detach().clone()
|
||||
|
||||
device_module = torch.get_device_module(x_t.device)
|
||||
if self._capture_stream is None:
|
||||
self._capture_stream = device_module.Stream(device=x_t.device)
|
||||
if self._graph_pool is None:
|
||||
self._graph_pool = get_or_create_global_graph_memory_pool(device_module)
|
||||
set_graph_pool_id(self._graph_pool)
|
||||
|
||||
# warm up lazy kernels and workspaces before capture
|
||||
device_module.synchronize()
|
||||
with device_module.stream(self._capture_stream), torch.inference_mode():
|
||||
step_fn(
|
||||
static_prefix_context,
|
||||
static_x_t,
|
||||
static_timestep,
|
||||
)
|
||||
self._capture_stream.synchronize()
|
||||
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with (
|
||||
device_module.graph(
|
||||
cuda_graph=graph,
|
||||
pool=self._graph_pool,
|
||||
stream=self._capture_stream,
|
||||
),
|
||||
torch.inference_mode(),
|
||||
):
|
||||
static_output = step_fn(
|
||||
static_prefix_context,
|
||||
static_x_t,
|
||||
static_timestep,
|
||||
)
|
||||
self._capture_stream.synchronize()
|
||||
|
||||
captured = _CapturedDenoiseGraph(
|
||||
graph=graph,
|
||||
static_prefix_context=static_prefix_context,
|
||||
static_x_t=static_x_t,
|
||||
static_timestep=static_timestep,
|
||||
static_output=static_output,
|
||||
current_context_id=id(prefix_context.past_key_values),
|
||||
current_context_digest=prefix_context.cache_key_digest,
|
||||
)
|
||||
self._captured[signature] = captured
|
||||
logger.info(
|
||||
"Captured VLA denoise CUDA graph: batch=%d prefix=%d action=%dx%d "
|
||||
"dtype=%s",
|
||||
signature.batch_size,
|
||||
signature.prefix_len,
|
||||
signature.action_horizon,
|
||||
signature.action_dim,
|
||||
signature.dtype,
|
||||
)
|
||||
return captured
|
||||
|
||||
def capture_or_run(
|
||||
self,
|
||||
signature: VLADenoiseGraphSignature,
|
||||
step_fn: Callable[..., torch.Tensor],
|
||||
prefix_context: PrefixContext,
|
||||
x_t: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
if not self.enabled or signature in self._disabled_signatures:
|
||||
return step_fn(prefix_context, x_t, timestep)
|
||||
|
||||
if x_t.device.type != "cuda":
|
||||
return step_fn(prefix_context, x_t, timestep)
|
||||
|
||||
captured = self._captured.get(signature)
|
||||
try:
|
||||
if captured is None:
|
||||
captured = self._capture(
|
||||
signature, step_fn, prefix_context, x_t, timestep
|
||||
)
|
||||
captured.graph.replay()
|
||||
else:
|
||||
self._sync_context_if_needed(captured, prefix_context)
|
||||
captured.static_x_t.copy_(x_t)
|
||||
captured.static_timestep.copy_(timestep)
|
||||
captured.graph.replay()
|
||||
return captured.static_output
|
||||
except Exception:
|
||||
self._disabled_signatures.add(signature)
|
||||
self._captured.pop(signature, None)
|
||||
logger.warning(
|
||||
"VLA denoise CUDA graph disabled for signature %s",
|
||||
signature,
|
||||
exc_info=True,
|
||||
)
|
||||
return step_fn(prefix_context, x_t, timestep)
|
||||
Reference in New Issue
Block a user