From 3875204c5eb8dc64da9ce0f483c166947683f2c7 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 11:44:57 +0200 Subject: [PATCH 01/15] Add CI workflow building TE against torch 2.1 Signed-off-by: Pawel Gadzinski --- .github/workflows/build_pytorch21.yml | 43 +++++++++++++++++++++++++++ 1 file changed, 43 insertions(+) create mode 100644 .github/workflows/build_pytorch21.yml diff --git a/.github/workflows/build_pytorch21.yml b/.github/workflows/build_pytorch21.yml new file mode 100644 index 0000000000..0f4aea1f0d --- /dev/null +++ b/.github/workflows/build_pytorch21.yml @@ -0,0 +1,43 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# Verify that the torch>=2.1 requirement declared in build_tools/pytorch.py +# actually holds: build against torch==2.1.2 and run the sanity import. +name: 'Build PyTorch 2.1' +on: + pull_request: + workflow_dispatch: +concurrency: + group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} + cancel-in-progress: true +jobs: + pytorch21: + name: 'PyTorch 2.1' + runs-on: ubuntu-latest + container: + image: nvcr.io/nvidia/cuda:12.8.0-devel-ubuntu22.04 + options: --user root + steps: + - name: 'Dependencies' + run: | + apt-get update + apt-get install -y git python3.9 pip cudnn9-cuda-12 + pip install torch==2.1.2 "numpy<2" + pip install cmake pybind11[global] ninja pydantic importlib-metadata>=1.0 packaging einops onnxscript "nvidia-cudnn-frontend>=1.25.0" + - name: 'Checkout' + uses: actions/checkout@v3 + with: + submodules: recursive + - name: ccache + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad + - name: 'Build' + run: NVTE_USE_CCACHE=1 NVTE_CCACHE_BIN=sccache pip install --no-build-isolation --no-deps . -v + env: + NVTE_FRAMEWORK: pytorch + # Single old arch (Volta) to keep memory/build size down + NVTE_CUDA_ARCHS: "70" + MAX_JOBS: 1 + SCCACHE_GHA_ENABLED: "true" + - name: 'Sanity check' + run: python3 tests/pytorch/test_sanity_import.py From fba3898529994d756ba4ac564203c370c24db6bb Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 12:14:50 +0200 Subject: [PATCH 02/15] Remove MAX_JOBS=1 from pytorch 2.1 build workflow Signed-off-by: Pawel Gadzinski --- .github/workflows/build_pytorch21.yml | 1 - 1 file changed, 1 deletion(-) diff --git a/.github/workflows/build_pytorch21.yml b/.github/workflows/build_pytorch21.yml index 0f4aea1f0d..462bbe6a20 100644 --- a/.github/workflows/build_pytorch21.yml +++ b/.github/workflows/build_pytorch21.yml @@ -37,7 +37,6 @@ jobs: NVTE_FRAMEWORK: pytorch # Single old arch (Volta) to keep memory/build size down NVTE_CUDA_ARCHS: "70" - MAX_JOBS: 1 SCCACHE_GHA_ENABLED: "true" - name: 'Sanity check' run: python3 tests/pytorch/test_sanity_import.py From 3b600ddd0b3fea3ef869aa056bcb0d3d19a8fe53 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 12:35:42 +0200 Subject: [PATCH 03/15] Use MAX_JOBS=2 in pytorch 2.1 workflow, unbounded build OOMs the runner Signed-off-by: Pawel Gadzinski --- .github/workflows/build_pytorch21.yml | 2 ++ 1 file changed, 2 insertions(+) diff --git a/.github/workflows/build_pytorch21.yml b/.github/workflows/build_pytorch21.yml index 462bbe6a20..b351849932 100644 --- a/.github/workflows/build_pytorch21.yml +++ b/.github/workflows/build_pytorch21.yml @@ -37,6 +37,8 @@ jobs: NVTE_FRAMEWORK: pytorch # Single old arch (Volta) to keep memory/build size down NVTE_CUDA_ARCHS: "70" + # Full parallelism OOMs the 7GB runner (exit 137) + MAX_JOBS: 2 SCCACHE_GHA_ENABLED: "true" - name: 'Sanity check' run: python3 tests/pytorch/test_sanity_import.py From dce613cf8070f6f1e8c8621cd3ba8c99a9e5450f Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 13:30:27 +0200 Subject: [PATCH 04/15] Point TE build at apt cudnn9, torch 2.1 pip deps shadow it with cudnn 8.9 Signed-off-by: Pawel Gadzinski --- .github/workflows/build_pytorch21.yml | 3 +++ 1 file changed, 3 insertions(+) diff --git a/.github/workflows/build_pytorch21.yml b/.github/workflows/build_pytorch21.yml index b351849932..f577f148da 100644 --- a/.github/workflows/build_pytorch21.yml +++ b/.github/workflows/build_pytorch21.yml @@ -37,6 +37,9 @@ jobs: NVTE_FRAMEWORK: pytorch # Single old arch (Volta) to keep memory/build size down NVTE_CUDA_ARCHS: "70" + # torch 2.1 pulls in pip cudnn 8.9; point the TE build at apt cudnn 9 + CUDNN_INCLUDE_PATH: /usr/include/x86_64-linux-gnu + CUDNN_LIBRARY_PATH: /usr/lib/x86_64-linux-gnu # Full parallelism OOMs the 7GB runner (exit 137) MAX_JOBS: 2 SCCACHE_GHA_ENABLED: "true" From 061834fc3484755a92ca82a553be8bae86fdba44 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 14:15:11 +0200 Subject: [PATCH 05/15] Avoid passing std::optional to at::get_generator_or_default, torch 2.1 needs c10::optional Signed-off-by: Pawel Gadzinski --- transformer_engine/pytorch/csrc/extensions/attention.cpp | 4 ++-- transformer_engine/pytorch/csrc/extensions/cast.cpp | 8 ++++---- transformer_engine/pytorch/csrc/extensions/dropout.cpp | 4 ++-- transformer_engine/pytorch/csrc/quantizer.cpp | 4 ++-- 4 files changed, 10 insertions(+), 10 deletions(-) diff --git a/transformer_engine/pytorch/csrc/extensions/attention.cpp b/transformer_engine/pytorch/csrc/extensions/attention.cpp index eb8813d4a0..0b3b882768 100644 --- a/transformer_engine/pytorch/csrc/extensions/attention.cpp +++ b/transformer_engine/pytorch/csrc/extensions/attention.cpp @@ -228,8 +228,8 @@ std::vector fused_attn_fwd( } // extract rng seed and offset - auto gen = at::get_generator_or_default( - rng_gen, at::cuda::detail::getDefaultCUDAGenerator()); + auto gen = at::check_generator( + rng_gen.has_value() ? *rng_gen : at::cuda::detail::getDefaultCUDAGenerator()); at::PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); auto options = torch::TensorOptions().dtype(torch::kInt64).device(torch::kCUDA); auto rng_state = torch::empty({2}, options); diff --git a/transformer_engine/pytorch/csrc/extensions/cast.cpp b/transformer_engine/pytorch/csrc/extensions/cast.cpp index 50cc5e1bd4..9f6a8d627e 100644 --- a/transformer_engine/pytorch/csrc/extensions/cast.cpp +++ b/transformer_engine/pytorch/csrc/extensions/cast.cpp @@ -173,8 +173,8 @@ void group_quantize_nvfp4_impl(const GroupedTensorWrapper &grouped_input_tensor, // number for different tensors in the group, so we only need to allocate one rng state const size_t rng_elts_per_thread = 1024 * num_tensors; rng_states_tensor = torch::empty({2}, opts); - auto gen = at::get_generator_or_default( - std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); + auto gen = + at::check_generator(at::cuda::detail::getDefaultCUDAGenerator()); at::PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); philox_unpack(philox_args, static_cast(rng_states_tensor.data_ptr())); @@ -1390,8 +1390,8 @@ static StochasticRngStateResources setup_stochastic_rounding_rng_states_helper( if (need_separate_rng_states) res.te_rng_state_list_colwise.reserve(num_tensors); for (size_t i = 0; i < num_tensors; ++i) { - auto gen = at::get_generator_or_default( - std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); + auto gen = + at::check_generator(at::cuda::detail::getDefaultCUDAGenerator()); // Rowwise RNG state at::PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); diff --git a/transformer_engine/pytorch/csrc/extensions/dropout.cpp b/transformer_engine/pytorch/csrc/extensions/dropout.cpp index bea8f3a7b5..0bf44eebd0 100644 --- a/transformer_engine/pytorch/csrc/extensions/dropout.cpp +++ b/transformer_engine/pytorch/csrc/extensions/dropout.cpp @@ -44,8 +44,8 @@ std::vector dropout_fwd(const py::handle &input, float dropout_proba auto mask_nvte = makeTransformerEngineTensor(mask_pyt); // RNG state tensor - auto gen = at::get_generator_or_default( - std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); + auto gen = + at::check_generator(at::cuda::detail::getDefaultCUDAGenerator()); at::PhiloxCudaState philox_args; { std::lock_guard lock(gen->mutex_); diff --git a/transformer_engine/pytorch/csrc/quantizer.cpp b/transformer_engine/pytorch/csrc/quantizer.cpp index 4704261790..d2ca6b635a 100644 --- a/transformer_engine/pytorch/csrc/quantizer.cpp +++ b/transformer_engine/pytorch/csrc/quantizer.cpp @@ -2603,8 +2603,8 @@ void NVFP4Quantizer::quantize_impl(const TensorWrapper& input, TensorWrapper& ou if (this->stochastic_rounding) { const size_t rng_elts_per_thread = 1024; // Wild guess, probably can be tightened - auto gen = at::get_generator_or_default( - std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); + auto gen = + at::check_generator(at::cuda::detail::getDefaultCUDAGenerator()); auto opts = at::TensorOptions().dtype(torch::kInt64).device(torch::kCUDA); // Generate RNG state for rowwise quantization From 32a28ef33390050a6cdcafcb54e2146baa1412a3 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 14:26:30 +0200 Subject: [PATCH 06/15] Gate NCCL EP in torch extension on nccl_ep_enabled, matching common CMake Signed-off-by: Pawel Gadzinski --- build_tools/pytorch.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/build_tools/pytorch.py b/build_tools/pytorch.py index 2bb238c522..98331ccbd8 100644 --- a/build_tools/pytorch.py +++ b/build_tools/pytorch.py @@ -15,6 +15,7 @@ cuda_version, get_cuda_include_dirs, debug_build_enabled, + nccl_ep_enabled, setup_mpi_flags, ) from typing import List @@ -89,7 +90,7 @@ def setup_pytorch_extension( # Mirror the NCCL EP gate from setup.py / common CMake. When disabled, the # ep.cpp source no-ops at the #ifdef boundary; without the define it would # produce undefined references to nvte_ep_*. - if bool(int(os.getenv("NVTE_WITH_NCCL_EP", "1"))): + if nccl_ep_enabled(): cxx_flags.append("-DNVTE_WITH_NCCL_EP") # PyTorch's symm-mem headers gate the NCCL_HAS_SYMMEM_* feature macros on # USE_NCCL. The EP extension shares the symm-mem NCCL comm with torch, so From 04b4b0f52843ede580a911059875fd802cee7cf1 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 14:42:36 +0200 Subject: [PATCH 07/15] Pass comm streams as raw cudaStream_t handles, torch 2.1 pybind lacks c10::Stream caster Signed-off-by: Pawel Gadzinski --- transformer_engine/pytorch/csrc/extensions.h | 8 ++++---- .../csrc/extensions/comm_gemm_overlap.cpp | 19 ++++++++++--------- .../pytorch/module/layernorm_linear.py | 8 ++++++-- .../pytorch/module/layernorm_mlp.py | 10 ++++++---- transformer_engine/pytorch/module/linear.py | 8 ++++++-- .../ops/fused/userbuffers_backward_linear.py | 8 ++++++-- 6 files changed, 38 insertions(+), 23 deletions(-) diff --git a/transformer_engine/pytorch/csrc/extensions.h b/transformer_engine/pytorch/csrc/extensions.h index 9114b1e453..b6d01cf7f8 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -755,8 +755,8 @@ void nvshmem_finalize(); * Comm+GEMM Overlap Wrappers **************************************************************************************************/ -void bulk_overlap_ag_with_external_gemm(CommOverlap &allgather_communicator, at::Stream send_stream, - at::Stream recv_stream); +void bulk_overlap_ag_with_external_gemm(CommOverlap &allgather_communicator, int64_t send_stream, + int64_t recv_stream); /*************************************************************************************************** * Newton-Schulz (cuSolverMp) @@ -843,7 +843,7 @@ class CommOverlap : torch::CustomClassHolder, public transformer_engine::CommOve at::Tensor get_buffer(bool local_chunk = false, std::optional> shape = std::nullopt); - std::pair get_communication_stream(); + std::pair get_communication_stream(); }; // CommOverlap @@ -876,7 +876,7 @@ class CommOverlapP2P : torch::CustomClassHolder, public transformer_engine::Comm at::Tensor get_buffer(bool local_chunk = false, std::optional> shape = std::nullopt); - std::pair get_communication_stream(); + std::pair get_communication_stream(); }; // CommOverlapP2P diff --git a/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp b/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp index 33237f0751..95e2fbb846 100644 --- a/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp +++ b/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp @@ -396,10 +396,9 @@ at::Tensor CommOverlap::get_buffer(bool local_chunk, std::optional CommOverlap::get_communication_stream() { +std::pair CommOverlap::get_communication_stream() { // Return the same stream for both send and recv - return {at::cuda::getStreamFromExternal(_stream_comm, at::cuda::current_device()), - at::cuda::getStreamFromExternal(_stream_comm, at::cuda::current_device())}; + return {reinterpret_cast(_stream_comm), reinterpret_cast(_stream_comm)}; } /*************************************************************************************************** @@ -499,14 +498,16 @@ at::Tensor CommOverlapP2P::get_buffer(bool local_chunk, std::optional CommOverlapP2P::get_communication_stream() { - return {at::cuda::getStreamFromExternal(_stream_send[0], at::cuda::current_device()), - at::cuda::getStreamFromExternal(_stream_recv, at::cuda::current_device())}; +std::pair CommOverlapP2P::get_communication_stream() { + return {reinterpret_cast(_stream_send[0]), reinterpret_cast(_stream_recv)}; } void transformer_engine::pytorch::bulk_overlap_ag_with_external_gemm( - CommOverlap &allgather_communicator, at::Stream send_stream, at::Stream recv_stream) { + CommOverlap &allgather_communicator, int64_t send_stream, int64_t recv_stream) { auto main_stream = at::cuda::getCurrentCUDAStream(); - allgather_communicator.bulk_overlap_external_ag(at::cuda::CUDAStream(send_stream), - at::cuda::CUDAStream(recv_stream), main_stream); + auto device = at::cuda::current_device(); + allgather_communicator.bulk_overlap_external_ag( + at::cuda::getStreamFromExternal(reinterpret_cast(send_stream), device), + at::cuda::getStreamFromExternal(reinterpret_cast(recv_stream), device), + main_stream); } diff --git a/transformer_engine/pytorch/module/layernorm_linear.py b/transformer_engine/pytorch/module/layernorm_linear.py index 561e813348..08a5047e91 100644 --- a/transformer_engine/pytorch/module/layernorm_linear.py +++ b/transformer_engine/pytorch/module/layernorm_linear.py @@ -956,7 +956,9 @@ def backward( # overlapping the AG operation with the dgrad GEMM. # Get the communication stream from the dgrad GEMM to use for the AG - dgrad_send_stream, dgrad_recv_stream = ub_obj_dgrad.get_communication_stream() + send_ptr, recv_ptr = ub_obj_dgrad.get_communication_stream() + dgrad_send_stream = torch.cuda.ExternalStream(send_ptr) + dgrad_recv_stream = torch.cuda.ExternalStream(recv_ptr) # This object is separate from the ub_obj_wgrad object which is passed to the GEMM ub_obj_overlap_wgrad = get_ub(ctx.ub_name + "_wgrad", ctx.fp8) @@ -976,7 +978,9 @@ def backward( # Allgather grad_outputs[0] using the dgrad streams so we can overlap with the fc2_dgrad gemm tex.bulk_overlap_ag_with_external_gemm( - ub_obj_overlap_wgrad, dgrad_send_stream, dgrad_recv_stream + ub_obj_overlap_wgrad, + dgrad_send_stream.cuda_stream, + dgrad_recv_stream.cuda_stream, ) # Prepare input tensor diff --git a/transformer_engine/pytorch/module/layernorm_mlp.py b/transformer_engine/pytorch/module/layernorm_mlp.py index 3ee0cda50c..030b410619 100644 --- a/transformer_engine/pytorch/module/layernorm_mlp.py +++ b/transformer_engine/pytorch/module/layernorm_mlp.py @@ -1321,9 +1321,9 @@ def backward( # overlapping the AG operation with the dgrad GEMM. # Get the communication stream from the dgrad GEMM to use for the AG - dgrad_send_stream, dgrad_recv_stream = ( - ub_obj_fc2_dgrad.get_communication_stream() - ) + send_ptr, recv_ptr = ub_obj_fc2_dgrad.get_communication_stream() + dgrad_send_stream = torch.cuda.ExternalStream(send_ptr) + dgrad_recv_stream = torch.cuda.ExternalStream(recv_ptr) ub_obj_fc2_wgrad = get_ub("fc2_wgrad", ctx.fp8) @@ -1342,7 +1342,9 @@ def backward( # Allgather grad_outputs[0] using the dgrad streams so we can overlap with the fc2_dgrad gemm tex.bulk_overlap_ag_with_external_gemm( - ub_obj_fc2_wgrad, dgrad_send_stream, dgrad_recv_stream + ub_obj_fc2_wgrad, + dgrad_send_stream.cuda_stream, + dgrad_recv_stream.cuda_stream, ) # Prepare input tensor diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 56622db5e6..7a1cc25f0b 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -1191,7 +1191,9 @@ def _linear_backward(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None], .. # overlapping the AG operation with the dgrad GEMM. # Get the communication stream from the dgrad GEMM to use for the AG - dgrad_send_stream, dgrad_recv_stream = ub_obj_dgrad.get_communication_stream() + send_ptr, recv_ptr = ub_obj_dgrad.get_communication_stream() + dgrad_send_stream = torch.cuda.ExternalStream(send_ptr) + dgrad_recv_stream = torch.cuda.ExternalStream(recv_ptr) # This object is separate from the ub_obj_wgrad object which is passed to the GEMM ub_obj_overlap_wgrad = get_ub(bwd_args.ub_name + "_wgrad", bwd_args.fp8) @@ -1211,7 +1213,9 @@ def _linear_backward(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None], .. # Allgather grad_outputs[0] using the dgrad streams so we can overlap with the fc2_dgrad gemm tex.bulk_overlap_ag_with_external_gemm( - ub_obj_overlap_wgrad, dgrad_send_stream, dgrad_recv_stream + ub_obj_overlap_wgrad, + dgrad_send_stream.cuda_stream, + dgrad_recv_stream.cuda_stream, ) if bwd_args.fp8 or bwd_args.debug: diff --git a/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py b/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py index c1070e38a6..f0c2c523b8 100644 --- a/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py +++ b/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py @@ -472,7 +472,9 @@ def _functional_backward( # overlapping the AG operation with the dgrad GEMM. # Get the communication stream from the dgrad GEMM to use for the AG - dgrad_send_stream, dgrad_recv_stream = ub_comm_dgrad.get_communication_stream() + send_ptr, recv_ptr = ub_comm_dgrad.get_communication_stream() + dgrad_send_stream = torch.cuda.ExternalStream(send_ptr) + dgrad_recv_stream = torch.cuda.ExternalStream(recv_ptr) ub_obj_overlap_wgrad = get_ub(ub_comm_name + "_wgrad", with_quantized_compute) @@ -491,7 +493,9 @@ def _functional_backward( # Allgather grad_outputs[0] using the dgrad streams so we can overlap with the fc2_dgrad gemm bulk_overlap_ag_with_external_gemm( - ub_obj_overlap_wgrad, dgrad_send_stream, dgrad_recv_stream + ub_obj_overlap_wgrad, + dgrad_send_stream.cuda_stream, + dgrad_recv_stream.cuda_stream, ) if tensor_parallel_mode == "column": From 4a3ad6630b0119aba1549ea3954f99792d6532f8 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 14:58:16 +0200 Subject: [PATCH 08/15] Revert "Pass comm streams as raw cudaStream_t handles, torch 2.1 pybind lacks c10::Stream caster" This reverts commit 04b4b0f52843ede580a911059875fd802cee7cf1. Signed-off-by: Pawel Gadzinski --- transformer_engine/pytorch/csrc/extensions.h | 8 ++++---- .../csrc/extensions/comm_gemm_overlap.cpp | 19 +++++++++---------- .../pytorch/module/layernorm_linear.py | 8 ++------ .../pytorch/module/layernorm_mlp.py | 10 ++++------ transformer_engine/pytorch/module/linear.py | 8 ++------ .../ops/fused/userbuffers_backward_linear.py | 8 ++------ 6 files changed, 23 insertions(+), 38 deletions(-) diff --git a/transformer_engine/pytorch/csrc/extensions.h b/transformer_engine/pytorch/csrc/extensions.h index b6d01cf7f8..9114b1e453 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -755,8 +755,8 @@ void nvshmem_finalize(); * Comm+GEMM Overlap Wrappers **************************************************************************************************/ -void bulk_overlap_ag_with_external_gemm(CommOverlap &allgather_communicator, int64_t send_stream, - int64_t recv_stream); +void bulk_overlap_ag_with_external_gemm(CommOverlap &allgather_communicator, at::Stream send_stream, + at::Stream recv_stream); /*************************************************************************************************** * Newton-Schulz (cuSolverMp) @@ -843,7 +843,7 @@ class CommOverlap : torch::CustomClassHolder, public transformer_engine::CommOve at::Tensor get_buffer(bool local_chunk = false, std::optional> shape = std::nullopt); - std::pair get_communication_stream(); + std::pair get_communication_stream(); }; // CommOverlap @@ -876,7 +876,7 @@ class CommOverlapP2P : torch::CustomClassHolder, public transformer_engine::Comm at::Tensor get_buffer(bool local_chunk = false, std::optional> shape = std::nullopt); - std::pair get_communication_stream(); + std::pair get_communication_stream(); }; // CommOverlapP2P diff --git a/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp b/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp index 95e2fbb846..33237f0751 100644 --- a/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp +++ b/transformer_engine/pytorch/csrc/extensions/comm_gemm_overlap.cpp @@ -396,9 +396,10 @@ at::Tensor CommOverlap::get_buffer(bool local_chunk, std::optional CommOverlap::get_communication_stream() { +std::pair CommOverlap::get_communication_stream() { // Return the same stream for both send and recv - return {reinterpret_cast(_stream_comm), reinterpret_cast(_stream_comm)}; + return {at::cuda::getStreamFromExternal(_stream_comm, at::cuda::current_device()), + at::cuda::getStreamFromExternal(_stream_comm, at::cuda::current_device())}; } /*************************************************************************************************** @@ -498,16 +499,14 @@ at::Tensor CommOverlapP2P::get_buffer(bool local_chunk, std::optional CommOverlapP2P::get_communication_stream() { - return {reinterpret_cast(_stream_send[0]), reinterpret_cast(_stream_recv)}; +std::pair CommOverlapP2P::get_communication_stream() { + return {at::cuda::getStreamFromExternal(_stream_send[0], at::cuda::current_device()), + at::cuda::getStreamFromExternal(_stream_recv, at::cuda::current_device())}; } void transformer_engine::pytorch::bulk_overlap_ag_with_external_gemm( - CommOverlap &allgather_communicator, int64_t send_stream, int64_t recv_stream) { + CommOverlap &allgather_communicator, at::Stream send_stream, at::Stream recv_stream) { auto main_stream = at::cuda::getCurrentCUDAStream(); - auto device = at::cuda::current_device(); - allgather_communicator.bulk_overlap_external_ag( - at::cuda::getStreamFromExternal(reinterpret_cast(send_stream), device), - at::cuda::getStreamFromExternal(reinterpret_cast(recv_stream), device), - main_stream); + allgather_communicator.bulk_overlap_external_ag(at::cuda::CUDAStream(send_stream), + at::cuda::CUDAStream(recv_stream), main_stream); } diff --git a/transformer_engine/pytorch/module/layernorm_linear.py b/transformer_engine/pytorch/module/layernorm_linear.py index 08a5047e91..561e813348 100644 --- a/transformer_engine/pytorch/module/layernorm_linear.py +++ b/transformer_engine/pytorch/module/layernorm_linear.py @@ -956,9 +956,7 @@ def backward( # overlapping the AG operation with the dgrad GEMM. # Get the communication stream from the dgrad GEMM to use for the AG - send_ptr, recv_ptr = ub_obj_dgrad.get_communication_stream() - dgrad_send_stream = torch.cuda.ExternalStream(send_ptr) - dgrad_recv_stream = torch.cuda.ExternalStream(recv_ptr) + dgrad_send_stream, dgrad_recv_stream = ub_obj_dgrad.get_communication_stream() # This object is separate from the ub_obj_wgrad object which is passed to the GEMM ub_obj_overlap_wgrad = get_ub(ctx.ub_name + "_wgrad", ctx.fp8) @@ -978,9 +976,7 @@ def backward( # Allgather grad_outputs[0] using the dgrad streams so we can overlap with the fc2_dgrad gemm tex.bulk_overlap_ag_with_external_gemm( - ub_obj_overlap_wgrad, - dgrad_send_stream.cuda_stream, - dgrad_recv_stream.cuda_stream, + ub_obj_overlap_wgrad, dgrad_send_stream, dgrad_recv_stream ) # Prepare input tensor diff --git a/transformer_engine/pytorch/module/layernorm_mlp.py b/transformer_engine/pytorch/module/layernorm_mlp.py index 030b410619..3ee0cda50c 100644 --- a/transformer_engine/pytorch/module/layernorm_mlp.py +++ b/transformer_engine/pytorch/module/layernorm_mlp.py @@ -1321,9 +1321,9 @@ def backward( # overlapping the AG operation with the dgrad GEMM. # Get the communication stream from the dgrad GEMM to use for the AG - send_ptr, recv_ptr = ub_obj_fc2_dgrad.get_communication_stream() - dgrad_send_stream = torch.cuda.ExternalStream(send_ptr) - dgrad_recv_stream = torch.cuda.ExternalStream(recv_ptr) + dgrad_send_stream, dgrad_recv_stream = ( + ub_obj_fc2_dgrad.get_communication_stream() + ) ub_obj_fc2_wgrad = get_ub("fc2_wgrad", ctx.fp8) @@ -1342,9 +1342,7 @@ def backward( # Allgather grad_outputs[0] using the dgrad streams so we can overlap with the fc2_dgrad gemm tex.bulk_overlap_ag_with_external_gemm( - ub_obj_fc2_wgrad, - dgrad_send_stream.cuda_stream, - dgrad_recv_stream.cuda_stream, + ub_obj_fc2_wgrad, dgrad_send_stream, dgrad_recv_stream ) # Prepare input tensor diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 7a1cc25f0b..56622db5e6 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -1191,9 +1191,7 @@ def _linear_backward(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None], .. # overlapping the AG operation with the dgrad GEMM. # Get the communication stream from the dgrad GEMM to use for the AG - send_ptr, recv_ptr = ub_obj_dgrad.get_communication_stream() - dgrad_send_stream = torch.cuda.ExternalStream(send_ptr) - dgrad_recv_stream = torch.cuda.ExternalStream(recv_ptr) + dgrad_send_stream, dgrad_recv_stream = ub_obj_dgrad.get_communication_stream() # This object is separate from the ub_obj_wgrad object which is passed to the GEMM ub_obj_overlap_wgrad = get_ub(bwd_args.ub_name + "_wgrad", bwd_args.fp8) @@ -1213,9 +1211,7 @@ def _linear_backward(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None], .. # Allgather grad_outputs[0] using the dgrad streams so we can overlap with the fc2_dgrad gemm tex.bulk_overlap_ag_with_external_gemm( - ub_obj_overlap_wgrad, - dgrad_send_stream.cuda_stream, - dgrad_recv_stream.cuda_stream, + ub_obj_overlap_wgrad, dgrad_send_stream, dgrad_recv_stream ) if bwd_args.fp8 or bwd_args.debug: diff --git a/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py b/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py index f0c2c523b8..c1070e38a6 100644 --- a/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py +++ b/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py @@ -472,9 +472,7 @@ def _functional_backward( # overlapping the AG operation with the dgrad GEMM. # Get the communication stream from the dgrad GEMM to use for the AG - send_ptr, recv_ptr = ub_comm_dgrad.get_communication_stream() - dgrad_send_stream = torch.cuda.ExternalStream(send_ptr) - dgrad_recv_stream = torch.cuda.ExternalStream(recv_ptr) + dgrad_send_stream, dgrad_recv_stream = ub_comm_dgrad.get_communication_stream() ub_obj_overlap_wgrad = get_ub(ub_comm_name + "_wgrad", with_quantized_compute) @@ -493,9 +491,7 @@ def _functional_backward( # Allgather grad_outputs[0] using the dgrad streams so we can overlap with the fc2_dgrad gemm bulk_overlap_ag_with_external_gemm( - ub_obj_overlap_wgrad, - dgrad_send_stream.cuda_stream, - dgrad_recv_stream.cuda_stream, + ub_obj_overlap_wgrad, dgrad_send_stream, dgrad_recv_stream ) if tensor_parallel_mode == "column": From 2087bd02b31f3c542eadeed113c852e5d8942dfd Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 14:58:16 +0200 Subject: [PATCH 09/15] Revert "Avoid passing std::optional to at::get_generator_or_default, torch 2.1 needs c10::optional" This reverts commit 061834fc3484755a92ca82a553be8bae86fdba44. Signed-off-by: Pawel Gadzinski --- transformer_engine/pytorch/csrc/extensions/attention.cpp | 4 ++-- transformer_engine/pytorch/csrc/extensions/cast.cpp | 8 ++++---- transformer_engine/pytorch/csrc/extensions/dropout.cpp | 4 ++-- transformer_engine/pytorch/csrc/quantizer.cpp | 4 ++-- 4 files changed, 10 insertions(+), 10 deletions(-) diff --git a/transformer_engine/pytorch/csrc/extensions/attention.cpp b/transformer_engine/pytorch/csrc/extensions/attention.cpp index 0b3b882768..eb8813d4a0 100644 --- a/transformer_engine/pytorch/csrc/extensions/attention.cpp +++ b/transformer_engine/pytorch/csrc/extensions/attention.cpp @@ -228,8 +228,8 @@ std::vector fused_attn_fwd( } // extract rng seed and offset - auto gen = at::check_generator( - rng_gen.has_value() ? *rng_gen : at::cuda::detail::getDefaultCUDAGenerator()); + auto gen = at::get_generator_or_default( + rng_gen, at::cuda::detail::getDefaultCUDAGenerator()); at::PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); auto options = torch::TensorOptions().dtype(torch::kInt64).device(torch::kCUDA); auto rng_state = torch::empty({2}, options); diff --git a/transformer_engine/pytorch/csrc/extensions/cast.cpp b/transformer_engine/pytorch/csrc/extensions/cast.cpp index 9f6a8d627e..50cc5e1bd4 100644 --- a/transformer_engine/pytorch/csrc/extensions/cast.cpp +++ b/transformer_engine/pytorch/csrc/extensions/cast.cpp @@ -173,8 +173,8 @@ void group_quantize_nvfp4_impl(const GroupedTensorWrapper &grouped_input_tensor, // number for different tensors in the group, so we only need to allocate one rng state const size_t rng_elts_per_thread = 1024 * num_tensors; rng_states_tensor = torch::empty({2}, opts); - auto gen = - at::check_generator(at::cuda::detail::getDefaultCUDAGenerator()); + auto gen = at::get_generator_or_default( + std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); at::PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); philox_unpack(philox_args, static_cast(rng_states_tensor.data_ptr())); @@ -1390,8 +1390,8 @@ static StochasticRngStateResources setup_stochastic_rounding_rng_states_helper( if (need_separate_rng_states) res.te_rng_state_list_colwise.reserve(num_tensors); for (size_t i = 0; i < num_tensors; ++i) { - auto gen = - at::check_generator(at::cuda::detail::getDefaultCUDAGenerator()); + auto gen = at::get_generator_or_default( + std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); // Rowwise RNG state at::PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); diff --git a/transformer_engine/pytorch/csrc/extensions/dropout.cpp b/transformer_engine/pytorch/csrc/extensions/dropout.cpp index 0bf44eebd0..bea8f3a7b5 100644 --- a/transformer_engine/pytorch/csrc/extensions/dropout.cpp +++ b/transformer_engine/pytorch/csrc/extensions/dropout.cpp @@ -44,8 +44,8 @@ std::vector dropout_fwd(const py::handle &input, float dropout_proba auto mask_nvte = makeTransformerEngineTensor(mask_pyt); // RNG state tensor - auto gen = - at::check_generator(at::cuda::detail::getDefaultCUDAGenerator()); + auto gen = at::get_generator_or_default( + std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); at::PhiloxCudaState philox_args; { std::lock_guard lock(gen->mutex_); diff --git a/transformer_engine/pytorch/csrc/quantizer.cpp b/transformer_engine/pytorch/csrc/quantizer.cpp index d2ca6b635a..4704261790 100644 --- a/transformer_engine/pytorch/csrc/quantizer.cpp +++ b/transformer_engine/pytorch/csrc/quantizer.cpp @@ -2603,8 +2603,8 @@ void NVFP4Quantizer::quantize_impl(const TensorWrapper& input, TensorWrapper& ou if (this->stochastic_rounding) { const size_t rng_elts_per_thread = 1024; // Wild guess, probably can be tightened - auto gen = - at::check_generator(at::cuda::detail::getDefaultCUDAGenerator()); + auto gen = at::get_generator_or_default( + std::nullopt, at::cuda::detail::getDefaultCUDAGenerator()); auto opts = at::TensorOptions().dtype(torch::kInt64).device(torch::kCUDA); // Generate RNG state for rowwise quantization From 690701a21be75ee10e92211ecb4b6672c4b323c4 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 14:59:01 +0200 Subject: [PATCH 10/15] Retarget CI workflow to torch 2.8 Signed-off-by: Pawel Gadzinski --- .../{build_pytorch21.yml => build_pytorch28.yml} | 13 ++++++------- 1 file changed, 6 insertions(+), 7 deletions(-) rename .github/workflows/{build_pytorch21.yml => build_pytorch28.yml} (81%) diff --git a/.github/workflows/build_pytorch21.yml b/.github/workflows/build_pytorch28.yml similarity index 81% rename from .github/workflows/build_pytorch21.yml rename to .github/workflows/build_pytorch28.yml index f577f148da..caa9392fbd 100644 --- a/.github/workflows/build_pytorch21.yml +++ b/.github/workflows/build_pytorch28.yml @@ -2,9 +2,8 @@ # # See LICENSE for license information. -# Verify that the torch>=2.1 requirement declared in build_tools/pytorch.py -# actually holds: build against torch==2.1.2 and run the sanity import. -name: 'Build PyTorch 2.1' +# Verify that TE builds and imports against torch==2.8. +name: 'Build PyTorch 2.8' on: pull_request: workflow_dispatch: @@ -12,8 +11,8 @@ concurrency: group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} cancel-in-progress: true jobs: - pytorch21: - name: 'PyTorch 2.1' + pytorch28: + name: 'PyTorch 2.8' runs-on: ubuntu-latest container: image: nvcr.io/nvidia/cuda:12.8.0-devel-ubuntu22.04 @@ -23,7 +22,7 @@ jobs: run: | apt-get update apt-get install -y git python3.9 pip cudnn9-cuda-12 - pip install torch==2.1.2 "numpy<2" + pip install torch==2.8.0 pip install cmake pybind11[global] ninja pydantic importlib-metadata>=1.0 packaging einops onnxscript "nvidia-cudnn-frontend>=1.25.0" - name: 'Checkout' uses: actions/checkout@v3 @@ -37,7 +36,7 @@ jobs: NVTE_FRAMEWORK: pytorch # Single old arch (Volta) to keep memory/build size down NVTE_CUDA_ARCHS: "70" - # torch 2.1 pulls in pip cudnn 8.9; point the TE build at apt cudnn 9 + # Prefer apt cudnn9 headers over the pip nvidia-cudnn-cu12 copy CUDNN_INCLUDE_PATH: /usr/include/x86_64-linux-gnu CUDNN_LIBRARY_PATH: /usr/lib/x86_64-linux-gnu # Full parallelism OOMs the 7GB runner (exit 137) From 7b7c6be418d2337db7c2a255d7f92b6cdf09d08c Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 15:49:35 +0200 Subject: [PATCH 11/15] Rename workflow to Minimum supported PyTorch, parametrize torch version Signed-off-by: Pawel Gadzinski --- .../{build_pytorch28.yml => minimum_pytorch.yml} | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) rename .github/workflows/{build_pytorch28.yml => minimum_pytorch.yml} (86%) diff --git a/.github/workflows/build_pytorch28.yml b/.github/workflows/minimum_pytorch.yml similarity index 86% rename from .github/workflows/build_pytorch28.yml rename to .github/workflows/minimum_pytorch.yml index caa9392fbd..ad24a02f7d 100644 --- a/.github/workflows/build_pytorch28.yml +++ b/.github/workflows/minimum_pytorch.yml @@ -2,17 +2,19 @@ # # See LICENSE for license information. -# Verify that TE builds and imports against torch==2.8. -name: 'Build PyTorch 2.8' +# Verify that TE builds and imports against the minimum supported PyTorch. +name: 'Minimum supported PyTorch' on: pull_request: workflow_dispatch: concurrency: group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} cancel-in-progress: true +env: + MIN_TORCH_VERSION: "2.8.0" jobs: - pytorch28: - name: 'PyTorch 2.8' + min-pytorch: + name: 'Minimum supported PyTorch' runs-on: ubuntu-latest container: image: nvcr.io/nvidia/cuda:12.8.0-devel-ubuntu22.04 @@ -22,7 +24,7 @@ jobs: run: | apt-get update apt-get install -y git python3.9 pip cudnn9-cuda-12 - pip install torch==2.8.0 + pip install torch==${MIN_TORCH_VERSION} pip install cmake pybind11[global] ninja pydantic importlib-metadata>=1.0 packaging einops onnxscript "nvidia-cudnn-frontend>=1.25.0" - name: 'Checkout' uses: actions/checkout@v3 From b86e86cd845cf9347e765358606936ed39d4d88b Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 16:21:03 +0200 Subject: [PATCH 12/15] Add Minimum supported JAX job, rename workflow to minimum_versions Signed-off-by: Pawel Gadzinski --- ...nimum_pytorch.yml => minimum_versions.yml} | 35 +++++++++++++++++-- 1 file changed, 33 insertions(+), 2 deletions(-) rename .github/workflows/{minimum_pytorch.yml => minimum_versions.yml} (59%) diff --git a/.github/workflows/minimum_pytorch.yml b/.github/workflows/minimum_versions.yml similarity index 59% rename from .github/workflows/minimum_pytorch.yml rename to .github/workflows/minimum_versions.yml index ad24a02f7d..216033548c 100644 --- a/.github/workflows/minimum_pytorch.yml +++ b/.github/workflows/minimum_versions.yml @@ -2,8 +2,8 @@ # # See LICENSE for license information. -# Verify that TE builds and imports against the minimum supported PyTorch. -name: 'Minimum supported PyTorch' +# Verify that TE builds and imports against the minimum supported framework versions. +name: 'Minimum supported versions' on: pull_request: workflow_dispatch: @@ -12,6 +12,7 @@ concurrency: cancel-in-progress: true env: MIN_TORCH_VERSION: "2.8.0" + MIN_JAX_VERSION: "0.5.3" jobs: min-pytorch: name: 'Minimum supported PyTorch' @@ -46,3 +47,33 @@ jobs: SCCACHE_GHA_ENABLED: "true" - name: 'Sanity check' run: python3 tests/pytorch/test_sanity_import.py + min-jax: + name: 'Minimum supported JAX' + runs-on: ubuntu-latest + container: + image: nvcr.io/nvidia/cuda:12.8.0-devel-ubuntu22.04 + options: --user root + steps: + - name: 'Dependencies' + run: | + apt-get update + apt-get install -y git python3.9 pip cudnn9-cuda-12 + pip install jax==${MIN_JAX_VERSION} "flax>=0.7.1" + pip install cmake pybind11[global] ninja "nvidia-cudnn-frontend>=1.25.0" + - name: 'Checkout' + uses: actions/checkout@v3 + with: + submodules: recursive + - name: ccache + uses: mozilla-actions/sccache-action@7d986dd989559c6ecdb630a3fd2557667be217ad + - name: 'Build' + run: NVTE_USE_CCACHE=1 NVTE_CCACHE_BIN=sccache pip install --no-build-isolation --no-deps . -v + env: + NVTE_FRAMEWORK: jax + NVTE_CUDA_ARCHS: "70" + CUDNN_INCLUDE_PATH: /usr/include/x86_64-linux-gnu + CUDNN_LIBRARY_PATH: /usr/lib/x86_64-linux-gnu + MAX_JOBS: 2 + SCCACHE_GHA_ENABLED: "true" + - name: 'Sanity check' + run: python3 tests/jax/test_sanity_import.py From 12f725fc07c8c719f89bb7d10d18a104c9577ba2 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 16:33:42 +0200 Subject: [PATCH 13/15] Add packaging to min-jax job deps Signed-off-by: Pawel Gadzinski --- .github/workflows/minimum_versions.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/minimum_versions.yml b/.github/workflows/minimum_versions.yml index 216033548c..efdeb1e834 100644 --- a/.github/workflows/minimum_versions.yml +++ b/.github/workflows/minimum_versions.yml @@ -59,7 +59,7 @@ jobs: apt-get update apt-get install -y git python3.9 pip cudnn9-cuda-12 pip install jax==${MIN_JAX_VERSION} "flax>=0.7.1" - pip install cmake pybind11[global] ninja "nvidia-cudnn-frontend>=1.25.0" + pip install cmake pybind11[global] ninja packaging "nvidia-cudnn-frontend>=1.25.0" - name: 'Checkout' uses: actions/checkout@v3 with: From f094de7a4cb0ef2fce3c94a839e6c959a2e1fea5 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 16:59:27 +0200 Subject: [PATCH 14/15] Skip NCCL EP in torch extension when torch lacks symm-mem headers; build min-torch job on sm90 Signed-off-by: Pawel Gadzinski --- .github/workflows/minimum_versions.yml | 7 ++++--- build_tools/pytorch.py | 25 ++++++++++++++++++++----- 2 files changed, 24 insertions(+), 8 deletions(-) diff --git a/.github/workflows/minimum_versions.yml b/.github/workflows/minimum_versions.yml index efdeb1e834..b7719394c3 100644 --- a/.github/workflows/minimum_versions.yml +++ b/.github/workflows/minimum_versions.yml @@ -24,7 +24,7 @@ jobs: - name: 'Dependencies' run: | apt-get update - apt-get install -y git python3.9 pip cudnn9-cuda-12 + apt-get install -y --allow-change-held-packages git python3.9 pip cudnn9-cuda-12 libnccl-dev libnccl2 pip install torch==${MIN_TORCH_VERSION} pip install cmake pybind11[global] ninja pydantic importlib-metadata>=1.0 packaging einops onnxscript "nvidia-cudnn-frontend>=1.25.0" - name: 'Checkout' @@ -37,8 +37,9 @@ jobs: run: NVTE_USE_CCACHE=1 NVTE_CCACHE_BIN=sccache pip install --no-build-isolation --no-deps . -v env: NVTE_FRAMEWORK: pytorch - # Single old arch (Volta) to keep memory/build size down - NVTE_CUDA_ARCHS: "70" + # Single arch to keep memory/build size down; sm90 so the NCCL EP + # path (and its torch-version gate) is exercised + NVTE_CUDA_ARCHS: "90" # Prefer apt cudnn9 headers over the pip nvidia-cudnn-cu12 copy CUDNN_INCLUDE_PATH: /usr/include/x86_64-linux-gnu CUDNN_LIBRARY_PATH: /usr/lib/x86_64-linux-gnu diff --git a/build_tools/pytorch.py b/build_tools/pytorch.py index 98331ccbd8..45a9055021 100644 --- a/build_tools/pytorch.py +++ b/build_tools/pytorch.py @@ -91,11 +91,26 @@ def setup_pytorch_extension( # ep.cpp source no-ops at the #ifdef boundary; without the define it would # produce undefined references to nvte_ep_*. if nccl_ep_enabled(): - cxx_flags.append("-DNVTE_WITH_NCCL_EP") - # PyTorch's symm-mem headers gate the NCCL_HAS_SYMMEM_* feature macros on - # USE_NCCL. The EP extension shares the symm-mem NCCL comm with torch, so - # it needs those macros visible. - cxx_flags.append("-DUSE_NCCL") + # ep.cpp additionally needs torch's NCCL symm-mem headers (torch >= 2.11). + from torch.utils.cpp_extension import include_paths + + symm_mem_header = "torch/csrc/distributed/c10d/symm_mem/nccl_dev_cap.hpp" + if any(os.path.exists(os.path.join(p, symm_mem_header)) for p in include_paths()): + cxx_flags.append("-DNVTE_WITH_NCCL_EP") + # PyTorch's symm-mem headers gate the NCCL_HAS_SYMMEM_* feature macros on + # USE_NCCL. The EP extension shares the symm-mem NCCL comm with torch, so + # it needs those macros visible. + cxx_flags.append("-DUSE_NCCL") + elif os.getenv("NVTE_WITH_NCCL_EP") == "1": + raise RuntimeError( + "NVTE_WITH_NCCL_EP=1 was set but the installed torch does not provide " + f"{symm_mem_header}. NCCL EP requires torch >= 2.11." + ) + else: + print( + f"[NCCL EP] Installed torch does not provide {symm_mem_header} " + "(torch >= 2.11 required); skipping NCCL EP in the torch extension." + ) library_dirs = [] libraries = [] From 98324dd53c4a08633a9322adefb1c44f6af5ce55 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 17:38:32 +0200 Subject: [PATCH 15/15] Fix min-versions jobs: pydantic for jax import, libcuda stub for sm90 import Signed-off-by: Pawel Gadzinski --- .github/workflows/minimum_versions.yml | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/.github/workflows/minimum_versions.yml b/.github/workflows/minimum_versions.yml index b7719394c3..ea50b5ee05 100644 --- a/.github/workflows/minimum_versions.yml +++ b/.github/workflows/minimum_versions.yml @@ -47,7 +47,10 @@ jobs: MAX_JOBS: 2 SCCACHE_GHA_ENABLED: "true" - name: 'Sanity check' - run: python3 tests/pytorch/test_sanity_import.py + # No GPU driver on the runner; the sm90 build links libcuda via NCCL EP + run: | + ln -s /usr/local/cuda/lib64/stubs/libcuda.so /usr/lib/x86_64-linux-gnu/libcuda.so.1 + python3 tests/pytorch/test_sanity_import.py min-jax: name: 'Minimum supported JAX' runs-on: ubuntu-latest @@ -60,7 +63,7 @@ jobs: apt-get update apt-get install -y git python3.9 pip cudnn9-cuda-12 pip install jax==${MIN_JAX_VERSION} "flax>=0.7.1" - pip install cmake pybind11[global] ninja packaging "nvidia-cudnn-frontend>=1.25.0" + pip install cmake pybind11[global] ninja packaging pydantic "nvidia-cudnn-frontend>=1.25.0" - name: 'Checkout' uses: actions/checkout@v3 with: