Day 87Module 9
in-technical-review

Custom PyTorch operators

A minimal wrapper can return the right numbers while breaking autograd, compilation, and stream ordering.

at::Tensor mm(at::Tensor a, at::Tensor b) {
    auto c = at::empty({a.size(0), b.size(1)}, a.options());
    launchWmmaGemm(halfPtr(a), halfPtr(b), c.data_ptr<float>(), ...);
    return c;
}

The output can match torch.matmul, yet c.sum().backward() reports no grad_fn. torch.compile splits the graph at the call, and a non-default stream can expose reads before the input-producing kernel finishes.

These failures occur because an ATen tensor carries a device, dtype, stride pattern, autograd history, and stream, while the raw kernel sees only a pointer. This lesson registers day 80's GEMM with an input contract, the caller's stream, a backward formula, and a fake implementation for compiler planning.

What a tensor carries that a pointer does not

A pytorch custom operator defines five properties that the wrapper above omits.

Device. A tensor knows it is on cuda:1. The CUDA runtime has a current device, set per host thread, and at::empty allocates on it, so an operator called on cuda:1 from a thread whose current device is 0 allocates its output on the wrong card.

at::cuda::CUDAGuard sets it for the scope and puts it back.

Dtype. The wrapper must reject unsupported inputs. This kernel takes FP16 in and writes FP32 out, because it accumulates in FP32 on the tensor cores and stores FP32 to preserve that mixed precision.

Passing FP32 input to data_ptr<at::Half>() would reinterpret its bytes as half-precision values.

Strides. x.t() swaps the shape and strides without copying data. Two tensors with the same shape can therefore have different layouts, and a row-major kernel misreads a transposed view without an error.

An operator must check contiguity or call .contiguous(), which makes a full copy.

The stream. PyTorch queues a model's work on a current stream and lets the host run ahead. A kernel launched on the default stream is not ordered after work on the caller's current stream.

A stream created inside the operator has the same problem.

at::cuda::getCurrentCUDAStream() supplies the required stream. This follows day 51's stream-ordering rule.

Autograd history. The output of a differentiable operation carries a grad_fn naming what produced it. A raw kernel call produces a tensor with none, so the graph ends there and every parameter upstream stops receiving gradients, even though training can continue.

Diagram: what each registration buys, and what fails without it. Three horizontal bands, one per stage, each showing the same call c = wmma_mm(a, b) on the left and three outcome boxes on the right, marked pass or fail: eager call, .backward(), torch.compile. Band 1, TORCH_LIBRARY plus a CUDA implementation: eager passes, the other two fail. Caption "1 of 3 registrations: the op runs and nothing else works." Band 2, plus register_fake: eager and torch.compile pass, .backward() still fails. Caption "2 of 3: the compiler can trace an operator it is not allowed to run." Band 3, plus register_autograd: all three pass. Caption "3 of 3." Alt text: "Registering a CUDA kernel with PyTorch takes three steps. One gets you an eager call. Two lets torch.compile keep a single graph. Only the third makes backward work, and the first two hide that."

The op is not the kernel

Registering the launch does not register shape inference or autograd. PyTorch asks for each part separately and reports a different failure for each missing registration.

The older recipe uses a torch.autograd.Function subclass with forward and backward static methods. test_op.py includes one to prove that gradcheck rejects a wrong formula.

A Function is opaque Python, so torch.compile breaks the graph at it, FakeTensor cannot infer the output shape without running the kernel, and the dispatcher has no per-device entry to route. The current documentation says to prefer the registration API for that reason ("Prefer this over directly using Python torch.autograd.Function", https://docs.pytorch.org/tutorials/advanced/cpp_custom_ops.html , checked 2026-09-01).

What the wrapper has to promise

Full program in code/day87-pytorch-op/, seven files listed with their jobs in the README.

The contract rejects invalid inputs. Wrong dtype, mixed devices, a non-contiguous input, a shape the 128 by 128 by 16 tile does not divide: each is a TORCH_CHECK raising a Python exception that names the file and line.

Rejecting non-contiguous input avoids a hidden full copy of the weight matrix. The caller can make that copy explicitly with .contiguous() and see it in a profile.

The kernel comes from day 80 without further tuning. It uses the same 128 by 128 block tile, same WMMA fragments, same FP32 accumulator. M and N are separate now, and there is no cuBLAS baseline in the file because PyTorch provides the comparison.

The backward uses two torch.matmul calls, following day 39's rule to use a library operation when it fits.

The extension neither exits nor synchronizes. Both .cu files carry CUDA_CHECK byte-identical to day 5's and the extension does not use it: std::exit would stop the interpreter without a traceback, so the launcher returns cudaGetLastError() and the binding raises. The missing cudaDeviceSynchronize() is deliberate the same way, since an operator that synchronizes makes every call a device-wide barrier.

The binding handles the device, allocation, and stream together:

    // The guard sets the current device to the tensors' device for the
    // rest of this scope and puts it back on the way out. Without it an
    // operator called on cuda:1 allocates its output on cuda:0.
    const at::cuda::CUDAGuard guard(a.device());
    at::Tensor c = at::empty({m, n}, a.options().dtype(at::kFloat));

    // The current stream, not a new one. PyTorch orders the model's work
    // on this stream; a stream created here would run beside it with no
    // dependency on the tensors it reads.
    const cudaError_t status = launchWmmaGemm(
        reinterpret_cast<const __half*>(a.data_ptr<at::Half>()),
        reinterpret_cast<const __half*>(b.data_ptr<at::Half>()),
        c.data_ptr<float>(), static_cast<size_t>(m), static_cast<size_t>(n),
        static_cast<size_t>(k), at::cuda::getCurrentCUDAStream());
    TORCH_CHECK(status == cudaSuccess,
                "day87::wmma_mm launch failed: ", cudaGetErrorString(status));

The schema string is the operator's public type, and PyTorch parses it to decide what autograd, the dispatcher and the compiler may assume:

TORCH_LIBRARY(day87, m) {
    m.def("wmma_mm(Tensor a, Tensor b) -> Tensor");
}

TORCH_LIBRARY_IMPL(day87, CUDA, m) {
    m.impl("wmma_mm", &wmmaMmCuda);
}

TORCH_LIBRARY_IMPL(day87, CPU, m) {
    m.impl("wmma_mm", &wmmaMmCpu);
}

The other two registrations are Python. The fake implementation derives the output's shape and dtype from the inputs' and reads no data, because it runs on tensors that have none. It is what keeps torch.compile in one graph.

@torch.library.register_fake("day87::wmma_mm")
def _(a, b):
    torch._check(a.dim() == 2 and b.dim() == 2)
    torch._check(a.shape[1] == b.shape[0])
    torch._check(a.dtype == b.dtype)
    torch._check(a.device == b.device)
    return a.new_empty((a.shape[0], b.shape[1]),
                       dtype=accumulate_dtype(a.dtype))

Then the backward. One output element is c[i][j], the sum over p of a[i][p]*b[p][j], so the loss gradient at a[i][p] is the sum over j of grad[i][j]*b[p][j], which is grad @ b^T, and the gradient at b[p][j] is a^T @ grad:

def _backward(ctx, grad):
    a, b = ctx.saved_tensors
    grad_a, grad_b = None, None
    if ctx.needs_input_grad[0]:
        wide_b = b.transpose(0, 1).to(grad.dtype)
        grad_a = torch.matmul(grad, wide_b).to(b.dtype)
    if ctx.needs_input_grad[1]:
        wide_a = a.transpose(0, 1).to(grad.dtype)
        grad_b = torch.matmul(wide_a, grad).to(a.dtype)
    return grad_a, grad_b

The dtype swap in those two lines is deliberate. setup_context saves a tensor only when its partner wants a gradient, so a is None on exactly the branch that computes grad_a, and the contract guarantees the two share a dtype anyway.

Note. The CPU implementation is not there for CPU users. It is there so gradcheck has a forward it can differentiate numerically: FP16's spacing at 1.0 is 2^-10, about 0.001, so a central difference in half precision is mostly rounding. The CPU path takes float64, so the formula is checked on a laptop and the kernel on the card.

Results

Verified on a Tesla T4, driver 580.173.02, CUDA 12.6 and torch 2.13.0+cu126. The standalone gate, extension build and full GPU harness all exited 0.

A separate CUDA_VISIBLE_DEVICES= run passed the CPU opcheck and both gradcheck rows, skipped all CUDA checks as designed and exited 0. The provenance capture records the installed torch distribution, its CUDA build, architecture list and literal nvcc --version output.

check what it compares pass or fail
contract four calls that must raise pass
forward, 256x512x384 kernel against a float64 CPU reference pass, 8.0% of bound
forward, 512x1024x128 kernel against a float64 CPU reference pass, 10.5% of bound
opcheck, cpu f64 and cuda f16 schema, fake impl, autograd pass
gradcheck the formula against a numerical jacobian pass
gradcheck rejects a grad_a scaled by 1.01 pass

The transcript checks four claims.

  1. Held. torch.cuda.get_arch_list() contained sm_75, and torch.version.cuda printed 12.6, matching the toolkit.
  2. Mixed. The 256x512x384 case used 8.0 percent of its bound, but the 512x1024x128 case used 10.5 percent. Both passed comfortably, but the prediction that both would stay below ten percent was narrowly refuted.
  3. Held. gradcheck passed and the deliberately wrong formula raised GradcheckError both with the T4 visible and under CUDA_VISIBLE_DEVICES=.
  4. Held. With no device visible, the contract, forward and CUDA opcheck rows skipped while CPU opcheck and both gradcheck rows passed; the process exited 0.

No timing appears here and none is planned. Day 88 profiles this operator inside a model and decides whether it earned its place.

Run it yourself

Use compute capability 7.5 or higher. Follow the README's exact environment: Python 3.12, then torch==2.13.0 from the cu126 index instead of PyPI. The default wheel uses CUDA 13.0.3, while the extension uses the installed nvcc; their CUDA major versions must match.

Day 44's mismatched build linked against libcublas.so.13 and could not run in its CUDA 12 environment.

Compiler Explorer cannot build this example because it needs two translation units, PyTorch headers, and libtorch.

Exercise

Add a second operator, day87::wmma_mm_bias, that folds a bias vector into the epilogue instead of leaving it to a separate elementwise kernel. Extend the schema, add the bias in the kernel's store loop, register a fake implementation and a backward formula for all three inputs, and say in one sentence what the bias gradient is and why.

Time: 30 to 45 minutes. Submit: your diff, the harness output for the new operator, and the sentence naming grad_bias.

Check: test_op.py already grades it. Its sixth check looks for torch.ops.day87.wmma_mm_bias, reports skipped, no such operator registered yet until you register one, then asks it the same three questions in float64 on the CPU and prints which failed: a forward that disagrees with a @ b + bias, a wrong grad_bias under gradcheck with the offending jacobian entries, or a fake implementation with the wrong shape under opcheck.

Hint 1

The kernel adds the bias once per output element and reads each column value many times. Compare this broadcast with the gradients derived above, then identify the axis that the backward pass must sum.

Hint 2

c[i][j] = (a @ b)[i][j] + bias[j]. Differentiate that with respect to bias[j] and count how many output elements bias[j] appears in. The answer to the exercise is that count's worth of terms, collapsed.

Solution

grad_bias = grad.sum(dim=0). bias[j] appears in every row of column j with partial derivative 1, so the chain rule sums the incoming gradient down the batch axis. The kernel change is one line in the store loop, and setup_context never saves the bias at all, because its gradient does not depend on its value.

A forward broadcast becomes a backward sum over the broadcast axes. The same rule applies to bias, scale, and layer-normalization epilogues.

Pitfalls

Training runs, but the loss does not move. The operator has no registered backward, so its output has no grad_fn and the graph ends there: parameters upstream keep their initial values while everything downstream trains normally. print(out.grad_fn) before you check anything else.

The result is wrong only when the caller transposes. A weight stored as (out_features, in_features) and used as w.t() is a stride swap and not a copy, and a kernel reading row-major from a column-major view gets a different matrix without an error. This operator refuses the call; a library operator calls .contiguous() and charges you a copy you did not see.

Your extension imports and then fails on an undefined symbol. The .so was linked against a different libtorch than the interpreter loaded, which happens the moment two PyTorch versions share a machine. Rebuild from the environment that will import it, and keep the build output: the linker line names the libtorch it used.

The kernel launches everywhere except the machine you shipped it to. CUDAExtension compiles for the build machine's card when TORCH_CUDA_ARCH_LIST is unset, so the .so carries one cubin and no PTX. That is day 69's gencode argument under a different name, and the fix is the same: name every architecture you support, or append +PTX to the last one so a newer card can JIT it.

torch.compile gets slower after you add your operator. Without a fake implementation the compiler cannot reason about the output without running the kernel, so it splits the graph there and loses every fusion that crossed it. Your operator is unchanged and everything around it got worse.

gradcheck passes and the formula is wrong. It differentiates your own forward, so a forward and a backward wrong in matching ways agree perfectly. Feed it a formula you know is broken and watch it fail, which is what the harness's fifth check is for.

Go deeper

Next

Day 88 puts this operator inside a model and profiles it down to the kernel. The profile shows whether the custom operator improves the workload. Day 90 then audits all four capstones one by one.