Profiling in PyTorch (Part 2): From nn.Linear to a Fused MLP

Part 2 of Hugging Face's PyTorch profiling series shows how to read torch.profiler traces from nn.Linear through compiled and hand-tuned MLP kernels, and what fuse decisions actually improve.

Thursday June 11, 2026 Source: huggingface.co
TL;DR — Quick Answer

In part 2 of the Profiling in PyTorch series, Hugging Face developers trace a single nn.Linear up to a fused GeGLU MLP on an NVIDIA A100. A linear with bias is already one cuBLAS GEMM because the bias is folded into the addmm epilogue, so torch.compile has nothing to fuse there. An eager GeGLU MLP runs five GPU kernels per forward — three GEMMs plus GeLU and a multiply — while compile fuses the GeLU, mul, and reshape into one Triton kernel, and hand-tuned Liger kernels from the Hugging Face Hub deliver the same fusion without Dynamo or recompilation.

Key Takeaways

Profiling in PyTorch (Part 2): From nn.Linear to a Fused MLP — AI news article illustration

This is the second post in Hugging Face’s Profiling in PyTorch series, which reads profiler traces to drive optimization. Where Part 1 used torch.add(torch.matmul(x, w), b), Part 2 climbs one rung: replace that hand-written pair with an nn.Linear, then stack three of them with an activation between to form an MLP block.

From matmul-add to nn.Linear

nn.Linear wraps the same multiplication and addition profiled in Part 1. A zoom into the trace shows an aten::t (transpose) before aten::addmm, but that transpose never launches a GPU kernel — it only rewrites tensor metadata as a view.

The bias, meanwhile, is never a separate add kernel: it is folded into the matrix multiplication using an epilogue, a small computation the GEMM runs just before writing results back to memory.

Why torch.compile Barely Moves a Single Linear

Compiling the single linear’s forward changes almost nothing on the GPU: the same cuBLAS GEMM kernel runs. Compile simply turns the aten::t view bookkeeping into direct aten::addmm calls with precomputed strides.

The kernel name — dominated by its _tn_ layout descriptor — is byte-for-byte identical in both runs; compile needs more than one op before it can fuse anything.

Profiling a GeGLU MLP

The authors profile a feed-forward network with a GeGLU activation, the pattern used heavily in practice. Forming an expectation before opening the trace is the core habit of the series — per forward, the GPU runs exactly five kernels.

OpKernel
gate_projampere_bf16_s16816gemm (128x128 tile)
up_projampere_bf16_s16816gemm (128x128 tile)
geluelementwise GeLU kernel
h * uelementwise multiply kernel
down_projampere_bf16_s16816gemm (128x256 tile)

Each GEMM is roughly 38.7 GFLOP; down_proj runs about 10 percent faster because a different tile gives better data reuse.

What torch.compile Actually Fuses

Compiling the MLP collapses the GeLU, the multiply, and a reshape into one fused Triton kernel — triton_poi_fused__unsafe_view_gelu_mul_0 — leaving the three GEMMs untouched. The win: the 50 MB [8192, 3072] intermediate stays in registers instead of round-tripping through HBM.

Hand-Tuned Liger Kernels

The third setup swaps in LigerGEGLUMLP from the Hugging Face Hub via the kernels library, which downloads a pre-built, version-pinned package. The Liger kernel runs the same fusion as a single Triton kernel with hardware-tuned launch parameters, and with none of compile’s Dynamo guards or recompilation risk. Measured at 92.8 µs versus Inductor’s shape-specialized 89.4 µs, the trade for a generic kernel is robustness to changing shapes.

The Takeaway: Guess First, Then Look

The real comparison is not slow human kernels versus fast compiled ones; it is a fast generic kernel versus one specialized for a single input shape. And the habit worth keeping is the one practiced on every trace: state what you expect the profiler to show, open it, and treat any mismatch as the most interesting thing on the screen.

Frequently Asked Questions

Why is there no separate add kernel in an nn.Linear trace?

The bias addition is folded into the matrix multiplication kernel as an epilogue. nn.Linear dispatches aten::addmm(bias, x, weight), so the cuBLAS GEMM kernel writes out x @ w.T + bias in one pass.

Does torch.compile speed up a single linear layer?

Barely. For a single GEMM-with-bias kernel, compile only removes CPU dispatch overhead such as the aten::t view work; the GPU runs the identical cuBLAS GEMM kernel.

What does torch.compile fuse in an MLP?

In a GeGLU MLP, torch.compile collapses the GeLU, the multiply, and a reshape into a single fused Triton kernel, keeping the intermediate tensor in registers instead of round-tripping through HBM.

Why use hand-tuned Liger kernels instead of torch.compile?

Liger kernels bake in the same fusion as a hand-written Triton kernel and ship pre-built and version-pinned from the Hugging Face Hub, avoiding Dynamo guards, compile latency, and per-shape recompilation while running any shape.

This article is based on the official announcement from huggingface.co . Read the original for full technical details.

Related Articles

Back to all news