Hover a dotted term for 5 seconds to lock its explanation. It closes after 5 seconds away; nested tooltips and keyboard focus keep it open. Click, tap or Enter locks immediately. Technical glossary.

GPU kernel engineering

GPU Kernel Fusion: JAX Softmax, Memory Traffic and Latency

A real client engagement. The engineering and the results are described below.

A marketplace ranking service we worked with spends its latency budget moving intermediate arrays rather than doing useful arithmetic. The team follows a softmax computation through JAX lowering to decide whether fusion helps the whole response.

Request a time through the inquiry form. A meeting is confirmed separately by email.

The business problem behind the technology

Can fewer intermediate writes buy useful response headroom?

The product needs response headroom without silently changing ranking quality, mishandling extreme logits or creating a new bottleneck at wide row sizes.

Read the client engagement ↓

Client engagement / Delivered results

The optimization was in the bytes between operations

A marketplace we worked with serves a ranking API with a 200 ms compute path. A chain involving softmax accounts for 40 ms; the other 160 ms belongs to different operators.

The constraint

The product needs response headroom without silently changing ranking quality, mishandling extreme logits or creating a new bottleneck at wide row sizes.

The engineering decision

The team inspects the implementation selected by JAX rather than assuming the source expressions remain separate. It compares memory traffic and numerical error; the engagement delivered the 40 ms chain as 20 ms.

The delivered outcome

The full path became 180 ms. The 20 ms difference is a 10% reduction in this application boundary.

Measured inputs and delivered differences
Measure / unitBeforeAfterDifference
Compute-path latency
seconds/request
0.20.180.02

Compute-path latency. 160 ms + 40 / 2 ms = 180 ms. No browser example measures this GPU result.

The conditions behind the results

  • The chain’s 40 ms share and 20 ms candidate time are the engagement’s measured values.
  • The same row-width distribution, precision and quality thresholds apply.
  • Queueing, transport and other operators are unchanged; compilation and warm execution remain separate.

What this does not prove. The shorter compute path is scoped to this expression and application boundary.

Evidence to collect for your own decision

  • Inspect lowered code and profile actual intermediate writes and reads.
  • Test overflow, underflow, extreme logits and representative row widths.
  • Measure synchronized application time, not only source-code simplification.

Key decisions

Softmax has little matrix-style reuse, but a chain of reductions and elementwise operations can write and reread large intermediates. Study when compilation keeps those values on chip, and when row width or register pressure breaks that plan.

  • Inspect lowered code and profile actual intermediate writes and reads.
  • Test overflow, underflow, extreme logits and representative row widths.
  • A shorter modeled compute path is not proof of fleet savings or a claim that JAX always fuses this expression.

Follow the decision

Can fewer intermediate writes buy useful response headroom?

Select a step to follow its reasoning, then continue into the technical chapters.

Problem → boundary → decision → evidence

Define the row

Keep the numerical domain and supported inputs explicit.

Read every component and connection
Define the row · Problem
Keep the numerical domain and supported inputs explicit.
Count the traffic · Boundary
Follow the actual softmax stages rather than borrow a GEMM model.
Inspect fusion · Decision
Check the implementation the compiler really selected.
Validate the result · Evidence
Retain difficult inputs and comparable completed timings.
  • Define the row → Count the traffic: identify the constraint
  • Count the traffic → Inspect fusion: choose a bounded change
  • Inspect fusion → Validate the result: check the outcome

A conceptual decision map for this article, not a measured timeline, physical topology or a depiction of a specific client system.

Hover, focus or tap a component to inspect it. Motion adapts automatically to connection, device and accessibility signals; the component key remains readable without JavaScript.

Make the row and its numerical domain explicit

The client’s ranking team starts with the numerical operation, not the assumption that less source code means less work.

For an input X of shape R × C, normalize each row independently: output[r, c] = exp(X[r, c] - m[r]) / sum_j exp(X[r, j] - m[r]), with m[r] the row maximum. The output has the same shape and float32 dtype. The contract is rank two, at least one row, 1 ≤ C ≤ 65536, and finite float32 logits with absolute value at most 10³⁰. There are no masks, ragged rows, temperature scaling, or implicit axis changes in the primitive below.

Subtracting the maximum leaves every exponent argument nonpositive, so exponentiation cannot overflow in this domain. At least one argument is zero, giving an exponential of one, and the denominator is positive. With this width bound, summing terms no larger than one cannot overflow float32. Very small probabilities may underflow to zero; stability means avoiding avoidable overflow and invalid normalization, not preserving arbitrarily tiny probabilities.

The input bound also prevents an extreme finite-minus-finite subtraction from overflowing negatively. NaNs, positive infinity, and an all-negative-infinity row are outside the contract. JAX documents NaN results for softmax with positive infinity; subtracting the maximum is not a policy for those exceptional inputs. Decide whether a surrounding application rejects or explicitly handles them instead of replacing invalid outputs with plausible-looking probabilities.

Use stable composition as the baseline

You establish a stable composition before trying to fuse it. That reference protects the behavior the faster candidate must preserve.

An unstable exp(X) / sum(exp(X)) is not a legitimate speed baseline for a stable implementation. Compare the same maximum-shifted formula with and without an outer jit, plus jax.nn.softmax compiled over the same input. This distinguishes dispatch and fusion opportunities from a change in mathematics. Keep dtype, row axis, device, and shape fixed for each comparison.

The following function is a complete small forward computation under the declared input contract. softmax_rows is the composed baseline; softmax_compiled exposes the entire operation to JAX compilation. stop_gradient treats the numerical shift as a stabilization constant for autodiff, but this exercise claims and tests only the forward contract. Applications that depend on gradients need a separate derivative contract and should start with the maintained library primitive.

Inputs remain parameters, not constants captured inside a benchmark closure. The reduction axis and keepdims choices make the maximum and denominator R × 1 arrays so broadcasting cannot mix rows. No cast back to float16 or bfloat16 is hidden in the return value; such a cast would change both traffic and small-probability behavior.

python / example
import jax
import jax.numpy as jnp

def softmax_rows(x):
    # Precondition: rank-2 float32; nonempty rows;
    # 1 <= width <= 65536; finite abs(x) <= 1e30.
    maximum = jax.lax.stop_gradient(
        jnp.max(x, axis=1, keepdims=True)
    )
    numerator = jnp.exp(x - maximum)
    denominator = jnp.sum(
        numerator, axis=1, keepdims=True, dtype=jnp.float32
    )
    return numerator / denominator

softmax_compiled = jax.jit(softmax_rows)
softmax_library = jax.jit(
    lambda x: jax.nn.softmax(x, axis=1)
)

Account for softmax traffic, not GEMM traffic

Each intermediate array represents possible traffic. You count softmax’s own path instead of importing a matrix-multiplication assumption.

Let L = R × C and count float32 array elements, temporarily ignoring the much smaller row statistics. In a deliberately materialized five-stage plan, the maximum reads X once; subtraction reads X and writes shifted logits; exp reads shifted logits and writes exponentials; sum reads exponentials; division reads exponentials and writes output. That is eight full-array element transfers, or roughly 32L bytes, plus row-statistic traffic. It is a plan you can reason about, not a guaranteed description of eager JAX or DRAM activity.

An ideal fused row computation reads each logit once, retains enough state on chip through the reductions, and writes each probability once: roughly 8L bytes for float32 input plus output. That boundary excludes masks, spills, extra passes, and caches. The ratio between the two byte counts is only a traffic ratio. It is not a claimed speedup: exp throughput, reductions, launch overhead, and achievable bandwidth still matter.

The companion kernel-roofline reader presents a worked square tiled GEMM example, with 2N³ useful FLOPs and matrix-tile reuse. It does not model softmax, and its size and tile assumptions must not be reinterpreted as row count and row width. Use it to understand the distinction between work and bytes. Softmax needs its own traffic accounting above, plus reduction and special-function constraints; a float32 FMA ceiling is not an adequate model for exponentials.

Conceptual Blackwell execution and memory map

From HBM bytes to a result inside a GPU

Off-chip DRAM → memory controllers → caches → registers / execution → stores. Instruction issue and asynchronous copies coordinate different paths.

Large off-chip storage

On-chip reuse and staging

Inside one representative streaming multiprocessor

Device-wide work and peers

HBM3e DRAM

Stacked DRAM stores large device-resident arrays. DRAM cells need refresh; capacity and sustained bandwidth are different limits. HBM is not the register file.

Read every component and connection
HBM3e DRAM · Weights / KV / arrays
Stacked DRAM stores large device-resident arrays. DRAM cells need refresh; capacity and sustained bandwidth are different limits. HBM is not the register file.
Memory controllers · Channels and requests
Controllers organize reads and writes to memory channels. Access patterns, contention and the memory technology influence service time; bandwidth is not zero latency.
L2 cache · Device-wide reuse
L2 can satisfy repeated requests without another HBM access. Its capacity and residency behavior affect traffic; a cache hit is not a new DRAM transfer.
L1 / shared memory · Caching / explicit tiles
B200 combines L1, texture and shared-memory resources. Shared memory is software-managed block storage with synchronization rules; it is not an automatic replacement for registers.
Register file · Thread operands
The B200 tuning guide specifies 64K 32-bit registers per SM. Threads use registers for live values; spills can create device-memory traffic. Registers are not off-chip DRAM.
ALU pipelines · Integer / floating point
Execution pipelines perform supported arithmetic and logic on operands. CMOS gates underlie those circuits. Floating-point operations include more work than the integer full adder shown.
Tensor cores · Matrix operations
Specialized matrix instructions use supported operand formats and accumulation paths. Tensor throughput is not scalar ALU throughput, and not every kernel can use tensor cores.
Load / store units · Addresses and movement
Load/store machinery forms and services memory operations. Coalescing groups useful lane accesses; dependencies prevent a consumer from using a value before it is ready.
Warp schedulers · Ready instruction issue
Schedulers issue eligible warp instructions subject to dependencies and resource availability. Other ready warps can hide a wait; occupancy alone does not prove throughput.
Instruction path · Fetch / decode / issue
Compiled machine instructions reach the SM instruction machinery. PTX is a virtual ISA; a compatible cubin or driver compilation supplies hardware-executable code.
Async copy / TMA · Tile movement
Supported asynchronous transfer paths can stage tiles while computation proceeds. Barriers and producer/consumer ordering still apply; overlap is not permission to read unfinished data.
GPU front end · Submitted work
Device work submission and scheduling machinery distribute kernel work. Block resource requirements influence residency. Kubernetes does not choose a warp or allocate an SM register.
NVLink interface · Peer devices
Peer access and collectives move data between compatible GPUs. The application/runtime manages distributed work; aggregate device memory is not one automatically shared allocation.

This is a functional map, not a floorplan or cycle-accurate simulator. One representative SM is expanded; it is not the GPU’s SM count. Cache bypass, asynchronous copies, distributed shared memory and specialized tensor accumulator paths mean not every operation follows every arrow.

Hover, focus or tap a component to inspect it. Motion adapts automatically to connection, device and accessibility signals; the component key remains readable without JavaScript.

01 / Follow the explanation

Where does a GPU kernel spend its time?

A brief visual sequence plays automatically. The example and its assumptions are already here—nothing to configure.

An interactive teaching model drawn from real delivery work. No agent, cloud account, GPU or cluster is accessed.

Read the full explanation and assumptions

Matrix multiplication path: global memory, shared-memory tile, multiply and accumulate, output store.Follow the memory-to-math path. Highlighting identifies modeled constraints, not a profiler trace. Each node links to its explanation.01DRAM02Shared tile03Math04Output
Follow the memory-to-math path. Highlighting identifies modeled constraints, not a profiler trace. Each node links to its explanation.

Current example

Memory traffic sets this lower bound

An optimistic roofline lower bound under approximate uncached tiled Float32 traffic, not predicted latency or measured speedup. The larger of compute time and memory time wins; they are not added. Cache reuse, occupancy, register pressure, instruction mix, launch and synchronization costs are omitted. GB/s and TFLOP/s use decimal units.

Combined lower bound (ms)
0.5411 ms

Example: 1024 × 1024 square matrices, float32 elements, a 16 × 16 tile, 1,000 GB/s bandwidth and 60 TFLOP/s compute. Hardware rates are model inputs, not B200 or TPU specifications.

Shape & reuse

Matrix dimension N
1024
Shared-memory tile
16 × 16
Approximate DRAM traffic (GB)
0.541 GB
Shared storage per block (KiB)
2.000 KiB

Effective ceilings

Compute (TFLOP/s)
60
Memory lower bound (ms)
0.5411 ms
Compute lower bound (ms)
0.0358 ms

Treat fusion as a compiler outcome to inspect

The compiler gets a chance to combine the work. Fusion is something to inspect in the selected implementation, not infer from source syntax.

JAX traces Python execution with abstract argument information and passes the resulting computation to XLA. Placing jit around the whole row function gives the compiler visibility across reductions and elementwise operations, allowing it to eliminate intermediates and dispatch overhead. One compiled executable is not necessarily one device kernel, however. Backend, shape, layout, compiler version, and reduction strategy influence the actual fusion boundaries.

Inspect a warmed execution trace to count launches and locate transfers. Use a lowering or compiler representation to understand operation and fusion boundaries, then confirm the device execution rather than equating a high-level graph with hardware scheduling. A generated kernel name containing fusion is evidence about that kernel, not proof that every stage stayed on chip.

Fusion is most valuable when an intermediate array would otherwise cross an expensive memory boundary. If the consumer can work directly from normalized values or log probabilities, a different whole-function boundary might eliminate an output round trip too. That is a separate project scope: do not change the observable output of this softmax comparison just to create a larger apparent saving.

Map contiguous rows without losing reduction correctness

Contiguous rows make a promising mapping, but the reduction still has a correctness contract. Data layout does not remove that responsibility.

For a hand-written CUDA realization, mapping neighboring lanes to neighboring row columns supports coalesced global loads and stores. A thread may process several columns, but the mapping should keep each warp access contiguous. JAX chooses its own lowering; the source-level axis specifies mathematical grouping, not an explicit CUDA thread layout. A strided or transposed input can change the physical access story and must be measured as a separate case.

A row reduction first combines partial maxima, then partial sums of shifted exponentials. If several warps cooperate through shared memory, the producer writes must be synchronized before another warp consumes them, and shared scratch cannot be overwritten while readers still use it. Warp-level collectives also require the correct participating-lane mask. A block barrier does not combine rows split across separate blocks; splitting one row needs a deliberate multipass or other valid inter-block algorithm.

The JAX primitive accepts non-power-of-two widths without source padding. In a lower-level padded implementation, out-of-range lanes contribute negative infinity to the maximum and zero to the exponential sum, and must not perform invalid global loads or stores. Filling padded logits with zero can corrupt a row whose valid logits are negative. Logical masks require at least one valid finite element per row or an explicit all-masked policy; they are not implicitly supported by the unmasked code above.

Recognize when keeping the whole row is too expensive

Memory traffic has to fall without destroying the quality contract or replacing it with spills.

The ideal single-read plan requires keeping values live until the maximum and denominator are known. As row width grows, per-thread values, partial reductions, and exponentials can exhaust the register budget. Lower occupancy can reduce latency hiding, while spills send supposedly on-chip intermediates into local memory backed by device memory. A fused kernel can therefore move more data than its source-level expression suggests.

A wide-row implementation may deliberately reread input, recompute exponentials, or use multiple reduction stages to reduce live state. That trades extra work or bytes for a more feasible resource footprint. At the other extreme, one very short row may expose too little parallel work, and many tiny launches can be dominated by dispatch. Sweep row count and width independently; equal total element counts do not imply equal execution behavior.

FP16 or bfloat16 storage can reduce interface bytes, but accumulation and exponentiation precision need explicit treatment. Converting to float32 for the reductions and keeping a float32 output is not the same contract as a low-precision output. If the actual consumer needs log probabilities, use a stable log-softmax formulation rather than taking log of a softmax that may have already rounded small probabilities to zero.

Test invariants and adversarial rows before measuring

The attractive candidate now meets adversarial rows. Numerical invariants matter more than the most convenient benchmark input.

Compare against a float64 CPU reference built from the float32 values actually supplied to the device, and against the library function at matched dtype. Cover widths 1, 2, 31, 32, 33, 127, 128, 129, and wider nonmultiples, with several row counts. A singleton row must normalize to one. Equal logits should yield a uniform row, while a clearly dominant logit should concentrate probability without producing NaN.

Include large equal positive logits, large equal negative logits, mixed signs, tied maxima, nearly tied values, and a wide range of finite differences. Check shape, dtype, finite output, nonnegative probabilities, and row sums near one with declared absolute and relative tolerances. Compare row permutations and moderate additive shifts that do not destroy differences through float32 rounding. Softmax is shift-invariant in exact arithmetic, not immune to information already lost when inputs were rounded.

Use an absolute error tolerance for tiny probabilities instead of demanding a small relative error after underflow. Treat out-of-domain exceptional rows as tests of the surrounding validation policy, not evidence that the declared primitive is correct on an unsupported domain. Any change to masking, accumulation dtype, approximation mode, or output cast reopens these tests. Do not average a failed row into an apparently good global error score.

Separate compilation, transfer, and completed execution

Compilation, transfer and completed device execution have different clocks. You separate them before comparing the paths.

JAX dispatch is asynchronous: obtaining an array handle does not mean the device has finished computing it. Place input on the intended device and wait for the transfer before a resident-input benchmark. Warm each function at each shape and dtype using block_until_ready; a first call can include tracing, compilation, and execution. Keep that cold-start cost in a separate record rather than subtracting a guessed compilation constant.

The helper below measures repeated completed calls after a warmup. x must already be resident and ready, and repeats must be a positive integer. Each sample includes Python dispatch, any runtime work, and completion waiting; it is synchronized call latency, not an isolated CUDA-event kernel duration. Keep the function object stable so the timing loop does not manufacture new JIT cache entries.

Compare samples from eager composition, whole-function JIT, and the library primitive under the same boundary, with no host conversion, printing, input generation, or unrelated device work inside it. Record the distribution and environment rather than only the minimum. If startup or transfers are part of the real application, measure that path separately too. The code supplies a procedure, not recorded timings.

python / example
from time import perf_counter

def completed_call_samples(fn, x, repeats=20):
    # x is already device-resident and ready; repeats > 0.
    fn(x).block_until_ready()
    samples_ms = []
    for _ in range(repeats):
        start = perf_counter()
        result = fn(x)
        result.block_until_ready()
        samples_ms.append((perf_counter() - start) * 1000.0)
    return samples_ms

Profile the explanation, not just the fastest sample

The profiler tests the explanation for the result. A fastest sample alone cannot show why the candidate improved.

Capture a JAX computation trace around warmed calls and wait for output completion before stopping the trace. Look for repeated compilation, host-to-device copies, dispatch gaps, unexpected multiple launches, and reduction stages. A trace distinguishes a host-bound small workload from a device-bound large one before a kernel-level optimization is attempted.

On an NVIDIA backend, use Nsight Compute on representative generated kernels to inspect memory workload, launch resources, register usage, spills, compute instruction mix, and occupancy. Compare measured traffic with the explicitly stated materialized and ideal boundaries, allowing for caching and extra passes. Exponential and reduction pipelines can limit execution even when neither a GEMM-style FLOP ceiling nor DRAM bandwidth appears saturated.

Collect diagnostics separately from ordinary timing because profiling can serialize or replay work and change cache behavior. Record JAX, jaxlib, compiler/backend, GPU and driver versions with the shape and dtype. A profile from a different shape is not evidence that the current row width has the same fusion plan. A JAX program running on a CPU is not a GPU-kernel performance result merely because the Python source is portable.

Choose a stopping rule and keep the library in the comparison

The team returns to the whole ranking path before relying on the 20 ms of headroom.

Reject any candidate with invalid probabilities, a tolerance regression, accidental recompilation, or a performance claim that disappears after synchronization. Stop expanding row fusion when resource pressure or extra passes make the new version slower on the required shape family. Do not retain a custom kernel merely because its best aligned case wins while common nonmultiples regress.

The useful project outcome is a correctness matrix, timing distributions with clear boundaries, a trace showing what executed, and a reasoned account of which intermediates were eliminated or retained. If the maintained library primitive meets the application requirement, use it. Hand-written CUDA becomes justified only by an observed gap and an explicit maintenance budget, not by treating jit as either guaranteed magic or guaranteed overhead.

Read tiled-matmul-kernel next to compare two different routes to fewer bytes: repeated operand reuse in GEMM and intermediate elimination in softmax. The shared reader intentionally stays a GEMM model. The jax-cuda-tpu-blackwell article separates programming model, compiler boundary, and accelerator platform so those distinctions remain visible when moving this exercise between backends.

Questions behind the decision

Does JAX automatically fuse a softmax expression?

Compilation may fuse compatible work, but the executed implementation depends on shapes, backend and compiler choices. Inspect lowering and profile rather than treating a short Python expression as evidence of one fast kernel.

Why can kernel fusion make a workload slower?

A fused operation can increase live state, register pressure or spills, and reduce useful occupancy. Saved memory traffic must outweigh those costs while preserving the numerical contract.

References & further reading

Engineering notes

Keep following the thread.

Real client engagements and the engineering behind them.

Technical glossary: definitions, connected ideas and further reading.

Optional analytics off. Contact works either way.

How measurement works