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.max_token_num = None
|
||||||
self.hidden_dim = None
|
self.hidden_dim = None
|
||||||
self.dtype = None
|
self.dtype = None
|
||||||
|
self.backend = None
|
||||||
|
self.use_fp32_lamport = None
|
||||||
self.initialized = False
|
self.initialized = False
|
||||||
# Track max sizes ever requested so the workspace only grows (fewer recreates)
|
# Track max sizes ever requested so the workspace only grows (fewer recreates)
|
||||||
self._max_token_num_seen: Optional[int] = None
|
self._max_token_num_seen: Optional[int] = None
|
||||||
@@ -523,19 +525,24 @@ class FlashInferWorkspaceManager:
|
|||||||
self.max_token_num = alloc_token_num
|
self.max_token_num = alloc_token_num
|
||||||
self.hidden_dim = alloc_hidden_dim
|
self.hidden_dim = alloc_hidden_dim
|
||||||
self.dtype = dtype or torch.bfloat16
|
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
|
self.initialized = True
|
||||||
|
|
||||||
backend_name = getattr(self.workspace, "backend", "unknown")
|
|
||||||
if not self._logged_init:
|
if not self._logged_init:
|
||||||
logger.info(
|
logger.info(
|
||||||
f"FlashInfer AllReduce Fusion enabled and workspace initialized: "
|
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}"
|
f"max_token_num={self.max_token_num}, hidden_dim={self.hidden_dim}"
|
||||||
)
|
)
|
||||||
self._logged_init = True
|
self._logged_init = True
|
||||||
else:
|
else:
|
||||||
logger.debug(
|
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}"
|
f"rank={rank}, world_size={world_size}"
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -611,6 +618,8 @@ class FlashInferWorkspaceManager:
|
|||||||
self.max_token_num = None
|
self.max_token_num = None
|
||||||
self.hidden_dim = None
|
self.hidden_dim = None
|
||||||
self.dtype = None
|
self.dtype = None
|
||||||
|
self.backend = None
|
||||||
|
self.use_fp32_lamport = None
|
||||||
self._logged_init = False
|
self._logged_init = False
|
||||||
|
|
||||||
|
|
||||||
@@ -923,6 +932,18 @@ def can_use_flashinfer_allreduce(
|
|||||||
and workspace_manager.dtype == input_.dtype
|
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(
|
return workspace_manager.is_buffer_size_sufficient(
|
||||||
token_num=token_num,
|
token_num=token_num,
|
||||||
hidden_dim=hidden_dim,
|
hidden_dim=hidden_dim,
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import contextlib
|
import contextlib
|
||||||
import types
|
import types
|
||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -16,9 +16,10 @@ register_cuda_ci(est_time=30, stage="base-c", runner_config="4-gpu-gb300")
|
|||||||
|
|
||||||
|
|
||||||
class _FakeWorkspace:
|
class _FakeWorkspace:
|
||||||
def __init__(self, backend, world_size):
|
def __init__(self, backend, world_size, dtype=torch.bfloat16):
|
||||||
self.backend = backend
|
self.backend = backend
|
||||||
self.world_size = world_size
|
self.world_size = world_size
|
||||||
|
self.metadata = {"use_fp32_lamport": dtype == torch.float32}
|
||||||
|
|
||||||
def is_buffer_size_sufficient(self, **_kwargs):
|
def is_buffer_size_sufficient(self, **_kwargs):
|
||||||
return True
|
return True
|
||||||
@@ -34,7 +35,9 @@ class _FakeFlashInferComm:
|
|||||||
|
|
||||||
def create_allreduce_fusion_workspace(self, **kwargs):
|
def create_allreduce_fusion_workspace(self, **kwargs):
|
||||||
self.calls.append(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(
|
def allreduce_fusion(
|
||||||
self,
|
self,
|
||||||
@@ -260,15 +263,17 @@ _OTHER_GROUP_KEY = ("other_device_group", "other_cpu_group")
|
|||||||
|
|
||||||
|
|
||||||
class TestFlashInferAllReduceOnly(CustomTestCase):
|
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 = fusion.FlashInferWorkspaceManager()
|
||||||
manager.workspace = _FakeWorkspace(None, world_size)
|
manager.workspace = _FakeWorkspace(backend, world_size)
|
||||||
manager.initialized = True
|
manager.initialized = True
|
||||||
manager.world_size = world_size
|
manager.world_size = world_size
|
||||||
manager.group = group_key
|
manager.group = group_key
|
||||||
manager.max_token_num = 2048
|
manager.max_token_num = 2048
|
||||||
manager.hidden_dim = 4096
|
manager.hidden_dim = 4096
|
||||||
manager.dtype = torch.float32
|
manager.dtype = torch.float32
|
||||||
|
manager.backend = backend
|
||||||
|
manager.use_fp32_lamport = True
|
||||||
return manager
|
return manager
|
||||||
|
|
||||||
@contextlib.contextmanager
|
@contextlib.contextmanager
|
||||||
@@ -306,7 +311,10 @@ class TestFlashInferAllReduceOnly(CustomTestCase):
|
|||||||
if not torch.cuda.is_available():
|
if not torch.cuda.is_available():
|
||||||
self.skipTest("CUDA required for flashinfer custom op")
|
self.skipTest("CUDA required for flashinfer custom op")
|
||||||
world_size = 4
|
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")
|
input_ = torch.randn(8, 16, dtype=torch.bfloat16, device="cuda")
|
||||||
expected = input_ * world_size
|
expected = input_ * world_size
|
||||||
|
|
||||||
@@ -359,12 +367,94 @@ class TestFlashInferAllReduceOnly(CustomTestCase):
|
|||||||
with self._patched_attn_workspace(self._make_manager(4)):
|
with self._patched_attn_workspace(self._make_manager(4)):
|
||||||
self.assertFalse(self._can_use(torch.randn(8, 16), world_size=2))
|
self.assertFalse(self._can_use(torch.randn(8, 16), world_size=2))
|
||||||
|
|
||||||
def test_rejects_when_token_num_exceeds_workspace_capacity(self):
|
def test_fp32_initialization_caches_allocated_lamport_mode(self):
|
||||||
"""Under Dynamo the capacity check replaces is_buffer_size_sufficient().
|
"""FP32 startup allocation must remain eligible for FP32 all-reduce.
|
||||||
|
|
||||||
_FakeWorkspace.is_buffer_size_sufficient() always says yes, so this only
|
The initialization API's legacy use_fp32_lamport argument defaults to
|
||||||
passes if the compiling branch consults the manager's own allocation.
|
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 = self._make_manager(4)
|
||||||
manager.max_token_num = 8
|
manager.max_token_num = 8
|
||||||
with self._patched_attn_workspace(manager):
|
with self._patched_attn_workspace(manager):
|
||||||
|
|||||||
Reference in New Issue
Block a user