From c3d829b4505df7a65c7d4772cbc71fedf99825ee Mon Sep 17 00:00:00 2001 From: "ningyunxiao.nyx" Date: Wed, 22 Jul 2026 06:52:38 -0500 Subject: [PATCH 1/5] Add stream-ordered CP gradient return primitive Signed-off-by: ningyunxiao.nyx --- build_tools/pytorch.py | 4 +- transformer_engine/pytorch/csrc/extensions.h | 8 ++ .../pytorch/csrc/extensions/nvshmem_comm.cpp | 111 ++++++++++++++++++ .../pytorch/csrc/extensions/pybind.cpp | 4 + 4 files changed, 126 insertions(+), 1 deletion(-) diff --git a/build_tools/pytorch.py b/build_tools/pytorch.py index 2bb238c522..a25403f323 100644 --- a/build_tools/pytorch.py +++ b/build_tools/pytorch.py @@ -97,7 +97,9 @@ def setup_pytorch_extension( cxx_flags.append("-DUSE_NCCL") library_dirs = [] - libraries = [] + # The CP gradient-return primitive uses stream-ordered CUDA Driver API + # writes and waits for its symmetric epoch protocol. + libraries = ["cuda"] if bool(int(os.getenv("NVTE_ENABLE_NVSHMEM", 0))): assert ( os.getenv("NVSHMEM_HOME") is not None diff --git a/transformer_engine/pytorch/csrc/extensions.h b/transformer_engine/pytorch/csrc/extensions.h index 6edfbdc00e..e2167effb3 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -121,6 +121,14 @@ std::vector fused_attn_bwd( const std::optional cu_seqlens_kv_padded, py::handle s_quantizer, py::handle dp_quantizer, py::handle dqkv_quantizer, bool cuda_graph); +std::vector nvshmem_cp_global_grad_return_execute( + at::Tensor dk_global, at::Tensor dv_global, at::Tensor key, at::Tensor value, + at::Tensor grad_key_return, at::Tensor grad_value_return, + at::Tensor grad_committed_epoch, + const std::vector &peer_grad_key_returns, + const std::vector &peer_grad_value_returns, + const std::vector &peer_grad_committed_epochs, int cp_size, int rank); + at::Tensor fa_prepare_fwd(at::Tensor qkvi); at::Tensor fa_prepare_bwd(at::Tensor q, at::Tensor k, at::Tensor v); diff --git a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp index ac68727ac8..053e9aa95e 100644 --- a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp +++ b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp @@ -5,6 +5,7 @@ ************************************************************************/ #include "../extensions.h" +#include "../../../common/util/cuda_driver.h" #ifdef NVTE_ENABLE_NVSHMEM #include @@ -17,8 +18,63 @@ #include #include +#include +#include +#include +#include + namespace transformer_engine::pytorch { +namespace { + +std::array, 4> cp_global_grad_return_epochs{}; + +at::Tensor cp_grad_return_slot(const at::Tensor &buffer, const at::Tensor &reference, + int cp_size, int writer_rank, const char *name) { + NVTE_CHECK(buffer.defined() && buffer.dim() == reference.dim() + 1, + name, " must have shape [CP, S, B, H, D]."); + NVTE_CHECK(buffer.size(0) == cp_size, name, " leading dimension must equal CP size."); + for (int dim = 0; dim < reference.dim(); ++dim) { + NVTE_CHECK(buffer.size(dim + 1) == reference.size(dim), + name, " trailing dimensions must match local K/V."); + } + return buffer.select(0, writer_rank); +} + +at::Tensor cp_epoch_slot(const at::Tensor &epochs, int writer_rank, const char *name) { + NVTE_CHECK(epochs.is_cuda() && epochs.scalar_type() == torch::kInt32 && + epochs.is_contiguous() && epochs.dim() == 1 && + epochs.size(0) > writer_rank, + name, " must be a contiguous CUDA int32 vector indexed by writer rank."); + return epochs.select(0, writer_rank); +} + +void cp_stream_write_epoch(const at::Tensor &epochs, int writer_rank, int64_t epoch) { + NVTE_CHECK(epoch > 0 && epoch <= static_cast(std::numeric_limits::max()), + "CP gradient epoch must fit in int32."); + at::Tensor slot = cp_epoch_slot(epochs, writer_rank, "peer_grad_committed_epoch"); + cudaStream_t stream = at::cuda::getCurrentCUDAStream(); + NVTE_CHECK_CUDA_DRIVER(cuStreamWriteValue32( + reinterpret_cast(stream), reinterpret_cast(slot.data_ptr()), + static_cast(epoch), 0)); +} + +void cp_stream_wait_epochs(const at::Tensor &epochs, int cp_size, int64_t epoch) { + NVTE_CHECK(epochs.is_cuda() && epochs.scalar_type() == torch::kInt32 && + epochs.is_contiguous() && epochs.dim() == 1 && epochs.size(0) == cp_size, + "grad_committed_epoch must be a contiguous CUDA int32 vector of CP size."); + cudaStream_t stream = at::cuda::getCurrentCUDAStream(); + const auto *base = epochs.data_ptr(); + for (int source = 0; source < cp_size; ++source) { + NVTE_CHECK_CUDA_DRIVER(cuStreamWaitValue32( + reinterpret_cast(stream), + reinterpret_cast(base + source), static_cast(epoch), + CU_STREAM_WAIT_VALUE_GEQ)); + } +} + +} // namespace + void init_nvshmem_backend(c10d::ProcessGroup *process_group) { #ifdef NVTE_ENABLE_NVSHMEM nvshmemx_init_attr_t attr = {}; @@ -119,6 +175,61 @@ void nvshmem_send_on_current_stream(torch::Tensor src, torch::Tensor dst, int pe "distributed process groups when TE is compiled with NVTE_ENABLE_NVSHMEM=1!"); #endif } + +std::vector nvshmem_cp_global_grad_return_execute( + at::Tensor dk_global, at::Tensor dv_global, at::Tensor key, at::Tensor value, + at::Tensor grad_key_return, at::Tensor grad_value_return, + at::Tensor grad_committed_epoch, + const std::vector &peer_grad_key_returns, + const std::vector &peer_grad_value_returns, + const std::vector &peer_grad_committed_epochs, int cp_size, int rank) { + NVTE_CHECK(cp_size == 4 && rank >= 0 && rank < cp_size, + "NVSHMEM global gradient return currently requires CP=4."); + NVTE_CHECK(key.is_cuda() && value.is_cuda() && dk_global.is_cuda() && dv_global.is_cuda(), + "NVSHMEM global gradient return requires CUDA tensors."); + NVTE_CHECK(key.sizes() == value.sizes() && dk_global.sizes() == dv_global.sizes(), + "K/V and global dK/dV pairs must have matching shapes."); + NVTE_CHECK(dk_global.dim() == key.dim() && dk_global.size(0) == key.size(0) * cp_size, + "Global dK/dV sequence length must be CP times local K/V sequence length."); + for (int dim = 1; dim < key.dim(); ++dim) { + NVTE_CHECK(dk_global.size(dim) == key.size(dim), + "Global dK/dV non-sequence dimensions must match local K/V."); + } + NVTE_CHECK(key.size(0) % 2 == 0, "Local K/V sequence length must be even."); + NVTE_CHECK(static_cast(peer_grad_key_returns.size()) == cp_size && + static_cast(peer_grad_value_returns.size()) == cp_size && + static_cast(peer_grad_committed_epochs.size()) == cp_size, + "Expected one symmetric gradient and epoch view per CP owner."); + + cp_grad_return_slot(grad_key_return, key, cp_size, rank, "grad_key_return"); + cp_grad_return_slot(grad_value_return, value, cp_size, rank, "grad_value_return"); + const int64_t half = key.size(0) / 2; + for (int owner = 0; owner < cp_size; ++owner) { + at::Tensor key_slot = cp_grad_return_slot( + peer_grad_key_returns[owner], key, cp_size, rank, "peer_grad_key_return"); + at::Tensor value_slot = cp_grad_return_slot( + peer_grad_value_returns[owner], value, cp_size, rank, "peer_grad_value_return"); + key_slot.narrow(0, 0, half).copy_(dk_global.narrow(0, owner * half, half)); + key_slot.narrow(0, half, half).copy_(dk_global.narrow(0, (7 - owner) * half, half)); + value_slot.narrow(0, 0, half).copy_(dv_global.narrow(0, owner * half, half)); + value_slot.narrow(0, half, half).copy_(dv_global.narrow(0, (7 - owner) * half, half)); + } + + const int64_t epoch = cp_global_grad_return_epochs[rank].fetch_add(1) + 1; + for (int owner = 0; owner < cp_size; ++owner) { + cp_stream_write_epoch(peer_grad_committed_epochs[owner], rank, epoch); + } + cp_stream_wait_epochs(grad_committed_epoch, cp_size, epoch); + + at::Tensor dk = grad_key_return.select(0, 0).clone(); + at::Tensor dv = grad_value_return.select(0, 0).clone(); + for (int source = 1; source < cp_size; ++source) { + dk.add_(grad_key_return.select(0, source)); + dv.add_(grad_value_return.select(0, source)); + } + return {dk.to(key.scalar_type()), dv.to(value.scalar_type())}; +} + void nvshmem_finalize() { #ifdef NVTE_ENABLE_NVSHMEM nvshmem_finalize(); diff --git a/transformer_engine/pytorch/csrc/extensions/pybind.cpp b/transformer_engine/pytorch/csrc/extensions/pybind.cpp index 7e9d114be8..83da1d2693 100644 --- a/transformer_engine/pytorch/csrc/extensions/pybind.cpp +++ b/transformer_engine/pytorch/csrc/extensions/pybind.cpp @@ -597,6 +597,10 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { "Wait for a signal value to be updated by a remote PE using NVSHMEM on the current CUDA " "stream", py::call_guard()); + m.def("nvshmem_cp_global_grad_return_execute", + &transformer_engine::pytorch::nvshmem_cp_global_grad_return_execute, + "Return global CP=4 K/V gradients to their symmetric owner buffers", + py::call_guard()); m.def("nvshmem_finalize", &transformer_engine::pytorch::nvshmem_finalize, "Clean up and finalize the NVSHMEM communication backend and free associated resources", py::call_guard()); From 79a026f727697fe66cdeaac767d5b66632e8bfa4 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 22 Jul 2026 12:43:59 +0000 Subject: [PATCH 2/5] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- transformer_engine/pytorch/csrc/extensions.h | 3 +- .../pytorch/csrc/extensions/nvshmem_comm.cpp | 45 +++++++++---------- 2 files changed, 22 insertions(+), 26 deletions(-) diff --git a/transformer_engine/pytorch/csrc/extensions.h b/transformer_engine/pytorch/csrc/extensions.h index e2167effb3..0aaa954991 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -123,8 +123,7 @@ std::vector fused_attn_bwd( std::vector nvshmem_cp_global_grad_return_execute( at::Tensor dk_global, at::Tensor dv_global, at::Tensor key, at::Tensor value, - at::Tensor grad_key_return, at::Tensor grad_value_return, - at::Tensor grad_committed_epoch, + at::Tensor grad_key_return, at::Tensor grad_value_return, at::Tensor grad_committed_epoch, const std::vector &peer_grad_key_returns, const std::vector &peer_grad_value_returns, const std::vector &peer_grad_committed_epochs, int cp_size, int rank); diff --git a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp index 053e9aa95e..ff3cea8ee4 100644 --- a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp +++ b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp @@ -4,8 +4,8 @@ * See LICENSE for license information. ************************************************************************/ -#include "../extensions.h" #include "../../../common/util/cuda_driver.h" +#include "../extensions.h" #ifdef NVTE_ENABLE_NVSHMEM #include @@ -29,22 +29,21 @@ namespace { std::array, 4> cp_global_grad_return_epochs{}; -at::Tensor cp_grad_return_slot(const at::Tensor &buffer, const at::Tensor &reference, - int cp_size, int writer_rank, const char *name) { - NVTE_CHECK(buffer.defined() && buffer.dim() == reference.dim() + 1, - name, " must have shape [CP, S, B, H, D]."); +at::Tensor cp_grad_return_slot(const at::Tensor &buffer, const at::Tensor &reference, int cp_size, + int writer_rank, const char *name) { + NVTE_CHECK(buffer.defined() && buffer.dim() == reference.dim() + 1, name, + " must have shape [CP, S, B, H, D]."); NVTE_CHECK(buffer.size(0) == cp_size, name, " leading dimension must equal CP size."); for (int dim = 0; dim < reference.dim(); ++dim) { - NVTE_CHECK(buffer.size(dim + 1) == reference.size(dim), - name, " trailing dimensions must match local K/V."); + NVTE_CHECK(buffer.size(dim + 1) == reference.size(dim), name, + " trailing dimensions must match local K/V."); } return buffer.select(0, writer_rank); } at::Tensor cp_epoch_slot(const at::Tensor &epochs, int writer_rank, const char *name) { - NVTE_CHECK(epochs.is_cuda() && epochs.scalar_type() == torch::kInt32 && - epochs.is_contiguous() && epochs.dim() == 1 && - epochs.size(0) > writer_rank, + NVTE_CHECK(epochs.is_cuda() && epochs.scalar_type() == torch::kInt32 && epochs.is_contiguous() && + epochs.dim() == 1 && epochs.size(0) > writer_rank, name, " must be a contiguous CUDA int32 vector indexed by writer rank."); return epochs.select(0, writer_rank); } @@ -54,22 +53,21 @@ void cp_stream_write_epoch(const at::Tensor &epochs, int writer_rank, int64_t ep "CP gradient epoch must fit in int32."); at::Tensor slot = cp_epoch_slot(epochs, writer_rank, "peer_grad_committed_epoch"); cudaStream_t stream = at::cuda::getCurrentCUDAStream(); - NVTE_CHECK_CUDA_DRIVER(cuStreamWriteValue32( - reinterpret_cast(stream), reinterpret_cast(slot.data_ptr()), - static_cast(epoch), 0)); + NVTE_CHECK_CUDA_DRIVER(cuStreamWriteValue32(reinterpret_cast(stream), + reinterpret_cast(slot.data_ptr()), + static_cast(epoch), 0)); } void cp_stream_wait_epochs(const at::Tensor &epochs, int cp_size, int64_t epoch) { - NVTE_CHECK(epochs.is_cuda() && epochs.scalar_type() == torch::kInt32 && - epochs.is_contiguous() && epochs.dim() == 1 && epochs.size(0) == cp_size, + NVTE_CHECK(epochs.is_cuda() && epochs.scalar_type() == torch::kInt32 && epochs.is_contiguous() && + epochs.dim() == 1 && epochs.size(0) == cp_size, "grad_committed_epoch must be a contiguous CUDA int32 vector of CP size."); cudaStream_t stream = at::cuda::getCurrentCUDAStream(); const auto *base = epochs.data_ptr(); for (int source = 0; source < cp_size; ++source) { NVTE_CHECK_CUDA_DRIVER(cuStreamWaitValue32( - reinterpret_cast(stream), - reinterpret_cast(base + source), static_cast(epoch), - CU_STREAM_WAIT_VALUE_GEQ)); + reinterpret_cast(stream), reinterpret_cast(base + source), + static_cast(epoch), CU_STREAM_WAIT_VALUE_GEQ)); } } @@ -178,8 +176,7 @@ void nvshmem_send_on_current_stream(torch::Tensor src, torch::Tensor dst, int pe std::vector nvshmem_cp_global_grad_return_execute( at::Tensor dk_global, at::Tensor dv_global, at::Tensor key, at::Tensor value, - at::Tensor grad_key_return, at::Tensor grad_value_return, - at::Tensor grad_committed_epoch, + at::Tensor grad_key_return, at::Tensor grad_value_return, at::Tensor grad_committed_epoch, const std::vector &peer_grad_key_returns, const std::vector &peer_grad_value_returns, const std::vector &peer_grad_committed_epochs, int cp_size, int rank) { @@ -205,10 +202,10 @@ std::vector nvshmem_cp_global_grad_return_execute( cp_grad_return_slot(grad_value_return, value, cp_size, rank, "grad_value_return"); const int64_t half = key.size(0) / 2; for (int owner = 0; owner < cp_size; ++owner) { - at::Tensor key_slot = cp_grad_return_slot( - peer_grad_key_returns[owner], key, cp_size, rank, "peer_grad_key_return"); - at::Tensor value_slot = cp_grad_return_slot( - peer_grad_value_returns[owner], value, cp_size, rank, "peer_grad_value_return"); + at::Tensor key_slot = cp_grad_return_slot(peer_grad_key_returns[owner], key, cp_size, rank, + "peer_grad_key_return"); + at::Tensor value_slot = cp_grad_return_slot(peer_grad_value_returns[owner], value, cp_size, + rank, "peer_grad_value_return"); key_slot.narrow(0, 0, half).copy_(dk_global.narrow(0, owner * half, half)); key_slot.narrow(0, half, half).copy_(dk_global.narrow(0, (7 - owner) * half, half)); value_slot.narrow(0, 0, half).copy_(dv_global.narrow(0, owner * half, half)); From f676f681a83304c0e8d08e1c57d3faa8fd77229d Mon Sep 17 00:00:00 2001 From: "ningyunxiao.nyx" Date: Wed, 22 Jul 2026 19:18:23 -0500 Subject: [PATCH 3/5] Use lazy CUDA driver calls for CP gradient return Signed-off-by: ningyunxiao.nyx --- build_tools/pytorch.py | 4 +--- .../pytorch/csrc/extensions/nvshmem_comm.cpp | 19 +++++++++++-------- 2 files changed, 12 insertions(+), 11 deletions(-) diff --git a/build_tools/pytorch.py b/build_tools/pytorch.py index a25403f323..2bb238c522 100644 --- a/build_tools/pytorch.py +++ b/build_tools/pytorch.py @@ -97,9 +97,7 @@ def setup_pytorch_extension( cxx_flags.append("-DUSE_NCCL") library_dirs = [] - # The CP gradient-return primitive uses stream-ordered CUDA Driver API - # writes and waits for its symmetric epoch protocol. - libraries = ["cuda"] + libraries = [] if bool(int(os.getenv("NVTE_ENABLE_NVSHMEM", 0))): assert ( os.getenv("NVSHMEM_HOME") is not None diff --git a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp index ff3cea8ee4..6241b6d18a 100644 --- a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp +++ b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp @@ -53,9 +53,9 @@ void cp_stream_write_epoch(const at::Tensor &epochs, int writer_rank, int64_t ep "CP gradient epoch must fit in int32."); at::Tensor slot = cp_epoch_slot(epochs, writer_rank, "peer_grad_committed_epoch"); cudaStream_t stream = at::cuda::getCurrentCUDAStream(); - NVTE_CHECK_CUDA_DRIVER(cuStreamWriteValue32(reinterpret_cast(stream), - reinterpret_cast(slot.data_ptr()), - static_cast(epoch), 0)); + NVTE_CALL_CHECK_CUDA_DRIVER(cuStreamWriteValue32, reinterpret_cast(stream), + reinterpret_cast(slot.data_ptr()), + static_cast(epoch), 0); } void cp_stream_wait_epochs(const at::Tensor &epochs, int cp_size, int64_t epoch) { @@ -65,9 +65,9 @@ void cp_stream_wait_epochs(const at::Tensor &epochs, int cp_size, int64_t epoch) cudaStream_t stream = at::cuda::getCurrentCUDAStream(); const auto *base = epochs.data_ptr(); for (int source = 0; source < cp_size; ++source) { - NVTE_CHECK_CUDA_DRIVER(cuStreamWaitValue32( - reinterpret_cast(stream), reinterpret_cast(base + source), - static_cast(epoch), CU_STREAM_WAIT_VALUE_GEQ)); + NVTE_CALL_CHECK_CUDA_DRIVER(cuStreamWaitValue32, reinterpret_cast(stream), + reinterpret_cast(base + source), + static_cast(epoch), CU_STREAM_WAIT_VALUE_GEQ); } } @@ -202,14 +202,17 @@ std::vector nvshmem_cp_global_grad_return_execute( cp_grad_return_slot(grad_value_return, value, cp_size, rank, "grad_value_return"); const int64_t half = key.size(0) / 2; for (int owner = 0; owner < cp_size; ++owner) { + // Native two-chunk CP layout: owner o owns half-chunks o and + // (2 * cp_size - 1 - o). For CP=4, the mirror index is 7 - o. + const int64_t mirror_half = 2 * static_cast(cp_size) - 1 - owner; at::Tensor key_slot = cp_grad_return_slot(peer_grad_key_returns[owner], key, cp_size, rank, "peer_grad_key_return"); at::Tensor value_slot = cp_grad_return_slot(peer_grad_value_returns[owner], value, cp_size, rank, "peer_grad_value_return"); key_slot.narrow(0, 0, half).copy_(dk_global.narrow(0, owner * half, half)); - key_slot.narrow(0, half, half).copy_(dk_global.narrow(0, (7 - owner) * half, half)); + key_slot.narrow(0, half, half).copy_(dk_global.narrow(0, mirror_half * half, half)); value_slot.narrow(0, 0, half).copy_(dv_global.narrow(0, owner * half, half)); - value_slot.narrow(0, half, half).copy_(dv_global.narrow(0, (7 - owner) * half, half)); + value_slot.narrow(0, half, half).copy_(dv_global.narrow(0, mirror_half * half, half)); } const int64_t epoch = cp_global_grad_return_epochs[rank].fetch_add(1) + 1; From 63ee452faff1640ca12037941a12c29ebac0fb73 Mon Sep 17 00:00:00 2001 From: "ningyunxiao.nyx" Date: Wed, 22 Jul 2026 21:16:48 -0500 Subject: [PATCH 4/5] Validate CP gradient return ownership and epochs Signed-off-by: ningyunxiao.nyx --- .../pytorch/csrc/extensions/nvshmem_comm.cpp | 33 ++++++++++++++++--- 1 file changed, 29 insertions(+), 4 deletions(-) diff --git a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp index 6241b6d18a..1fe9603bfb 100644 --- a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp +++ b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp @@ -18,16 +18,28 @@ #include #include -#include -#include #include #include +#include +#include namespace transformer_engine::pytorch { namespace { -std::array, 4> cp_global_grad_return_epochs{}; +std::mutex cp_global_grad_return_epoch_mutex; +std::unordered_map cp_global_grad_return_epochs; + +int64_t cp_next_global_grad_return_epoch(const at::Tensor &epochs) { + // Scope each epoch sequence to its symmetric buffer so independent CP + // backend instances in one process cannot satisfy each other's waits. + const void *buffer = epochs.data_ptr(); + std::lock_guard lock(cp_global_grad_return_epoch_mutex); + int64_t &epoch = cp_global_grad_return_epochs[buffer]; + NVTE_CHECK(epoch < static_cast(std::numeric_limits::max()), + "CP gradient epoch exhausted for this symmetric epoch buffer."); + return ++epoch; +} at::Tensor cp_grad_return_slot(const at::Tensor &buffer, const at::Tensor &reference, int cp_size, int writer_rank, const char *name) { @@ -200,6 +212,19 @@ std::vector nvshmem_cp_global_grad_return_execute( cp_grad_return_slot(grad_key_return, key, cp_size, rank, "grad_key_return"); cp_grad_return_slot(grad_value_return, value, cp_size, rank, "grad_value_return"); + cp_grad_return_slot(peer_grad_key_returns[rank], key, cp_size, rank, + "peer_grad_key_returns[rank]"); + cp_grad_return_slot(peer_grad_value_returns[rank], value, cp_size, rank, + "peer_grad_value_returns[rank]"); + cp_epoch_slot(grad_committed_epoch, rank, "grad_committed_epoch"); + cp_epoch_slot(peer_grad_committed_epochs[rank], rank, "peer_grad_committed_epochs[rank]"); + NVTE_CHECK(peer_grad_key_returns[rank].data_ptr() == grad_key_return.data_ptr(), + "peer_grad_key_returns[rank] must alias grad_key_return."); + NVTE_CHECK(peer_grad_value_returns[rank].data_ptr() == grad_value_return.data_ptr(), + "peer_grad_value_returns[rank] must alias grad_value_return."); + NVTE_CHECK(peer_grad_committed_epochs[rank].data_ptr() == grad_committed_epoch.data_ptr(), + "peer_grad_committed_epochs[rank] must alias grad_committed_epoch."); + const int64_t half = key.size(0) / 2; for (int owner = 0; owner < cp_size; ++owner) { // Native two-chunk CP layout: owner o owns half-chunks o and @@ -215,7 +240,7 @@ std::vector nvshmem_cp_global_grad_return_execute( value_slot.narrow(0, half, half).copy_(dv_global.narrow(0, mirror_half * half, half)); } - const int64_t epoch = cp_global_grad_return_epochs[rank].fetch_add(1) + 1; + const int64_t epoch = cp_next_global_grad_return_epoch(grad_committed_epoch); for (int owner = 0; owner < cp_size; ++owner) { cp_stream_write_epoch(peer_grad_committed_epochs[owner], rank, epoch); } From 059f2f10da55b613dda00a5b236c119b324a1cef Mon Sep 17 00:00:00 2001 From: "ningyunxiao.nyx" Date: Wed, 22 Jul 2026 22:29:51 -0500 Subject: [PATCH 5/5] Validate CP gradient return tensor devices Signed-off-by: ningyunxiao.nyx --- .../pytorch/csrc/extensions/nvshmem_comm.cpp | 22 ++++++++++++++----- 1 file changed, 16 insertions(+), 6 deletions(-) diff --git a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp index 1fe9603bfb..9059f334c8 100644 --- a/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp +++ b/transformer_engine/pytorch/csrc/extensions/nvshmem_comm.cpp @@ -41,10 +41,17 @@ int64_t cp_next_global_grad_return_epoch(const at::Tensor &epochs) { return ++epoch; } +void cp_check_current_device(const at::Tensor &tensor, const char *name) { + NVTE_CHECK(tensor.defined() && tensor.is_cuda(), name, " must be a CUDA tensor."); + const int current_device = c10::cuda::current_device(); + NVTE_CHECK(tensor.get_device() == current_device, name, " must be mapped on current CUDA device ", + current_device, "."); +} + at::Tensor cp_grad_return_slot(const at::Tensor &buffer, const at::Tensor &reference, int cp_size, int writer_rank, const char *name) { - NVTE_CHECK(buffer.defined() && buffer.dim() == reference.dim() + 1, name, - " must have shape [CP, S, B, H, D]."); + cp_check_current_device(buffer, name); + NVTE_CHECK(buffer.dim() == reference.dim() + 1, name, " must have shape [CP, S, B, H, D]."); NVTE_CHECK(buffer.size(0) == cp_size, name, " leading dimension must equal CP size."); for (int dim = 0; dim < reference.dim(); ++dim) { NVTE_CHECK(buffer.size(dim + 1) == reference.size(dim), name, @@ -54,8 +61,9 @@ at::Tensor cp_grad_return_slot(const at::Tensor &buffer, const at::Tensor &refer } at::Tensor cp_epoch_slot(const at::Tensor &epochs, int writer_rank, const char *name) { - NVTE_CHECK(epochs.is_cuda() && epochs.scalar_type() == torch::kInt32 && epochs.is_contiguous() && - epochs.dim() == 1 && epochs.size(0) > writer_rank, + cp_check_current_device(epochs, name); + NVTE_CHECK(epochs.scalar_type() == torch::kInt32 && epochs.is_contiguous() && epochs.dim() == 1 && + epochs.size(0) > writer_rank, name, " must be a contiguous CUDA int32 vector indexed by writer rank."); return epochs.select(0, writer_rank); } @@ -194,8 +202,10 @@ std::vector nvshmem_cp_global_grad_return_execute( const std::vector &peer_grad_committed_epochs, int cp_size, int rank) { NVTE_CHECK(cp_size == 4 && rank >= 0 && rank < cp_size, "NVSHMEM global gradient return currently requires CP=4."); - NVTE_CHECK(key.is_cuda() && value.is_cuda() && dk_global.is_cuda() && dv_global.is_cuda(), - "NVSHMEM global gradient return requires CUDA tensors."); + cp_check_current_device(key, "key"); + cp_check_current_device(value, "value"); + cp_check_current_device(dk_global, "dk_global"); + cp_check_current_device(dv_global, "dv_global"); NVTE_CHECK(key.sizes() == value.sizes() && dk_global.sizes() == dv_global.sizes(), "K/V and global dK/dV pairs must have matching shapes."); NVTE_CHECK(dk_global.dim() == key.dim() && dk_global.size(0) == key.size(0) * cp_size,