| |
|
|
| #ifndef CUDA_UTILS_CUH |
| #define CUDA_UTILS_CUH |
|
|
| #include "cuda_common.h" |
|
|
| |
| |
| |
| |
| |
|
|
| template<class ElementType> |
| struct alignas(16) Packed128 { |
| Packed128() = default; |
| __device__ explicit Packed128(int4 bits) { |
| static_assert(sizeof(bits) == sizeof(payload), "Size mismatch."); |
| memcpy(&payload, &bits, sizeof(bits)); |
| } |
|
|
| __device__ static Packed128 constant(ElementType value) { |
| Packed128 result; |
| for(int k = 0; k < size; ++k) { |
| result.payload[k] = value; |
| } |
| return result; |
| } |
| __device__ static Packed128 zeros() { |
| return constant(0.f); |
| } |
| __device__ static Packed128 ones() { |
| return constant(1.f); |
| } |
|
|
| __device__ ElementType& operator[](int index) { |
| return payload[index]; |
| } |
| __device__ const ElementType& operator[](int index) const { |
| return payload[index]; |
| } |
| __device__ int4 get_bits() const { |
| int4 bits; |
| static_assert(sizeof(bits) == sizeof(payload), "Size mismatch."); |
| memcpy(&bits, &payload, sizeof(bits)); |
| return bits; |
| } |
| static constexpr const size_t size = sizeof(int4) / sizeof(ElementType); |
| ElementType payload[size]; |
| }; |
|
|
| |
| template<class ElementType> |
| __device__ Packed128<ElementType> load128(const ElementType* address) { |
| return Packed128<ElementType>{*reinterpret_cast<const int4*>(address)}; |
| } |
| |
| template<class ElementType> |
| __device__ Packed128<ElementType> load128cs(const ElementType* address) { |
| return Packed128<ElementType>{__ldcs(reinterpret_cast<const int4*>(address))}; |
| } |
| |
| template<class ElementType> |
| __device__ void store128(ElementType* target, Packed128<ElementType> value) { |
| *reinterpret_cast<int4*>(target) = value.get_bits(); |
| } |
| |
| template<class ElementType> |
| __device__ void store128cs(ElementType* target, Packed128<ElementType> value) { |
| __stcs(reinterpret_cast<int4*>(target), value.get_bits()); |
| } |
| |
| template<class ElementType> |
| __device__ void store128cg(ElementType* target, Packed128<ElementType> value) { |
| __stcg(reinterpret_cast<int4*>(target), value.get_bits()); |
| } |
|
|
| |
| typedef Packed128<float> f128; |
| typedef Packed128<floatX> x128; |
|
|
| |
| |
|
|
| |
| enum class DType : uint8_t { |
| FP32, FP16, BF16 |
| }; |
|
|
| |
| |
| size_t sizeof_dtype(DType type) { |
| switch (type) { |
| case DType::FP32: |
| return sizeof(float); |
| case DType::FP16: |
| return sizeof(half); |
| case DType::BF16: |
| return sizeof(nv_bfloat16); |
| default: |
| fprintf(stderr, "Unknown datatype\n"); |
| exit(EXIT_FAILURE); |
| } |
| } |
|
|
| DType dtype_of(float* f) { return DType::FP32; } |
| DType dtype_of(nv_bfloat16 * f) { return DType::BF16; } |
| DType dtype_of(half * f) { return DType::FP16; } |
|
|
|
|
|
|
| |
| |
|
|
| |
| template<typename Td, typename Ts> |
| __device__ Td cast_value(Ts val); |
|
|
| template<> |
| __device__ float cast_value<float, float>(float val) { |
| return val; |
| } |
|
|
| template<> |
| __device__ float cast_value<float, half>(half val) { |
| return __half2float(val); |
| } |
|
|
| template<> |
| __device__ float cast_value<float, __nv_bfloat16>(__nv_bfloat16 val) { |
| return __bfloat162float(val); |
| } |
|
|
| template<typename Td, typename Ts> |
| __global__ void copy_and_cast_kernel(Td* dst, const Ts* src, size_t n, ptrdiff_t stride_dst, ptrdiff_t stride_src) { |
| int idx = blockIdx.x * blockDim.x + threadIdx.x; |
| |
| if (idx < n) { |
| dst[idx + stride_dst * blockIdx.y] = cast_value<Td, Ts>(src[idx + stride_src * blockIdx.y]); |
| } |
| } |
|
|
| |
| |
|
|
| |
| __device__ inline float warpReduceSum(float val) { |
| for (int offset = 16; offset > 0; offset /= 2) { |
| val += __shfl_xor_sync(0xFFFFFFFF, val, offset); |
| } |
| return val; |
| } |
| |
| __device__ inline float warpReduceMax(float val) { |
| for (int offset = 16; offset > 0; offset /= 2) { |
| val = fmaxf(val, __shfl_xor_sync(0xFFFFFFFF, val, offset)); |
| } |
| return val; |
| } |
| |
| |
| |
| |
| using reduction_func_t = float (*) (float); |
| template<reduction_func_t warp_reduction> |
| __device__ inline float blockReduce(float val, bool final_sync=false, float out_of_bounds=0.0f) { |
| |
| |
| __shared__ float shared_val[WARP_SIZE]; |
| const int lane_id = threadIdx.x % WARP_SIZE; |
| const int warp_id = threadIdx.x / WARP_SIZE; |
| const int num_warps = blockDim.x / WARP_SIZE; |
|
|
| float warp_val = warp_reduction(val); |
| if (lane_id == 0) { shared_val[warp_id] = warp_val; } |
| __syncthreads(); |
| warp_val = (lane_id < num_warps) ? shared_val[lane_id] : out_of_bounds; |
| float block_val = warp_reduction(warp_val); |
|
|
| if (final_sync) { |
| __syncthreads(); |
| } |
| return block_val; |
| } |
|
|
| |
| |
| template<class Float> |
| __global__ void global_sum_single_block_kernel(float* result, const Float* values, size_t count) { |
| assert(gridDim.x == 1); |
| float thread_sum = 0; |
| for(size_t index = threadIdx.x; index < count; index += blockDim.x) { |
| thread_sum += (float)values[index]; |
| } |
|
|
| float reduction = blockReduce<warpReduceSum>(thread_sum, true); |
| if(threadIdx.x == 0) { |
| *result = reduction; |
| } |
| } |
|
|
| template<class Float> |
| void global_sum_deterministic(float* result, const Float* values, int count, cudaStream_t stream) { |
| global_sum_single_block_kernel<<<1, 1024, 0, stream>>>(result, values, count); |
| cudaCheck(cudaGetLastError()); |
| } |
|
|
| |
| |
|
|
| |
| |
| int cudaMallocConditionallyManaged(void** out, size_t bytes, const char *file, int line) { |
| |
| cudaError_t err = cudaMalloc(out, bytes); |
| if(err == cudaErrorMemoryAllocation) { |
| |
| cudaGetLastError(); |
| cudaCheck_(cudaMallocManaged(out, bytes), file, line); |
| cudaCheck_(cudaMemAdvise(*out, bytes, cudaMemAdviseSetPreferredLocation, cudaCpuDeviceId), file, line); |
| return 1; |
| } else { |
| cudaCheck_(err, file, line); |
| return 0; |
| } |
| } |
|
|
| #define cudaMallocConditionallyManaged(out, bytes)\ |
| (cudaMallocConditionallyManaged((void**)out, bytes, __FILE__, __LINE__)) |
|
|
| |
| |
|
|
| |
| |
| |
| |
| __device__ __host__ constexpr unsigned int SquirrelNoise5(unsigned int positionX, unsigned int seed) |
| { |
| constexpr unsigned int SQ5_BIT_NOISE1 = 0xd2a80a3f; |
| constexpr unsigned int SQ5_BIT_NOISE2 = 0xa884f197; |
| constexpr unsigned int SQ5_BIT_NOISE3 = 0x6C736F4B; |
| constexpr unsigned int SQ5_BIT_NOISE4 = 0xB79F3ABB; |
| constexpr unsigned int SQ5_BIT_NOISE5 = 0x1b56c4f5; |
| unsigned int mangledBits = positionX; |
| mangledBits *= SQ5_BIT_NOISE1; |
| mangledBits += seed; |
| mangledBits ^= (mangledBits >> 9); |
| mangledBits += SQ5_BIT_NOISE2; |
| mangledBits ^= (mangledBits >> 11); |
| mangledBits *= SQ5_BIT_NOISE3; |
| mangledBits ^= (mangledBits >> 13); |
| mangledBits += SQ5_BIT_NOISE4; |
| mangledBits ^= (mangledBits >> 15); |
| mangledBits *= SQ5_BIT_NOISE5; |
| mangledBits ^= (mangledBits >> 17); |
| return mangledBits; |
| } |
| __device__ __host__ constexpr unsigned int Get2dNoiseUint(int indexX, int indexY, unsigned int seed) |
| { |
| constexpr unsigned int PRIME_NUMBER = 198491317u; |
| unsigned int x = static_cast<unsigned int>(indexX); |
| unsigned int y = static_cast<unsigned int>(indexY); |
|
|
| return SquirrelNoise5(x + (PRIME_NUMBER * y), seed); |
| } |
|
|
| |
| __device__ __forceinline__ void stochastic_rounding(float in, __nv_bfloat16 *out, unsigned int seed) { |
| |
| |
| unsigned int random = Get2dNoiseUint(threadIdx.x, blockIdx.x * blockDim.x + blockIdx.y, seed); |
| unsigned int threshold = random & 0xFFFF; |
| unsigned int float_bits = __float_as_uint(in); |
| unsigned int rounded_bits = float_bits & 0x0000FFFF; |
| float_bits = (rounded_bits > threshold) ? (float_bits | 0xFFFF) : (float_bits & ~0xFFFF); |
| *out = __float2bfloat16_rn(__uint_as_float(float_bits)); |
| } |
| __device__ __forceinline__ void stochastic_rounding(float in, half *out, unsigned int random) { |
| *out = (float)in; |
| } |
| __device__ __forceinline__ void stochastic_rounding(float in, float *out, unsigned int random) { |
| *out = in; |
| } |
|
|
| #endif |