refactor(runner): split BaseRunner (shared) from BaseCudaGraphRunner (#28385)
This commit is contained in:
@@ -6,8 +6,11 @@ delegates capture/replay mechanics to a pluggable
|
|||||||
BaseCudaGraphBackend chosen via cuda_graph_config.
|
BaseCudaGraphBackend chosen via cuda_graph_config.
|
||||||
|
|
||||||
Public API:
|
Public API:
|
||||||
- BaseCudaGraphRunner — abstract base; shared init + bucket
|
- BaseRunner — minimal abstract base shared by the cuda-graph runners
|
||||||
padding + capture-loop scaffolding.
|
and the eager runner (shared __init__ + abstract
|
||||||
|
can_run_graph/load_batch/execute).
|
||||||
|
- BaseCudaGraphRunner — abstract cuda-graph base; bucket padding +
|
||||||
|
capture-loop scaffolding on top of BaseRunner.
|
||||||
- DecodeCudaGraphRunner — concrete decode-phase runner.
|
- DecodeCudaGraphRunner — concrete decode-phase runner.
|
||||||
- PrefillCudaGraphRunner — concrete prefill-phase runner.
|
- PrefillCudaGraphRunner — concrete prefill-phase runner.
|
||||||
- Buffer dataclasses, capture-mode flags, the global memory pool,
|
- Buffer dataclasses, capture-mode flags, the global memory pool,
|
||||||
@@ -22,6 +25,7 @@ from sglang.srt.model_executor.runner.base_cuda_graph_runner import ( # noqa: F
|
|||||||
freeze_gc,
|
freeze_gc,
|
||||||
get_batch_sizes_to_capture,
|
get_batch_sizes_to_capture,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.runner.base_runner import BaseRunner # noqa: F401
|
||||||
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
||||||
DecodeCudaGraphRunner,
|
DecodeCudaGraphRunner,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -18,18 +18,15 @@ from __future__ import annotations
|
|||||||
import bisect
|
import bisect
|
||||||
import gc
|
import gc
|
||||||
import logging
|
import logging
|
||||||
from abc import ABC, abstractmethod
|
from abc import abstractmethod
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from typing import TYPE_CHECKING, Any, List, Sequence, Tuple
|
from typing import TYPE_CHECKING, Any, List, Sequence, Tuple
|
||||||
|
|
||||||
import torch
|
from sglang.srt.model_executor.runner.base_runner import BaseRunner
|
||||||
|
|
||||||
from sglang.srt.batch_overlap.two_batch_overlap import TboCudaGraphRunnerPlugin
|
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import require_gathered_buffer
|
from sglang.srt.utils import require_gathered_buffer
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
|
||||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
|
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
|
||||||
@@ -103,7 +100,7 @@ def get_batch_sizes_to_capture(
|
|||||||
return capture_bs, compile_bs
|
return capture_bs, compile_bs
|
||||||
|
|
||||||
|
|
||||||
class BaseCudaGraphRunner(ABC):
|
class BaseCudaGraphRunner(BaseRunner):
|
||||||
"""Abstract base for phase-specific cuda-graph runners.
|
"""Abstract base for phase-specific cuda-graph runners.
|
||||||
|
|
||||||
A subclass (DecodeCudaGraphRunner / PrefillCudaGraphRunner) owns one
|
A subclass (DecodeCudaGraphRunner / PrefillCudaGraphRunner) owns one
|
||||||
@@ -112,19 +109,18 @@ class BaseCudaGraphRunner(ABC):
|
|||||||
selection, static buffer population, attention metadata init,
|
selection, static buffer population, attention metadata init,
|
||||||
replay dispatch, and output slicing.
|
replay dispatch, and output slicing.
|
||||||
|
|
||||||
Methods:
|
Adds the capture/shape machinery on top of BaseRunner:
|
||||||
- can_run_graph(forward_batch) — should forward_batch go through cuda
|
|
||||||
graph replay (vs eager fallback)?
|
|
||||||
- capture_prepare(size, ...) — build the dummy ForwardBatch and
|
- capture_prepare(size, ...) — build the dummy ForwardBatch and
|
||||||
per-shape local state needed by capture_one_shape.
|
per-shape local state needed by capture_one_shape.
|
||||||
- capture() — one-time setup; iterates over shapes and calls
|
- capture() — one-time setup; iterates over shapes and calls
|
||||||
capture_one_shape for each.
|
capture_one_shape for each.
|
||||||
- capture_one_shape(size, ...) — drive one model forward at this
|
- capture_one_shape(size, ...) — drive one model forward at this
|
||||||
shape into the backend's captured artifact.
|
shape into the backend's captured artifact.
|
||||||
- load_batch(forward_batch, ...) — pad to the nearest captured
|
- _pad_to_bucket(...) — round a raw shape up to the nearest captured
|
||||||
bucket, populate static input buffers, init attention metadata.
|
bucket.
|
||||||
- execute(forward_batch, ...) — dispatch one batch through cuda
|
|
||||||
graph replay.
|
Inherits from BaseRunner: __init__ and the abstract
|
||||||
|
can_run_graph / load_batch / execute.
|
||||||
|
|
||||||
Notes:
|
Notes:
|
||||||
- buffers and backend are populated by the subclass before
|
- buffers and backend are populated by the subclass before
|
||||||
@@ -135,22 +131,11 @@ class BaseCudaGraphRunner(ABC):
|
|||||||
buffers: ForwardInputBuffers
|
buffers: ForwardInputBuffers
|
||||||
backend: BaseCudaGraphBackend
|
backend: BaseCudaGraphBackend
|
||||||
|
|
||||||
def __init__(self, model_runner: ModelRunner) -> None:
|
|
||||||
self.model_runner = model_runner
|
|
||||||
self.device = model_runner.device
|
|
||||||
self.device_module = torch.get_device_module(self.device)
|
|
||||||
self.tp_size = model_runner.server_args.tp_size
|
|
||||||
self.dp_size = model_runner.server_args.dp_size
|
|
||||||
self.pp_size = model_runner.server_args.pp_size
|
|
||||||
self.attn_tp_size = get_parallel().attn_tp_size
|
|
||||||
self.attn_tp_rank = get_parallel().attn_tp_rank
|
|
||||||
self.tbo_plugin = TboCudaGraphRunnerPlugin()
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _pad_to_bucket(raw_size: int, buckets: Sequence[int]) -> int:
|
def _pad_to_bucket(raw_size: int, buckets: Sequence[int]) -> int:
|
||||||
"""Return the smallest buckets[i] >= raw_size.
|
"""Return the smallest buckets[i] >= raw_size.
|
||||||
|
|
||||||
Caller's can_run must reject raw_size > max(buckets) before
|
Caller's can_run_graph must reject raw_size > max(buckets) before
|
||||||
reaching load_batch; this assertion makes the contract
|
reaching load_batch; this assertion makes the contract
|
||||||
explicit (bisect_left returns len(buckets) when the value
|
explicit (bisect_left returns len(buckets) when the value
|
||||||
exceeds all buckets, which would otherwise IndexError below
|
exceeds all buckets, which would otherwise IndexError below
|
||||||
@@ -158,14 +143,11 @@ class BaseCudaGraphRunner(ABC):
|
|||||||
"""
|
"""
|
||||||
assert raw_size <= buckets[-1], (
|
assert raw_size <= buckets[-1], (
|
||||||
f"size {raw_size} exceeds max captured bucket {buckets[-1]}; "
|
f"size {raw_size} exceeds max captured bucket {buckets[-1]}; "
|
||||||
f"can_run should have rejected this batch"
|
f"can_run_graph should have rejected this batch"
|
||||||
)
|
)
|
||||||
index = bisect.bisect_left(buckets, raw_size)
|
index = bisect.bisect_left(buckets, raw_size)
|
||||||
return buckets[index]
|
return buckets[index]
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def can_run_graph(self, forward_batch: ForwardBatch) -> bool: ...
|
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def capture_prepare(self, size: int, *args, **kwargs) -> Any: ...
|
def capture_prepare(self, size: int, *args, **kwargs) -> Any: ...
|
||||||
|
|
||||||
@@ -174,17 +156,3 @@ class BaseCudaGraphRunner(ABC):
|
|||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def capture_one_shape(self, size: int, *args, **kwargs) -> Any: ...
|
def capture_one_shape(self, size: int, *args, **kwargs) -> Any: ...
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def load_batch(
|
|
||||||
self,
|
|
||||||
forward_batch: ForwardBatch,
|
|
||||||
**kwargs,
|
|
||||||
) -> Any: ...
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def execute(
|
|
||||||
self,
|
|
||||||
forward_batch: ForwardBatch,
|
|
||||||
**kwargs,
|
|
||||||
) -> Any: ...
|
|
||||||
|
|||||||
@@ -0,0 +1,61 @@
|
|||||||
|
# Copyright 2023-2026 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.
|
||||||
|
# ==============================================================================
|
||||||
|
"""Base class shared by EagerRunner and BaseCudaGraphRunner."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.batch_overlap.two_batch_overlap import TboCudaGraphRunnerPlugin
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class BaseRunner(ABC):
|
||||||
|
def __init__(self, model_runner: ModelRunner) -> None:
|
||||||
|
self.model_runner = model_runner
|
||||||
|
self.device = model_runner.device
|
||||||
|
self.device_module = torch.get_device_module(self.device)
|
||||||
|
self.tp_size = model_runner.server_args.tp_size
|
||||||
|
self.dp_size = model_runner.server_args.dp_size
|
||||||
|
self.pp_size = model_runner.server_args.pp_size
|
||||||
|
self.attn_tp_size = get_parallel().attn_tp_size
|
||||||
|
self.attn_tp_rank = get_parallel().attn_tp_rank
|
||||||
|
self.tbo_plugin = TboCudaGraphRunnerPlugin()
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def can_run_graph(self, forward_batch: ForwardBatch) -> bool: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def load_batch(
|
||||||
|
self,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
**kwargs,
|
||||||
|
) -> Any: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def execute(
|
||||||
|
self,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
**kwargs,
|
||||||
|
) -> Any: ...
|
||||||
Reference in New Issue
Block a user