Re-enable SM90 FlashInfer allreduce fusion with safe backend defaults (#28789)

This commit is contained in:
Mohammad Miadh Angkad
2026-06-23 01:29:19 -07:00
committed by GitHub
parent 854c688121
commit 7b1a20344c
3 changed files with 117 additions and 47 deletions
@@ -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()