refactor: rewrite common_device_functions.cuh — bf16 wrappers now pure f32 identity functions
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -6,111 +6,55 @@
|
||||
*/
|
||||
|
||||
/* ------------------------------------------------------------------ */
|
||||
/* BF16 Mixed Precision Support (H100 Tensor Core Path) */
|
||||
/* F32/TF32 Precision — All GPU buffers are native float. */
|
||||
/* ------------------------------------------------------------------ */
|
||||
/* H100 SM90: 989 TFLOPS BF16 tensor cores vs 67 TFLOPS F32 scalar.
|
||||
* Strategy: BF16 storage + BF16 matmul inputs + F32 accumulation.
|
||||
* This matches NVIDIA's native tensor core accumulate mode. */
|
||||
#include <cuda_bf16.h>
|
||||
/* TF32 tensor cores (19-bit mantissa) activated via cublasLtMatmul */
|
||||
/* with CUBLAS_COMPUTE_32F. Storage is pure f32 everywhere. */
|
||||
|
||||
/* ── BF16 native math wrappers ─────────────────────────────────────── */
|
||||
/* Thin wrappers around F32 transcendentals for BF16 arguments. */
|
||||
/* The cast is hidden inside — kernel code reads as pure BF16. */
|
||||
/* Arithmetic (+, -, *, /, >, <) uses native __nv_bfloat16 ops (SM80+)*/
|
||||
/* ── F32 math wrappers (kept as thin aliases for kernel readability) ── */
|
||||
__device__ __forceinline__ float bf16_sqrt(float x) { return sqrtf(x); }
|
||||
__device__ __forceinline__ float bf16_log(float x) { return logf(x); }
|
||||
__device__ __forceinline__ float bf16_exp(float x) { return expf(x); }
|
||||
__device__ __forceinline__ float bf16_pow(float x, float p) { return powf(x, p); }
|
||||
__device__ __forceinline__ float bf16_fabs(float x) { return fabsf(x); }
|
||||
__device__ __forceinline__ float bf16_fmax(float a, float b) { return fmaxf(a, b); }
|
||||
__device__ __forceinline__ float bf16_fmin(float a, float b) { return fminf(a, b); }
|
||||
__device__ __forceinline__ float bf16_cos(float x) { return cosf(x); }
|
||||
__device__ __forceinline__ float bf16_sin(float x) { return sinf(x); }
|
||||
__device__ __forceinline__ float bf16_tanh(float x) { return tanhf(x); }
|
||||
__device__ __forceinline__ float bf16_floor(float x) { return floorf(x); }
|
||||
__device__ __forceinline__ float bf16_powf(float x, float p) { return powf(x, p); }
|
||||
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_sqrt(__nv_bfloat16 x) {
|
||||
return __float2bfloat16(sqrtf(__bfloat162float(x)));
|
||||
}
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_log(__nv_bfloat16 x) {
|
||||
return __float2bfloat16(logf(__bfloat162float(x)));
|
||||
}
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_exp(__nv_bfloat16 x) {
|
||||
return __float2bfloat16(expf(__bfloat162float(x)));
|
||||
}
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_pow(__nv_bfloat16 x, __nv_bfloat16 p) {
|
||||
return __float2bfloat16(powf(__bfloat162float(x), __bfloat162float(p)));
|
||||
}
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_fabs(__nv_bfloat16 x) {
|
||||
return __float2bfloat16(fabsf(__bfloat162float(x)));
|
||||
}
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_fmax(__nv_bfloat16 a, __nv_bfloat16 b) {
|
||||
return __float2bfloat16(fmaxf(__bfloat162float(a), __bfloat162float(b)));
|
||||
}
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_fmin(__nv_bfloat16 a, __nv_bfloat16 b) {
|
||||
return __float2bfloat16(fminf(__bfloat162float(a), __bfloat162float(b)));
|
||||
}
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_cos(__nv_bfloat16 x) {
|
||||
return __float2bfloat16(cosf(__bfloat162float(x)));
|
||||
}
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_sin(__nv_bfloat16 x) {
|
||||
return __float2bfloat16(sinf(__bfloat162float(x)));
|
||||
}
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_tanh(__nv_bfloat16 x) {
|
||||
return __float2bfloat16(tanhf(__bfloat162float(x)));
|
||||
}
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_floor(__nv_bfloat16 x) {
|
||||
return __float2bfloat16(floorf(__bfloat162float(x)));
|
||||
}
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_powf(__nv_bfloat16 x, float p) {
|
||||
return __float2bfloat16(powf(__bfloat162float(x), p));
|
||||
}
|
||||
__device__ __forceinline__ float bf16_zero() { return 0.0f; }
|
||||
__device__ __forceinline__ float bf16_one() { return 1.0f; }
|
||||
__device__ __forceinline__ float bf16(float x) { return x; }
|
||||
|
||||
/** BF16 zero constant */
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_zero() {
|
||||
return __float2bfloat16(0.0f);
|
||||
__device__ __forceinline__ float bf16_shfl_xor(unsigned mask, float val, int offset) {
|
||||
return __shfl_xor_sync(mask, val, offset);
|
||||
}
|
||||
/** BF16 one constant */
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_one() {
|
||||
return __float2bfloat16(1.0f);
|
||||
__device__ __forceinline__ float bf16_shfl_down(unsigned mask, float val, int offset) {
|
||||
return __shfl_down_sync(mask, val, offset);
|
||||
}
|
||||
/** BF16 from float scalar (for host-passed params like lr, gamma) */
|
||||
__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) {
|
||||
__device__ __forceinline__ float bf16_warp_sum(float val) {
|
||||
for (int offset = 16; offset > 0; offset >>= 1)
|
||||
val = val + bf16_shfl_xor(0xFFFFFFFF, val, offset);
|
||||
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
||||
return val;
|
||||
}
|
||||
/** BF16 warp-level max reduction. */
|
||||
__device__ __forceinline__ __nv_bfloat16 bf16_warp_max(__nv_bfloat16 val) {
|
||||
__device__ __forceinline__ float bf16_warp_max(float val) {
|
||||
for (int offset = 16; offset > 0; offset >>= 1)
|
||||
val = bf16_fmax(val, bf16_shfl_xor(0xFFFFFFFF, val, offset));
|
||||
val = fmaxf(val, __shfl_xor_sync(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;
|
||||
__device__ __forceinline__ float leaky_relu_bf16(float x) {
|
||||
return (x > 0.0f) ? x : 0.01f * x;
|
||||
}
|
||||
|
||||
/** Convert F32 to BF16 (truncation, matching PyTorch default). */
|
||||
__device__ __forceinline__ __nv_bfloat16 f32_to_bf16(float x) {
|
||||
return __float2bfloat16(x);
|
||||
}
|
||||
__device__ __forceinline__ float f32_to_bf16(float x) { return x; }
|
||||
|
||||
/** BF16 atomicAdd via CAS loop (no native BF16 atomicAdd on any SM). */
|
||||
__device__ __forceinline__ void atomicAddBF16(__nv_bfloat16* addr, __nv_bfloat16 val) {
|
||||
unsigned short* addr_as_us = (unsigned short*)addr;
|
||||
unsigned short old = *addr_as_us;
|
||||
unsigned short assumed;
|
||||
do {
|
||||
assumed = old;
|
||||
__nv_bfloat16 sum = __float2bfloat16(
|
||||
__bfloat162float(*(__nv_bfloat16*)&assumed) + __bfloat162float(val)
|
||||
);
|
||||
old = atomicCAS(addr_as_us, assumed, *(unsigned short*)&sum);
|
||||
} while (assumed != old);
|
||||
/* atomicAdd for float is native on SM30+ — no CAS loop needed */
|
||||
__device__ __forceinline__ void atomicAddBF16(float* addr, float val) {
|
||||
atomicAdd(addr, val);
|
||||
}
|
||||
|
||||
/* ------------------------------------------------------------------ */
|
||||
|
||||
Reference in New Issue
Block a user