PRIMER: Writing Custom CUDA Operators for PyTorch Models Core Requirements - Replace operators with a custom CUDA implementation, typically inside a torch.autograd.Function or as a direct extension function. - No try/except or fallback logic — let assertions crash. - Output format is critical: the entire answer must be raw Python source, starting with the first character and continuing until the end. If the instruction says “output only the code,” the response must be solely the code block—no introductory or trailing text and no Markdown fences. - All methods must be fully implemented, with no placeholders. Token Budget & Code-First Strategy - Code output must come immediately. Spending the budget on reasoning before the code can lead to truncation, leaving no answer at all. - Prefer a minimal, correct kernel, such as one thread per output element or a serial scan, to remain within the token limit. - For inference-only tasks, skip backward computation by raising NotImplementedError. - Target roughly 60–80 lines of model code; simplify kernels that grow substantially larger. Symbol Visibility Before Binding (Critical) Every function referenced by m.def must be known at the point of registration. The extension binding code is a plain C++ translation unit and does not permit linking later to unresolved symbols. Two safe patterns are: 1. Monolithic source (recommended): define all functions before the PYBIND11_MODULE block inside a single CUDA source string and leave cpp_sources=[]. 2. Multi-file: declare the wrapper function in a header included by the binding file. Its definition in a separate .cu file must exactly match the declaration. The host function signature must be compatible with torch::wrap_pybind_function, conceptually a std::function accepting the specified argument types and returning a torch::Tensor. Using &my_func directly is simpler and preferred. Monolithic Template (Safe) cuda_src = r''' #include __global__ void my_kernel(...) { ... } void my_op(torch::Tensor x, torch::Tensor y) { ... /* launch kernel */ } PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("my_op", &my_op, "doc"); } ''' ext = load_inline(name="my_ext", cpp_sources=[], cuda_sources=[cuda_src], verbose=False) General Kernel Design Patterns - One thread per output for convolution, pooling, and similar operations: launch one thread per output element, loop over kernel and channel dimensions, and use __ldg for read-only data. - Scan-type operations such as cumulative sums: default to one thread per independent slice, looping sequentially over the scan dimension. This is simple, token-efficient, and reliable. - Reverse cumulative sums: loop from right to left rather than composing flip, cumulative sum, and another flip. - Multidimensional indexing: parenthesize expressions aggressively to avoid compilation failures caused by missing parentheses. Implementing Scan Operations - Treat the tensor as (num_vectors, L), moving the target dimension to the final position when necessary. - Launch num_vectors threads, each performing a serial inclusive scan. - For a reverse cumulative sum, iterate from right to left. - Let the host wrapper move the target dimension to the end, launch on a two-dimensional view, and permute the result back. Extension Building with load_inline - Exactly one PYBIND11_MODULE must appear, with TORCH_EXTENSION_NAME as its first argument. - The name supplied to load_inline must match the module name represented by the macro. - Use extra_cflags=["-O3"] and extra_cuda_cflags=["-O3"]. - Compile once and cache the module at the class level to avoid recompilation for every instance. Lazy Compilation (Class-Level Caching) class ModelNew(nn.Module): _ext = None def __init__(self, dim): if ModelNew._ext is None: ModelNew._ext = load_inline( name="op", cpp_sources=[], cuda_sources=[cuda_src], verbose=False) self.ext = ModelNew._ext Testing Builds Locally Verify that a small extension compiles and executes before integrating it: ext = load_inline(name="test", cpp_sources=[], cuda_sources=[cuda_src], verbose=True) x = torch.randn(1, 1, 8, 8, 8).cuda() ext.my_op(x, ...) Common Failure Modes Symptom | Root cause and fix SyntaxError at the start of the compiled string | Markdown fences or leading text were included; output pure Python source. No code output; token budget exhausted | The model spent the budget on reasoning; emit code immediately. redefinition of PyInit_xxx | Multiple PYBIND11_MODULE blocks; retain exactly one. my_func was not declared | The function is not visible before m.def; define it first or include its declaration. TypeError: SupportsFloat | The host function expects raw pointers; accept torch::Tensor and extract pointers internally. Wrong output or segmentation fault | Non-contiguous strides were ignored; call .contiguous() or handle strides explicitly. CUDA expected a semicolon | A complex index expression has unbalanced parentheses; introduce intermediate variables and balance the expression. Extension-name mismatch | The supplied name and binding module differ; use PYBIND11_MODULE(TORCH_EXTENSION_NAME, ...). Verification Checklist - Output is pure Python source, without fences or commentary when these are forbidden. - Code is emitted immediately without prolonged reasoning. - Exactly one PYBIND11_MODULE is present, and every referenced host function is visible beforehand. - Host functions accept PyTorch-compatible types rather than raw pointers. - No fallback logic appears in the optimized implementation. - Tensor contiguity is checked or handled. - Reverse scans traverse in reverse rather than composing flips. - Multidimensional index expressions are balanced and tested. - The extension is tested on a small tensor before a full benchmark run.