diff --git a/crates/ml/src/cuda_pipeline/common_device_functions.cuh b/crates/ml/src/cuda_pipeline/common_device_functions.cuh index 49e8f7aa1..e6b1f235a 100644 --- a/crates/ml/src/cuda_pipeline/common_device_functions.cuh +++ b/crates/ml/src/cuda_pipeline/common_device_functions.cuh @@ -56,6 +56,27 @@ __device__ __forceinline__ __nv_bfloat16 bf16(float x) { return __float2bfloat16(x); } +/** BF16 warp shuffle — wraps __shfl_xor_sync for __nv_bfloat16. */ +__device__ __forceinline__ __nv_bfloat16 bf16_shfl_xor(unsigned mask, __nv_bfloat16 val, int offset) { + return __float2bfloat16(__shfl_xor_sync(mask, __bfloat162float(val), offset)); +} +/** BF16 warp shuffle down. */ +__device__ __forceinline__ __nv_bfloat16 bf16_shfl_down(unsigned mask, __nv_bfloat16 val, int offset) { + return __float2bfloat16(__shfl_down_sync(mask, __bfloat162float(val), offset)); +} +/** BF16 warp-level sum reduction (16 lanes). */ +__device__ __forceinline__ __nv_bfloat16 bf16_warp_sum(__nv_bfloat16 val) { + for (int offset = 16; offset > 0; offset >>= 1) + val = val + bf16_shfl_xor(0xFFFFFFFF, val, offset); + return val; +} +/** BF16 warp-level max reduction. */ +__device__ __forceinline__ __nv_bfloat16 bf16_warp_max(__nv_bfloat16 val) { + for (int offset = 16; offset > 0; offset >>= 1) + val = bf16_fmax(val, bf16_shfl_xor(0xFFFFFFFF, val, offset)); + return val; +} + /** BF16 leaky ReLU — operates on raw bf16, returns bf16. */ __device__ __forceinline__ __nv_bfloat16 leaky_relu_bf16(__nv_bfloat16 x) { return (x > bf16_zero()) ? x : bf16(0.01f) * x;