Repository navigation
Conversation
am17an
left a comment
There was a problem hiding this comment.
Thanks, nice change. Tested on RTX Pro 6000, results are all positive or neutral.
ORippler
left a comment
There was a problem hiding this comment.
@IMbackK did you benchmark both wider loads and keeping elements in registers independently? Wider loads are not really expected to help unless you run into an IPC issue on AMD - the memory access pattern is fully coalesced. This kernel is by definition latency-limited (moving 43 kB inputs is not going to saturate any recent HW BW imo, and on CUDA we launch just 1/ at most very few CTAs meaning we cannot saturate memory BW).
The above were generated on gemma4-26b on master, not this PR.
If vpt/keeping things in registers alone is sufficient, it would simplify the PR and reduce templating.
Other observations:
rms_norm_mul_rope_f32/other kernels with the same pattern could benefit from this as well most likely.- For the invocation I opened, we actually could hide latency on CUDA a bit by emitting wider loads (2816 cols / 1024 threads = 2.5 loads per thread on master, and the compiler decided to do the classic 1/4/8/16 optimization based on provided
block_size, so we have a loop and pay 3x the latency of which would be needed). But based on Aman's numbers and previous PRs/attempts E2E performance gains are negligible for NVGPUs (#20520)
| // The vec path reads 16B per lane, so the base pointer and every row offset have to be 16B aligned. | ||
| static bool rms_norm_f32_vec_aligned( | ||
| const void * ptr, const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, | ||
| const int64_t nrows, const int64_t nchannels, const int64_t nsamples) { | ||
| return ((uintptr_t) ptr & 15) == 0 | ||
| && (nrows <= 1 || stride_row % 4 == 0) | ||
| && (nchannels <= 1 || stride_channel % 4 == 0) | ||
| && (nsamples <= 1 || stride_sample % 4 == 0); | ||
| } |
There was a problem hiding this comment.
Why can't we use ggml_cuda_is_aligned from common.cuh here? It exists purely to check base alignment of poitner and rows at the ggml_tensor level
| #define RMS_NORM_LAUNCH(BLOCK_SIZE, VPT, VEC) \ | ||
| do { \ | ||
| const dim3 block_dims(BLOCK_SIZE, 1, 1); \ | ||
| const ggml_cuda_kernel_launch_params launch_params{ blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float) : 0, stream }; \ |
There was a problem hiding this comment.
I'm not an AMD guy but WARP_SIZE resolves to 32, and I was guided to use an IHV-agnostic helper in the past. Not sure if extending block_reduce and this kernel to use full warps on AMD HW may be an alternative path
There was a problem hiding this comment.
dosent matter, it wastes half the wave but is not compute bound anyhow
| static_assert(!do_scale || !do_multiply, "fusing scale is not supported with multiplying"); | ||
|
|
||
| // dst is written at row * ncols, i.e. always at a multiple of ncols, so only its base needs checking here | ||
| bool vec = ncols % 4 == 0 && ncols > 0 && ((uintptr_t) dst & 15) == 0 |
There was a problem hiding this comment.
might make sense to use braces to separate dst and x checks that are both happening here
| if (n4 <= 32) { | ||
| RMS_NORM_LAUNCH(32, 1, true); | ||
| } else if (n4 <= 128) { |
There was a problem hiding this comment.
Scheduling < 4 full warps tends to be an anti-pattern in cuda, as it often reduces occupancy
There was a problem hiding this comment.
i dont think this branch is actually practically used in real models so i dont think its worth anz complexity to optimize it
| float4 r = make_float4(scale*xs[k].x, scale*xs[k].y, scale*xs[k].z, scale*xs[k].w); | ||
| if constexpr (do_multiply) { | ||
| const float4 m = mul4[i]; | ||
| r.x *= m.x; r.y *= m.y; r.z *= m.z; r.w *= m.w; |
There was a problem hiding this comment.
| r.x *= m.x; r.y *= m.y; r.z *= m.z; r.w *= m.w; | |
| r.x *= m.x; | |
| r.y *= m.y; | |
| r.z *= m.z; | |
| r.w *= m.w; |
I find it more legible uncollapsed
Nope i was actually after wider loads and the x in registers just came as an aside. anyhow looks like that x in register actually has the dominant effect but the wider loads help too, numbers from test-backend-ops perf cases:
using half the waves in the reduction should not matter here but the n4 < 32 case is indeed pretty bad. |
|
There are parallel efforts to fix issues with limited CUDA grid sizes: #28175 . The changes are largely orthogonal to each other though. |
Overview
While profiling gemma4 on mi100 at batch 1 decode i noticed that the rms_norm kernels where takeing a large amount of time, relative to mem bandwidth, so i gave float4 loads a try. On the way i noticed that x fits into registers here, which might help on cdna/gcn too since its caches are puny.
Additional information
Benchmarks (benefits mostly on CDNA, RDNA more neutral)
GPU-f980f88ef31631a2
GPU-86308d5dff4ce29e
GPU-575dfb9c7cb709fd
Requirements