Skip to content

Batch norm updates and fused kernel - #1689

Draft
elliottslaughter wants to merge 6 commits into
flexflow:masterfrom
elliottslaughter:batch-norm
Draft

elliottslaughter wants to merge 6 commits into
flexflow:masterfrom
elliottslaughter:batch-norm

Conversation

@elliottslaughter

@elliottslaughter elliottslaughter commented Oct 1, 2026 •

Copy link
Copy Markdown
Collaborator

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 Reviewable

elliottslaughter and others added 6 commits September 30, 2026 16:28
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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant