Richael
← Back to blog

2026-08-30

Cross-Entropy: The Benchmark That Nearly Fooled Me

7th kernel: cross-entropy loss, forward and backward fused into one kernel. The gradient is just softmax(logits) - onehot(target), which depends on nothing but the logits — so you can compute it during the forward pass and write it straight back over the input buffer. PyTorch has to allocate a second [batch, vocab] tensor for the gradient. This doesn't.

First numbers were 2.5x faster and 2.5x less memory, and I nearly stopped there. My baseline was F.cross_entropy(x.float(), targets), and that .float() quietly makes a full fp32 copy of the logits my kernel never pays for. Against an honest fp16 baseline it's 1.51x faster at vocab 131072 (15.9ms vs 24.0ms) and 1.67x less peak memory — exactly 1.67x at every vocab size, because PyTorch keeps five copies of the logits alive through fwd+bwd and this keeps three. Worth catching, since F.cross_entropy is a real fused kernel: this is the first week I'm not measuring against unfused eager PyTorch.

The catch: at vocab 32k, 74.3% of the gradient values it writes are exactly zero. Each one is about 1/(vocab x batch) ≈ 6e-8, and fp16's smallest subnormal is 5.96e-8, so the tail underflows on the way into the buffer. GradScaler can't save it either — that scale arrives in backward(), after the kernel has already stored them.

Next up is probably fused linear + cross-entropy — chunking the lm_head projection into the loss so the [batch, vocab] logits never exist at all.

Two line charts comparing Triton fused cross-entropy against torch F.cross_entropy on a T4 across vocab sizes 4096 to 131072. Left panel, forward plus backward time: Triton reaches about 15.9ms at vocab 131072 versus PyTorch's 24.0ms. Right panel, peak memory: Triton stays consistently below PyTorch, ending at about 1611MB versus 2684MB.