Skip to content

Commit be872eb

Browse files
authored
[CUDA] implement Hadamard transform (#3179)
1 parent 3b3590b commit be872eb

6 files changed

Lines changed: 432 additions & 4 deletions

File tree

mlx/backend/cuda/CMakeLists.txt

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@ target_sources(
3030
${CMAKE_CURRENT_SOURCE_DIR}/gemms/gemv.cu
3131
${CMAKE_CURRENT_SOURCE_DIR}/gemms/cublas_gemm.cpp
3232
${CMAKE_CURRENT_SOURCE_DIR}/gemms/grouped_gemm_unaligned.cu
33+
${CMAKE_CURRENT_SOURCE_DIR}/hadamard.cu
3334
${CMAKE_CURRENT_SOURCE_DIR}/jit_module.cpp
3435
${CMAKE_CURRENT_SOURCE_DIR}/indexing.cpp
3536
${CMAKE_CURRENT_SOURCE_DIR}/kernel_utils.cu
Lines changed: 184 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,184 @@
1+
// Copyright © 2025 Apple Inc.
2+
3+
#pragma once
4+
5+
#include "mlx/backend/cuda/device/utils.cuh"
6+
7+
namespace mlx::core::cu {
8+
9+
__device__ __forceinline__ void hadamard_radix_m(float* x);
10+
11+
template <int N>
12+
struct Pow2Log2 {
13+
static_assert(
14+
(N > 0) && ((N & (N - 1)) == 0),
15+
"N must be a positive power of two.");
16+
static constexpr int value = 1 + Pow2Log2<N / 2>::value;
17+
};
18+
19+
template <>
20+
struct Pow2Log2<1> {
21+
static constexpr int value = 0;
22+
};
23+
24+
template <int R>
25+
__device__ __forceinline__ void hadamard_radix_pow2(float* x) {
26+
constexpr int kLogR = Pow2Log2<R>::value;
27+
int h = 1;
28+
#pragma unroll
29+
for (int s = 0; s < kLogR; ++s) {
30+
#pragma unroll
31+
for (int i = 0; i < R / 2; ++i) {
32+
int k = i & (h - 1);
33+
int j = ((i - k) << 1) + k;
34+
float a = x[j];
35+
float b = x[j + h];
36+
x[j] = a + b;
37+
x[j + h] = a - b;
38+
}
39+
h <<= 1;
40+
}
41+
}
42+
43+
template <typename T, int N, int max_radix, int read_width, int stride = 1>
44+
__global__ void
45+
hadamard_n(const T* in, T* out, float scale, long long num_transforms) {
46+
constexpr int kNumThreads = N / max_radix;
47+
constexpr int kLogN = Pow2Log2<N>::value;
48+
constexpr int kLogR = Pow2Log2<max_radix>::value;
49+
constexpr int kNumSteps = kLogN / kLogR;
50+
constexpr int kLogFinal = kLogN % kLogR;
51+
constexpr int kFinalRadix = 1 << kLogFinal;
52+
53+
if (threadIdx.x >= kNumThreads) {
54+
return;
55+
}
56+
57+
__shared__ T buf[N];
58+
int i = threadIdx.x;
59+
60+
for (long long transform = blockIdx.x; transform < num_transforms;
61+
transform += gridDim.x) {
62+
long long base = (transform / stride) * static_cast<long long>(N) * stride +
63+
(transform % stride);
64+
65+
if constexpr (stride == 1) {
66+
#pragma unroll
67+
for (int j = 0; j < max_radix / read_width; ++j) {
68+
int index = j * read_width * kNumThreads + i * read_width;
69+
#pragma unroll
70+
for (int r = 0; r < read_width; ++r) {
71+
buf[index + r] = in[base + index + r];
72+
}
73+
}
74+
} else {
75+
#pragma unroll
76+
for (int j = 0; j < max_radix; ++j) {
77+
buf[j * kNumThreads + i] = in[base + (j * kNumThreads + i) * stride];
78+
}
79+
}
80+
__syncthreads();
81+
82+
float x[max_radix];
83+
int h = 1;
84+
85+
#pragma unroll
86+
for (int s = 0; s < kNumSteps; ++s) {
87+
int k = i & (h - 1);
88+
int j = ((i - k) << kLogR) + k;
89+
90+
#pragma unroll
91+
for (int r = 0; r < max_radix; ++r) {
92+
x[r] = static_cast<float>(buf[j + h * r]);
93+
}
94+
95+
hadamard_radix_pow2<max_radix>(x);
96+
97+
#pragma unroll
98+
for (int r = 0; r < max_radix; ++r) {
99+
buf[j + h * r] = static_cast<T>(x[r]);
100+
}
101+
102+
h <<= kLogR;
103+
__syncthreads();
104+
}
105+
106+
if constexpr (kFinalRadix > 1) {
107+
#pragma unroll
108+
for (int t = 0; t < max_radix / kFinalRadix; ++t) {
109+
int index = i + t * kNumThreads;
110+
int k = index & (h - 1);
111+
int j = ((index - k) << kLogFinal) + k;
112+
#pragma unroll
113+
for (int r = 0; r < kFinalRadix; ++r) {
114+
x[r] = static_cast<float>(buf[j + h * r]);
115+
}
116+
117+
hadamard_radix_pow2<kFinalRadix>(x);
118+
119+
#pragma unroll
120+
for (int r = 0; r < kFinalRadix; ++r) {
121+
buf[j + h * r] = static_cast<T>(x[r]);
122+
}
123+
}
124+
__syncthreads();
125+
}
126+
127+
if constexpr (stride == 1) {
128+
#pragma unroll
129+
for (int j = 0; j < max_radix / read_width; ++j) {
130+
int index = j * read_width * kNumThreads + i * read_width;
131+
#pragma unroll
132+
for (int r = 0; r < read_width; ++r) {
133+
float val = static_cast<float>(buf[index + r]);
134+
out[base + index + r] = static_cast<T>(val * scale);
135+
}
136+
}
137+
} else {
138+
#pragma unroll
139+
for (int j = 0; j < max_radix; ++j) {
140+
out[base + (j * kNumThreads + i) * stride] = buf[j * kNumThreads + i];
141+
}
142+
}
143+
144+
__syncthreads();
145+
}
146+
}
147+
148+
template <typename T, int N, int M, int read_width>
149+
__global__ void
150+
hadamard_m(const T* in, T* out, float scale, long long num_tasks) {
151+
constexpr int kTasksPerBatch = N / read_width;
152+
153+
for (long long task = blockIdx.x * blockDim.x + threadIdx.x; task < num_tasks;
154+
task += blockDim.x * gridDim.x) {
155+
long long i = task % kTasksPerBatch;
156+
long long batch = task / kTasksPerBatch;
157+
long long base = batch * static_cast<long long>(M) * N;
158+
159+
float x[read_width][M];
160+
#pragma unroll
161+
for (int c = 0; c < M; ++c) {
162+
#pragma unroll
163+
for (int r = 0; r < read_width; ++r) {
164+
x[r][c] = static_cast<float>(in[base + c * N + i * read_width + r]);
165+
}
166+
}
167+
168+
#pragma unroll
169+
for (int r = 0; r < read_width; ++r) {
170+
hadamard_radix_m(x[r]);
171+
}
172+
173+
#pragma unroll
174+
for (int c = 0; c < M; ++c) {
175+
#pragma unroll
176+
for (int r = 0; r < read_width; ++r) {
177+
out[base + c * N + i * read_width + r] =
178+
static_cast<T>(x[r][c] * scale);
179+
}
180+
}
181+
}
182+
}
183+
184+
} // namespace mlx::core::cu

0 commit comments

Comments
 (0)