From 3cf49f249b2f8cdebb2bf2e5d93fadbc8c9a9b3c Mon Sep 17 00:00:00 2001 From: Aidyn-A Date: Thu, 16 Jul 2026 11:48:57 +0400 Subject: [PATCH 1/5] Fix GroupNorm NHWC one-pass backward execution on Thor --- .../group_norm/group_norm_nhwc_bwd_one_pass.h | 5 +- .../group_norm_nhwc_bwd_one_pass_kernel.cuh | 79 +++++++++++++++++++ 2 files changed, 82 insertions(+), 2 deletions(-) diff --git a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h index c85181a52..fa503df1b 100644 --- a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h +++ b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h @@ -128,8 +128,9 @@ void group_norm_nhwc_bwd_one_pass_setup(Group_norm_nhwc_bwd_params& params, size // The number of blocks per grid. int max_blocks_per_grid = blocks_per_sm * props.multiProcessorCount; - // Make sure we are safe to run that many blocks - assert(blocks_per_slice <= max_blocks_per_grid); + // Cooperative kernels require all blocks to be resident concurrently. Blocks process + // additional activation tiles in a grid-stride loop when the full grid does not fit. + blocks_per_slice = std::min(blocks_per_slice, max_blocks_per_grid); // The number of blocks per slice is the X dimension of the grid. grid.x = blocks_per_slice; diff --git a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh index 008520414..7ed65ba15 100644 --- a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh +++ b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh @@ -173,6 +173,47 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa mean_2 += dx_norm_y; } + // A cooperative launch may use fewer blocks than activation tiles. Accumulate any + // additional tiles assigned to this block without increasing its register footprint. + for (int extra_hwi = hwi + gridDim.x * params.acts_per_block; extra_hwi < params.hw; + extra_hwi += gridDim.x * params.acts_per_block) { +#pragma unroll + for (int ii = 0; ii < ACTS_PER_THREAD; ++ii) { + int hwj = extra_hwi + ii * ACTS_PER_LOOP; + IOType2 x_extra = IOTraits::zero(); + IOType2 dy_extra = IOTraits::zero(); + if (is_active && hwj < params.hw) { + x_extra = *reinterpret_cast(&x_ptr[hwj * params.c]); + dy_extra = *reinterpret_cast(&dy_ptr[hwj * params.c]); + } + + float2 x_f2 = IOTraits::unpack(x_extra); + float2 dy_f2 = IOTraits::unpack(dy_extra); + + float x_norm_x = (x_f2.x - x_mean) * rcp_x_stddev; + float x_norm_y = (x_f2.y - x_mean) * rcp_x_stddev; + + if (params.with_swish) { + float x_gn_x = x_norm_x * gamma_f2.x + beta_f2.x; + float x_gn_y = x_norm_y * gamma_f2.y + beta_f2.y; + float s_x = sigmoid(x_gn_x); + float s_y = sigmoid(x_gn_y); + dy_f2.x = dy_f2.x * s_x * (1.f + x_gn_x * (1.f - s_x)); + dy_f2.y = dy_f2.y * s_y * (1.f + x_gn_y * (1.f - s_y)); + } + + dgamma_dbeta.x += dy_f2.x * x_norm_x; + dgamma_dbeta.y += dy_f2.y * x_norm_y; + dgamma_dbeta.z += dy_f2.x; + dgamma_dbeta.w += dy_f2.y; + + float dx_norm_x = dy_f2.x * gamma_f2.x; + float dx_norm_y = dy_f2.y * gamma_f2.y; + mean_1 += dx_norm_x * x_norm_x + dx_norm_y * x_norm_y; + mean_2 += dx_norm_x + dx_norm_y; + } + } + // Pack valid gradients. float2 sums = make_float2(0.f, 0.f); if (ACTIVE_THREADS == THREADS_PER_BLOCK || is_active) { @@ -308,6 +349,44 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa *reinterpret_cast(&dx_ptr[hwj * params.c]) = IOTraits::pack(dx); } } + + // Store gradients for any additional activation tiles assigned to this block. + for (int extra_hwi = hwi + gridDim.x * params.acts_per_block; extra_hwi < params.hw; + extra_hwi += gridDim.x * params.acts_per_block) { +#pragma unroll + for (int ii = 0; ii < ACTS_PER_THREAD; ++ii) { + int hwj = extra_hwi + ii * ACTS_PER_LOOP; + if (!is_active || hwj >= params.hw) { + continue; + } + + float2 x_f2 = IOTraits::unpack(*reinterpret_cast(&x_ptr[hwj * params.c])); + float2 dy_f2 = IOTraits::unpack(*reinterpret_cast(&dy_ptr[hwj * params.c])); + + float2 x_norm; + x_norm.x = (x_f2.x - x_mean) * rcp_x_stddev; + x_norm.y = (x_f2.y - x_mean) * rcp_x_stddev; + + if (params.with_swish) { + float x_gn_x = x_norm.x * gamma_f2.x + beta_f2.x; + float x_gn_y = x_norm.y * gamma_f2.y + beta_f2.y; + float s_x = sigmoid(x_gn_x); + float s_y = sigmoid(x_gn_y); + dy_f2.x = dy_f2.x * s_x * (1.f + x_gn_x * (1.f - s_x)); + dy_f2.y = dy_f2.y * s_y * (1.f + x_gn_y * (1.f - s_y)); + } + + float2 dx_norm; + dx_norm.x = dy_f2.x * gamma_f2.x; + dx_norm.y = dy_f2.y * gamma_f2.y; + + float2 dx; + dx.x = (dx_norm.x - (x_norm.x * mean_1 + mean_2)) * rcp_x_stddev; + dx.y = (dx_norm.y - (x_norm.y * mean_1 + mean_2)) * rcp_x_stddev; + + *reinterpret_cast(&dx_ptr[hwj * params.c]) = IOTraits::pack(dx); + } + } } // The completion barrier. From 661aee82f18402ddcaa0ed92e1e0c232eaa980e9 Mon Sep 17 00:00:00 2001 From: Aidyn-A Date: Tue, 4 Aug 2026 09:11:13 +0400 Subject: [PATCH 2/5] Guard against sm_110 and sm_120 --- .../csrc/group_norm/group_norm_nhwc_bwd_one_pass.h | 11 ++++++++--- .../group_norm_nhwc_bwd_one_pass_kernel.cuh | 4 ++++ 2 files changed, 12 insertions(+), 3 deletions(-) diff --git a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h index fa503df1b..3478936e4 100644 --- a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h +++ b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h @@ -128,9 +128,14 @@ void group_norm_nhwc_bwd_one_pass_setup(Group_norm_nhwc_bwd_params& params, size // The number of blocks per grid. int max_blocks_per_grid = blocks_per_sm * props.multiProcessorCount; - // Cooperative kernels require all blocks to be resident concurrently. Blocks process - // additional activation tiles in a grid-stride loop when the full grid does not fit. - blocks_per_slice = std::min(blocks_per_slice, max_blocks_per_grid); + if (props.major == 11 || props.major == 12) { + // Cooperative kernels require all blocks to be resident concurrently. Blocks process + // additional activation tiles in a grid-stride loop when the full grid does not fit. + blocks_per_slice = std::min(blocks_per_slice, max_blocks_per_grid); + } else { + // Make sure we are safe to run that many blocks + assert(blocks_per_slice <= max_blocks_per_grid); + } // The number of blocks per slice is the X dimension of the grid. grid.x = blocks_per_slice; diff --git a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh index 7ed65ba15..605794be2 100644 --- a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh +++ b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh @@ -173,6 +173,7 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa mean_2 += dx_norm_y; } +#if __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 // A cooperative launch may use fewer blocks than activation tiles. Accumulate any // additional tiles assigned to this block without increasing its register footprint. for (int extra_hwi = hwi + gridDim.x * params.acts_per_block; extra_hwi < params.hw; @@ -213,6 +214,7 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa mean_2 += dx_norm_x + dx_norm_y; } } +#endif // __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 // Pack valid gradients. float2 sums = make_float2(0.f, 0.f); @@ -350,6 +352,7 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa } } +#if __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 // Store gradients for any additional activation tiles assigned to this block. for (int extra_hwi = hwi + gridDim.x * params.acts_per_block; extra_hwi < params.hw; extra_hwi += gridDim.x * params.acts_per_block) { @@ -387,6 +390,7 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa *reinterpret_cast(&dx_ptr[hwj * params.c]) = IOTraits::pack(dx); } } +#endif // __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 } // The completion barrier. From bc0bcd094d58a6d4af96d495400ac5b4b67c324e Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 5 Aug 2026 05:49:14 +0000 Subject: [PATCH 3/5] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh index 605794be2..58046c660 100644 --- a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh +++ b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh @@ -214,7 +214,7 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa mean_2 += dx_norm_x + dx_norm_y; } } -#endif // __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 +#endif // __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 // Pack valid gradients. float2 sums = make_float2(0.f, 0.f); @@ -390,7 +390,7 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa *reinterpret_cast(&dx_ptr[hwj * params.c]) = IOTraits::pack(dx); } } -#endif // __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 +#endif // __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 } // The completion barrier. From 34ec79438de2c8f8b204e0d9db8bc8a0ef73b2de Mon Sep 17 00:00:00 2001 From: Aidyn-A Date: Fri, 7 Aug 2026 16:14:11 +0400 Subject: [PATCH 4/5] Add sm_8x and unskip test --- apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h | 2 +- .../csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh | 4 ++-- apex/contrib/test/group_norm/test_group_norm.py | 4 ---- 3 files changed, 3 insertions(+), 7 deletions(-) diff --git a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h index 3478936e4..dc54570f4 100644 --- a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h +++ b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass.h @@ -128,7 +128,7 @@ void group_norm_nhwc_bwd_one_pass_setup(Group_norm_nhwc_bwd_params& params, size // The number of blocks per grid. int max_blocks_per_grid = blocks_per_sm * props.multiProcessorCount; - if (props.major == 11 || props.major == 12) { + if (props.major == 8 || props.major == 11 || props.major == 12) { // Cooperative kernels require all blocks to be resident concurrently. Blocks process // additional activation tiles in a grid-stride loop when the full grid does not fit. blocks_per_slice = std::min(blocks_per_slice, max_blocks_per_grid); diff --git a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh index 58046c660..501dbf377 100644 --- a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh +++ b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh @@ -352,7 +352,7 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa } } -#if __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 +#if __CUDA_ARCH__ / 100 == 8 || __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 // Store gradients for any additional activation tiles assigned to this block. for (int extra_hwi = hwi + gridDim.x * params.acts_per_block; extra_hwi < params.hw; extra_hwi += gridDim.x * params.acts_per_block) { @@ -390,7 +390,7 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa *reinterpret_cast(&dx_ptr[hwj * params.c]) = IOTraits::pack(dx); } } -#endif // __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 +#endif // __CUDA_ARCH__ / 100 == 8 || __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 } // The completion barrier. diff --git a/apex/contrib/test/group_norm/test_group_norm.py b/apex/contrib/test/group_norm/test_group_norm.py index 99bc7072b..5e82e4d79 100644 --- a/apex/contrib/test/group_norm/test_group_norm.py +++ b/apex/contrib/test/group_norm/test_group_norm.py @@ -107,10 +107,6 @@ def _has_sufficient_cuda_memory(required_bytes: int, *, safety_factor: float = 0 return required_bytes <= int(free_bytes * safety_factor) -@unittest.skipIf( - torch.cuda.get_device_properties().multi_processor_count < 16, - "GroupNorm is unsupported on low SM count devices", -) @unittest.skipIf(SKIP_TEST, f"{SKIP_TEST}") class GroupNormTest(unittest.TestCase): def setUp(self, seed=0): From 569d275b05105506ad356d7eb6f682e66cd0e6e4 Mon Sep 17 00:00:00 2001 From: Aidyn-A <31858918+Aidyn-A@users.noreply.github.com> Date: Wed, 12 Aug 2026 16:23:20 +0400 Subject: [PATCH 5/5] add missing `__CUDA_ARCH__ / 100 == 8` --- .../csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh index 501dbf377..5b2ddf698 100644 --- a/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh +++ b/apex/contrib/csrc/group_norm/group_norm_nhwc_bwd_one_pass_kernel.cuh @@ -173,7 +173,7 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa mean_2 += dx_norm_y; } -#if __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 +#if __CUDA_ARCH__ / 100 == 8 || __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 // A cooperative launch may use fewer blocks than activation tiles. Accumulate any // additional tiles assigned to this block without increasing its register footprint. for (int extra_hwi = hwi + gridDim.x * params.acts_per_block; extra_hwi < params.hw; @@ -214,7 +214,7 @@ __global__ __launch_bounds__(THREADS_PER_BLOCK_) void group_norm_nhwc_bwd_one_pa mean_2 += dx_norm_x + dx_norm_y; } } -#endif // __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 +#endif // __CUDA_ARCH__ / 100 == 8 || __CUDA_ARCH__ / 100 == 11 || __CUDA_ARCH__ / 100 == 12 // Pack valid gradients. float2 sums = make_float2(0.f, 0.f);