Batch norm updates and fused kernel - #1689
Draft
elliottslaughter wants to merge 6 commits into
Draft
elliottslaughter wants to merge 6 commits into
elliottslaughter wants to merge 6 commits into
Conversation
The batch_norm portion of a0890b3 ("Update kernels for conv_2d, batch_norm, pool_2d, concat, reshape, split, batch_matmul, transpose and upsample."), plus the create_4d_accessor_w_with_contents fix its GPU test needs. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
The element_unary portion of a0890b3 ("Update kernels for conv_2d, batch_norm, pool_2d, concat, reshape, split, batch_matmul, transpose and upsample."). Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
A batch norm followed by an elementwise activation writes its output only for the activation to read it straight back and write it again, and because the backward pass then needs that value, it is read a third and fourth time there. On YOLOv10x those values come to 3145 MiB per iteration. Replaces BatchNormAttrs::relu, which was never implemented -- the forward kernel silently ignored it while the backward kernel applied it, so the kernels asserted it off -- with an optional Activation that the kernels do apply. The builders already took an activation argument to set the old flag from, so they now pass it through instead, and no longer refuse anything but relu. The forward kernel normalises and activates in registers, so the value between the two is never written down. The backward kernel then reconstructs it from the input and the saved statistics, at the price of a multiply-add rather than a pass over memory, which is what lets the forward kernel get away with not writing it. Measured across every batch-norm site in YOLOv10x: forward 13.70 -> 9.38 ms, backward 22.79 -> 13.80 ms. One block per channel, 512 threads. 512 beat 256 (by 32% forward, 21% backward) and 1024 (8%, 5%): below it there are not enough threads in flight to hide the memory latency, above it the block-wide reductions cost more than the occupancy buys. Splitting each channel across several blocks, with the same memory traffic but full occupancy, was tried and came out 14% slower. Statistics come out of Welford's method rather than E[x^2] - E[x]^2, whose cancellation loses most of the significant digits when the mean is large next to the standard deviation, and can return a negative variance and hence a NaN inverse standard deviation. Reconstructing the activation's input needs the shift, so beta joins the backward kernel's signature. The new test checks the fused kernels against the two operators they replace rather than against a table of numbers, since what matters is that folding them together did not change the answer. Which activations a batch norm can apply, and in which modes, lives in op-attrs so that the pass deciding to write one into the attrs and the kernel asked to run it cannot disagree. Sliced onto master without the change to have backward kernels overwrite their gradients, so the fused backward kernel accumulates into its gradients like the cuDNN path does, and its test starts every gradient at zero. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
The fused kernels put one block on each channel, which is all the parallelism
the statistics reduction can have but more than the normalization that follows
it needs. At YOLOv10x's 80-channel sites that left 80 blocks on 82 SMs: a
single wave, sixteen warps per SM, and 58-60% of the bandwidth the same access
pattern reaches elsewhere in the network, against 88-126% where the channel
count is larger.
Two changes, each chosen per call site from the extents:
- below two blocks per SM, split the pass in two, so the normalization runs
over a grid sized to the tensor rather than to the channel count. It is two
thirds of the traffic and needs none of the reduction's structure.
- walk four floats at a time where the spatial size leaves every thread work.
Below that the vector walk divides the work by four and loses.
Batch norm goes from 28.95 to 21.01 ms per iteration, and the iteration from
153.09 to 146.68 ms (4.2%), measured with evaluation/benchmark_flexflow.sh on
the deploy build. The validation suite passes all four stages.
Sliced onto master without the change to have backward kernels overwrite
their gradients, so these kernels accumulate into their gradients too.
That means the split backward pass can no longer hand its per-channel sums
from the stats kernel to the apply kernel through gamma_grad and beta_grad,
so they go through a gradSums scratch buffer in BatchNormPerDeviceState,
allocated alongside the saved statistics.
A new test runs the fused kernels at sizes that reach the vectorized and
split paths, which the existing one is too small for, starting both the
fused and separate paths from the same non-zero gradients.
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Slices out the portions of #1676 for BatchNorm, along with ElementUnary fixes for SiLU and a fused kernel BatchNorm+SiLU along with the op-attrs infrastructure to express the fused kernels.
Mainly Claude-generated, still waiting on me to review the diff. Should at least pass tests.
This change is