@@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.attention.flashinfer_backend import (
|
from sglang.srt.layers.attention.flashinfer_backend import (
|
||||||
FlashInferAttnBackend,
|
FlashInferAttnBackend,
|
||||||
FlashInferMultiStepDraftBackend,
|
FlashInferMultiStepDraftBackend,
|
||||||
@@ -32,9 +33,9 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.speculative.spec_info import SpecInput
|
from sglang.srt.speculative.spec_info import SpecInput
|
||||||
|
|
||||||
# Constants
|
# Constants
|
||||||
DEFAULT_WORKSPACE_SIZE_MB = (
|
# Default workspace size in MB for TRTLLM MHA
|
||||||
512 # Memory workspace size in MB, todo(Yingyi): read from config
|
# Can be configured via SGLANG_FLASHINFER_WORKSPACE_SIZE environment variable
|
||||||
)
|
DEFAULT_WORKSPACE_SIZE_MB = 512
|
||||||
|
|
||||||
# Reuse this workspace buffer across all TRTLLM MHA wrappers
|
# Reuse this workspace buffer across all TRTLLM MHA wrappers
|
||||||
global_zero_init_workspace_buffer = None
|
global_zero_init_workspace_buffer = None
|
||||||
@@ -67,6 +68,15 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
kv_last_page_len_buf: Optional[torch.Tensor] = None,
|
kv_last_page_len_buf: Optional[torch.Tensor] = None,
|
||||||
speculative_step_id: int = 0,
|
speculative_step_id: int = 0,
|
||||||
):
|
):
|
||||||
|
# Capture workspace size before super().__init__() to preserve user's
|
||||||
|
# SGLANG_FLASHINFER_WORKSPACE_SIZE setting (may be overridden by parent)
|
||||||
|
env_var = envs.SGLANG_FLASHINFER_WORKSPACE_SIZE
|
||||||
|
workspace_size_bytes = (
|
||||||
|
env_var.get()
|
||||||
|
if env_var.is_set()
|
||||||
|
else DEFAULT_WORKSPACE_SIZE_MB * 1024 * 1024
|
||||||
|
)
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
model_runner, skip_prefill, kv_indptr_buf, kv_last_page_len_buf
|
model_runner, skip_prefill, kv_indptr_buf, kv_last_page_len_buf
|
||||||
)
|
)
|
||||||
@@ -85,7 +95,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
|
|
||||||
# Workspace allocation
|
# Workspace allocation
|
||||||
self.workspace_size = DEFAULT_WORKSPACE_SIZE_MB * 1024 * 1024
|
self.workspace_size = workspace_size_bytes
|
||||||
# Allocate buffers
|
# Allocate buffers
|
||||||
global global_zero_init_workspace_buffer
|
global global_zero_init_workspace_buffer
|
||||||
if global_zero_init_workspace_buffer is None:
|
if global_zero_init_workspace_buffer is None:
|
||||||
|
|||||||
Reference in New Issue
Block a user