Suppress expected FlashInfer TRT-LLM workspace warnings (#34921)
This commit is contained in:
@@ -389,6 +389,8 @@ class FlashInferWorkspaceManager:
|
||||
self.max_token_num = None
|
||||
self.hidden_dim = None
|
||||
self.dtype = None
|
||||
self.backend = None
|
||||
self.use_fp32_lamport = None
|
||||
self.initialized = False
|
||||
# Track max sizes ever requested so the workspace only grows (fewer recreates)
|
||||
self._max_token_num_seen: Optional[int] = None
|
||||
@@ -523,19 +525,24 @@ class FlashInferWorkspaceManager:
|
||||
self.max_token_num = alloc_token_num
|
||||
self.hidden_dim = alloc_hidden_dim
|
||||
self.dtype = dtype or torch.bfloat16
|
||||
self.backend = self.workspace.backend
|
||||
self.use_fp32_lamport = (
|
||||
self.workspace.metadata["use_fp32_lamport"]
|
||||
if self.backend == "trtllm"
|
||||
else None
|
||||
)
|
||||
self.initialized = True
|
||||
|
||||
backend_name = getattr(self.workspace, "backend", "unknown")
|
||||
if not self._logged_init:
|
||||
logger.info(
|
||||
f"FlashInfer AllReduce Fusion enabled and workspace initialized: "
|
||||
f"backend={backend_name}, rank={rank}, world_size={world_size}, "
|
||||
f"backend={self.backend}, rank={rank}, world_size={world_size}, "
|
||||
f"max_token_num={self.max_token_num}, hidden_dim={self.hidden_dim}"
|
||||
)
|
||||
self._logged_init = True
|
||||
else:
|
||||
logger.debug(
|
||||
f"FlashInfer workspace re-initialized: backend={backend_name}, "
|
||||
f"FlashInfer workspace re-initialized: backend={self.backend}, "
|
||||
f"rank={rank}, world_size={world_size}"
|
||||
)
|
||||
except Exception as e:
|
||||
@@ -611,6 +618,8 @@ class FlashInferWorkspaceManager:
|
||||
self.max_token_num = None
|
||||
self.hidden_dim = None
|
||||
self.dtype = None
|
||||
self.backend = None
|
||||
self.use_fp32_lamport = None
|
||||
self._logged_init = False
|
||||
|
||||
|
||||
@@ -923,6 +932,18 @@ def can_use_flashinfer_allreduce(
|
||||
and workspace_manager.dtype == input_.dtype
|
||||
)
|
||||
|
||||
# TRT-LLM logs a warning when its validator rejects an expected fallback.
|
||||
# Mirror only the conditions that prove its workspace is insufficient:
|
||||
# total element capacity and FP32-Lamport compatibility. MNNVL has a
|
||||
# byte/strategy-based contract and does not emit that warning, so its
|
||||
# authoritative validator remains the source of truth below.
|
||||
if workspace_manager.backend == "trtllm" and (
|
||||
token_num * hidden_dim
|
||||
> workspace_manager.max_token_num * workspace_manager.hidden_dim
|
||||
or workspace_manager.use_fp32_lamport != (input_.dtype == torch.float32)
|
||||
):
|
||||
return False
|
||||
|
||||
return workspace_manager.is_buffer_size_sufficient(
|
||||
token_num=token_num,
|
||||
hidden_dim=hidden_dim,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import contextlib
|
||||
import types
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
@@ -16,9 +16,10 @@ register_cuda_ci(est_time=30, stage="base-c", runner_config="4-gpu-gb300")
|
||||
|
||||
|
||||
class _FakeWorkspace:
|
||||
def __init__(self, backend, world_size):
|
||||
def __init__(self, backend, world_size, dtype=torch.bfloat16):
|
||||
self.backend = backend
|
||||
self.world_size = world_size
|
||||
self.metadata = {"use_fp32_lamport": dtype == torch.float32}
|
||||
|
||||
def is_buffer_size_sufficient(self, **_kwargs):
|
||||
return True
|
||||
@@ -34,7 +35,9 @@ class _FakeFlashInferComm:
|
||||
|
||||
def create_allreduce_fusion_workspace(self, **kwargs):
|
||||
self.calls.append(kwargs)
|
||||
return _FakeWorkspace(kwargs["backend"], kwargs["world_size"])
|
||||
return _FakeWorkspace(
|
||||
kwargs["backend"], kwargs["world_size"], dtype=kwargs["dtype"]
|
||||
)
|
||||
|
||||
def allreduce_fusion(
|
||||
self,
|
||||
@@ -260,15 +263,17 @@ _OTHER_GROUP_KEY = ("other_device_group", "other_cpu_group")
|
||||
|
||||
|
||||
class TestFlashInferAllReduceOnly(CustomTestCase):
|
||||
def _make_manager(self, world_size, group_key=_GROUP_KEY):
|
||||
def _make_manager(self, world_size, group_key=_GROUP_KEY, backend="trtllm"):
|
||||
manager = fusion.FlashInferWorkspaceManager()
|
||||
manager.workspace = _FakeWorkspace(None, world_size)
|
||||
manager.workspace = _FakeWorkspace(backend, world_size)
|
||||
manager.initialized = True
|
||||
manager.world_size = world_size
|
||||
manager.group = group_key
|
||||
manager.max_token_num = 2048
|
||||
manager.hidden_dim = 4096
|
||||
manager.dtype = torch.float32
|
||||
manager.backend = backend
|
||||
manager.use_fp32_lamport = True
|
||||
return manager
|
||||
|
||||
@contextlib.contextmanager
|
||||
@@ -306,7 +311,10 @@ class TestFlashInferAllReduceOnly(CustomTestCase):
|
||||
if not torch.cuda.is_available():
|
||||
self.skipTest("CUDA required for flashinfer custom op")
|
||||
world_size = 4
|
||||
with self._patched_attn_workspace(self._make_manager(world_size)):
|
||||
manager = self._make_manager(world_size)
|
||||
manager.dtype = torch.bfloat16
|
||||
manager.use_fp32_lamport = False
|
||||
with self._patched_attn_workspace(manager):
|
||||
input_ = torch.randn(8, 16, dtype=torch.bfloat16, device="cuda")
|
||||
expected = input_ * world_size
|
||||
|
||||
@@ -359,12 +367,94 @@ class TestFlashInferAllReduceOnly(CustomTestCase):
|
||||
with self._patched_attn_workspace(self._make_manager(4)):
|
||||
self.assertFalse(self._can_use(torch.randn(8, 16), world_size=2))
|
||||
|
||||
def test_rejects_when_token_num_exceeds_workspace_capacity(self):
|
||||
"""Under Dynamo the capacity check replaces is_buffer_size_sufficient().
|
||||
def test_fp32_initialization_caches_allocated_lamport_mode(self):
|
||||
"""FP32 startup allocation must remain eligible for FP32 all-reduce.
|
||||
|
||||
_FakeWorkspace.is_buffer_size_sufficient() always says yes, so this only
|
||||
passes if the compiling branch consults the manager's own allocation.
|
||||
The initialization API's legacy use_fp32_lamport argument defaults to
|
||||
False, while FlashInfer derives the allocated TRT-LLM mode from dtype.
|
||||
Eligibility must follow the workspace metadata rather than that input.
|
||||
"""
|
||||
fake_comm = _FakeFlashInferComm()
|
||||
manager = fusion.FlashInferWorkspaceManager()
|
||||
with (
|
||||
patch.object(fusion, "_flashinfer_comm", fake_comm),
|
||||
patch.object(
|
||||
fusion,
|
||||
"_create_allreduce_fusion_workspace",
|
||||
fake_comm.create_allreduce_fusion_workspace,
|
||||
),
|
||||
patch.object(
|
||||
fusion, "_preflight_check_workspace_memory", return_value=True
|
||||
),
|
||||
):
|
||||
manager.initialize(
|
||||
world_size=4,
|
||||
rank=0,
|
||||
max_token_num=8,
|
||||
hidden_dim=4096,
|
||||
backend="trtllm",
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
self.assertTrue(manager.use_fp32_lamport)
|
||||
with self._patched_attn_workspace(manager):
|
||||
self.assertTrue(
|
||||
self._can_use(
|
||||
torch.randn(8, 16, dtype=torch.float32), group_key=(None, None)
|
||||
)
|
||||
)
|
||||
|
||||
def test_rejects_when_token_num_exceeds_workspace_capacity(self):
|
||||
"""Oversized all-reduces fall back without triggering a warning.
|
||||
|
||||
TRT-LLM warns whenever its size validator rejects an operation. The
|
||||
total element count already proves that this operation cannot use the
|
||||
workspace, so the validator must not be invoked.
|
||||
"""
|
||||
manager = self._make_manager(4)
|
||||
manager.max_token_num = 8
|
||||
manager.workspace.is_buffer_size_sufficient = MagicMock(return_value=True)
|
||||
with self._patched_attn_workspace(manager):
|
||||
self.assertFalse(self._can_use(torch.randn(9, 4096)))
|
||||
|
||||
manager.workspace.is_buffer_size_sufficient.assert_not_called()
|
||||
|
||||
def test_reshaped_input_within_total_capacity_reaches_validator(self):
|
||||
manager = self._make_manager(4)
|
||||
manager.max_token_num = 8
|
||||
manager.workspace.is_buffer_size_sufficient = MagicMock(return_value=True)
|
||||
input_ = torch.randn(9, 16)
|
||||
with self._patched_attn_workspace(manager):
|
||||
self.assertTrue(self._can_use(input_))
|
||||
|
||||
manager.workspace.is_buffer_size_sufficient.assert_called_once_with(
|
||||
tp_size=4,
|
||||
num_tokens=9,
|
||||
hidden_dim=16,
|
||||
dtype=input_.dtype,
|
||||
)
|
||||
|
||||
def test_non_fp32_dtype_change_reaches_validator(self):
|
||||
manager = self._make_manager(4)
|
||||
manager.dtype = torch.bfloat16
|
||||
manager.use_fp32_lamport = False
|
||||
manager.workspace.is_buffer_size_sufficient = MagicMock(return_value=True)
|
||||
input_ = torch.randn(8, 16, dtype=torch.float16)
|
||||
with self._patched_attn_workspace(manager):
|
||||
self.assertTrue(self._can_use(input_))
|
||||
|
||||
manager.workspace.is_buffer_size_sufficient.assert_called_once()
|
||||
|
||||
def test_mnnvl_capacity_decision_reaches_validator(self):
|
||||
manager = self._make_manager(4, backend="mnnvl")
|
||||
manager.max_token_num = 8
|
||||
manager.workspace.is_buffer_size_sufficient = MagicMock(return_value=False)
|
||||
with self._patched_attn_workspace(manager):
|
||||
self.assertFalse(self._can_use(torch.randn(9, 4096)))
|
||||
|
||||
manager.workspace.is_buffer_size_sufficient.assert_called_once()
|
||||
|
||||
def test_compiling_rejects_when_token_num_exceeds_workspace_capacity(self):
|
||||
manager = self._make_manager(4)
|
||||
manager.max_token_num = 8
|
||||
with self._patched_attn_workspace(manager):
|
||||
|
||||
Reference in New Issue
Block a user