Adding Torch Custom Ops in cuDNN Frontend
Best practices for wrapping cuDNN graph ops as PyTorch custom ops with minimal CPU overhead.
File location
Custom-op implementations live with their owning operation family, for example
normalization ops under python/cudnn/ops/norm/, GEMM ops under
python/cudnn/gemm/ops/, and SDPA under python/cudnn/sdpa/.
Experimental APIs may be re-exported lazily from
python/cudnn/experimental/ops/__init__.py while they mature.
Registration: use torch.Library, NOT @torch.library.custom_op
@torch.library.custom_op adds ~36 us per-call overhead vs ~10 us for direct
torch.Library registration (PyTorch issue #139500). Always use the direct API:
User calls via: torch.ops.cudnn.my_op(x, w, eps) or wrap in a public function.
Graph caching
Build graphs once per unique (shape, stride, dtype, config) tuple. Cache as module-level dict:
Include device in the key — different GPUs may get different engine plans.
Graph building pattern
Use tensor_like() for automatic shape/stride/dtype inference from DLPack tensors:
Execute
graph.execute(uid_to_tensor, workspace, handle=handle) is the only form you need.
It already caches the backend’s operand order and reuses one pointer array, so the
sorted-pointer path is what runs underneath — there is nothing faster to reach for,
and hand-rolling it costs you the dynamic-shape overrides and the python-engine
dispatch that execute() handles.
Do NOT cache workspace tensors — they can race on different CUDA streams. Allocate per-call; PyTorch’s caching allocator recycles the allocation without hitting cudaMalloc (~2 us overhead). This matches PyTorch’s own conv/SDPA pattern.
UID management
Use an IntEnum for explicit UIDs — makes the code self-documenting and cache keys stable:
Set UIDs explicitly when building the graph:
Or use tensor_like which auto-assigns UIDs in insertion order (1, 2, 3, …).
Handle management
Cache one handle per device. Call set_stream once per call (costs ~5 us):
Future optimization: skip set_stream if stream hasn’t changed since last call.
Autograd pattern
Public API wrapper
Provide a user-friendly function that matches PyTorch conventions:
Performance checklist
- Use
torch.Library.define/impl, not@torch.library.custom_op - Cache the graph and workspace size in a bounded, thread-safe cache
- Use
graph.execute(uid_to_tensor, workspace, handle=handle)— it takes the sorted-pointer path internally - Use explicit UIDs (IntEnum) for stable cache keys
- Cache cuDNN handle per device
- Allocate workspace per-call (do NOT cache — stream safety)
- Consider
out=param for output tensors in performance-critical paths - Include device in cache key
- Test with tiny tensors to measure CPU overhead in isolation
CPU overhead budget (per call, Blackwell release build)
The ~34 us gap vs native ATen is torch.ops dispatcher + autograd overhead.
For inference without torch.compile, calling _my_op_impl directly saves ~25 us.