From a3a48113bd3d546869e592b16d30649fc47209b6 Mon Sep 17 00:00:00 2001 From: lihongyang1990 Date: Fri, 7 Aug 2026 16:03:24 +0800 Subject: [PATCH] fix(plugin): add TP comm overlap backend ops --- .../plugin/core/backends/vendor/cuda/cuda.py | 12 +++++++++ .../core/backends/vendor/cuda/register_ops.py | 25 +++++++++++++++++++ .../core/backends/vendor/hygon/hygon.py | 12 +++++++++ .../backends/vendor/hygon/register_ops.py | 24 ++++++++++++++++++ 4 files changed, 73 insertions(+) diff --git a/transformer_engine/plugin/core/backends/vendor/cuda/cuda.py b/transformer_engine/plugin/core/backends/vendor/cuda/cuda.py index 3be294fe57..28c22fa94b 100644 --- a/transformer_engine/plugin/core/backends/vendor/cuda/cuda.py +++ b/transformer_engine/plugin/core/backends/vendor/cuda/cuda.py @@ -1992,6 +1992,7 @@ def create_comm_overlap_p2p( aggregate: bool = False, ) -> "CommOverlapP2P": tex = self._get_tex() + comm_type = tex.CommOverlapType(int(comm_type)) if comm_type is not None else None return tex.CommOverlapP2P( buffer_shape, buffer_dtype, @@ -2008,3 +2009,14 @@ def create_comm_overlap_p2p( use_ce, aggregate, ) + + def device_supports_multicast(self, device_id=-1): + return self._get_tex().device_supports_multicast(device_id) + + def ubuf_built_with_mpi(self): + tex = self._get_tex() + return tex.ubuf_built_with_mpi() + + def get_stream_priority_range(self, device_id=-1): + tex = self._get_tex() + return tex.get_stream_priority_range(device_id) diff --git a/transformer_engine/plugin/core/backends/vendor/cuda/register_ops.py b/transformer_engine/plugin/core/backends/vendor/cuda/register_ops.py index 5fac3e34c4..1b891d6a99 100644 --- a/transformer_engine/plugin/core/backends/vendor/cuda/register_ops.py +++ b/transformer_engine/plugin/core/backends/vendor/cuda/register_ops.py @@ -1150,6 +1150,31 @@ def register_builtins(registry) -> None: vendor="CUDA", priority=100, ), + OpImpl( + op_name="device_supports_multicast", + impl_id="vendor.cuda", + kind=BackendImplKind.VENDOR, + fn=_bind_is_available(backend.device_supports_multicast, is_avail), + vendor="CUDA", + priority=100, + ), + + OpImpl( + op_name="ubuf_built_with_mpi", + impl_id="vendor.cuda", + kind=BackendImplKind.VENDOR, + fn=_bind_is_available(backend.ubuf_built_with_mpi, is_avail), + vendor="CUDA", + priority=100, + ), + OpImpl( + op_name="get_stream_priority_range", + impl_id="vendor.cuda", + kind=BackendImplKind.VENDOR, + fn=_bind_is_available(backend.get_stream_priority_range, is_avail), + vendor="CUDA", + priority=100, + ), # FlashAttention class getter OpImpl( op_name="get_flash_attention_class", diff --git a/transformer_engine/plugin/core/backends/vendor/hygon/hygon.py b/transformer_engine/plugin/core/backends/vendor/hygon/hygon.py index 69ca8608ed..f0e87d14a0 100644 --- a/transformer_engine/plugin/core/backends/vendor/hygon/hygon.py +++ b/transformer_engine/plugin/core/backends/vendor/hygon/hygon.py @@ -1954,6 +1954,7 @@ def create_comm_overlap_p2p( aggregate: bool = False, ) -> "CommOverlapP2P": tex = self._get_tex() + comm_type = tex.CommOverlapType(int(comm_type)) if comm_type is not None else None return tex.CommOverlapP2P( buffer_shape, buffer_dtype, @@ -1970,3 +1971,14 @@ def create_comm_overlap_p2p( use_ce, aggregate, ) + + def device_supports_multicast(self, device_id=-1): + return self._get_tex().device_supports_multicast(device_id) + + def ubuf_built_with_mpi(self): + tex = self._get_tex() + return tex.ubuf_built_with_mpi() + + def get_stream_priority_range(self, device_id=-1): + tex = self._get_tex() + return tex.get_stream_priority_range(device_id) diff --git a/transformer_engine/plugin/core/backends/vendor/hygon/register_ops.py b/transformer_engine/plugin/core/backends/vendor/hygon/register_ops.py index 2b0bbc8aa0..c7e732f8fb 100644 --- a/transformer_engine/plugin/core/backends/vendor/hygon/register_ops.py +++ b/transformer_engine/plugin/core/backends/vendor/hygon/register_ops.py @@ -1084,6 +1084,30 @@ def register_builtins(registry) -> None: vendor="HYGON", priority=100, ), + OpImpl( + op_name="device_supports_multicast", + impl_id="vendor.hygon", + kind=BackendImplKind.VENDOR, + fn=_bind_is_available(backend.device_supports_multicast, is_avail), + vendor="HYGON", + priority=100, + ), + OpImpl( + op_name="ubuf_built_with_mpi", + impl_id="vendor.hygon", + kind=BackendImplKind.VENDOR, + fn=_bind_is_available(backend.ubuf_built_with_mpi, is_avail), + vendor="HYGON", + priority=100, + ), + OpImpl( + op_name="get_stream_priority_range", + impl_id="vendor.hygon", + kind=BackendImplKind.VENDOR, + fn=_bind_is_available(backend.get_stream_priority_range, is_avail), + vendor="HYGON", + priority=100, + ), # FlashAttention class getter OpImpl( op_name="get_flash_attention_class",