Re-enable SM90 FlashInfer allreduce fusion with safe backend defaults (#28789)
This commit is contained in:
@@ -80,36 +80,95 @@ class TestFlashInferCommFusion(unittest.TestCase):
|
||||
flashinfer_allreduce_fusion_backend="auto", nnodes=2
|
||||
)
|
||||
|
||||
# Blackwell: mnnvl regardless of node count.
|
||||
# Blackwell: trtllm on single-node, mnnvl on multi-node.
|
||||
with patch.object(fusion, "is_sm100_supported", return_value=True):
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node), "mnnvl"
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node),
|
||||
"trtllm",
|
||||
)
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node), "mnnvl"
|
||||
)
|
||||
|
||||
# SM90: mnnvl on single-node, trtllm fallback on multi-node.
|
||||
# SM90: auto uses trtllm on single-node, multi-node is unsupported.
|
||||
with (
|
||||
patch.object(fusion, "is_sm100_supported", return_value=False),
|
||||
patch.object(fusion, "is_sm90_supported", return_value=True),
|
||||
):
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node), "mnnvl"
|
||||
)
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node), "trtllm"
|
||||
)
|
||||
|
||||
# Pre-SM90: trtllm everywhere.
|
||||
with (
|
||||
patch.object(fusion, "is_sm100_supported", return_value=False),
|
||||
patch.object(fusion, "is_sm90_supported", return_value=False),
|
||||
):
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node),
|
||||
"trtllm",
|
||||
)
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node)
|
||||
|
||||
# Architectures outside SM90/SM10X are unsupported. Both pre-SM90
|
||||
# and post-SM10X devices (e.g. SM120) must fail closed.
|
||||
for arch in ("pre_sm90", "post_sm10x"):
|
||||
with (
|
||||
self.subTest(arch=arch),
|
||||
patch.object(fusion, "is_sm100_supported", return_value=False),
|
||||
patch.object(fusion, "is_sm90_supported", return_value=False),
|
||||
):
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node)
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node)
|
||||
|
||||
def test_explicit_backend_validation(self):
|
||||
single_node_mnnvl = types.SimpleNamespace(
|
||||
flashinfer_allreduce_fusion_backend="mnnvl", nnodes=1
|
||||
)
|
||||
multi_node_mnnvl = types.SimpleNamespace(
|
||||
flashinfer_allreduce_fusion_backend="mnnvl", nnodes=2
|
||||
)
|
||||
single_node_trtllm = types.SimpleNamespace(
|
||||
flashinfer_allreduce_fusion_backend="trtllm", nnodes=1
|
||||
)
|
||||
multi_node_trtllm = types.SimpleNamespace(
|
||||
flashinfer_allreduce_fusion_backend="trtllm", nnodes=2
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(fusion, "is_sm100_supported", return_value=False),
|
||||
patch.object(fusion, "is_sm90_supported", return_value=True),
|
||||
):
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node_mnnvl),
|
||||
"mnnvl",
|
||||
)
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(single_node_trtllm),
|
||||
"trtllm",
|
||||
)
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_mnnvl)
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_trtllm)
|
||||
|
||||
with patch.object(fusion, "is_sm100_supported", return_value=True):
|
||||
self.assertEqual(
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_mnnvl),
|
||||
"mnnvl",
|
||||
)
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(multi_node_trtllm)
|
||||
|
||||
for arch in ("pre_sm90", "post_sm10x"):
|
||||
with (
|
||||
self.subTest(arch=arch),
|
||||
patch.object(fusion, "is_sm100_supported", return_value=False),
|
||||
patch.object(fusion, "is_sm90_supported", return_value=False),
|
||||
):
|
||||
for args in (
|
||||
single_node_mnnvl,
|
||||
multi_node_mnnvl,
|
||||
single_node_trtllm,
|
||||
multi_node_trtllm,
|
||||
):
|
||||
with self.subTest(backend=args.flashinfer_allreduce_fusion_backend):
|
||||
with self.assertRaises(ValueError):
|
||||
fusion.resolve_flashinfer_allreduce_fusion_backend(args)
|
||||
|
||||
def test_allreduce_fusion_backends_match_torch_baseline(self):
|
||||
fake_comm = _FakeFlashInferComm()
|
||||
|
||||
Reference in New Issue
Block a user