Python-native cudnn.pygraph and pluggable execution backends
Python-native cudnn.pygraph and pluggable execution backends
What this is
cudnn.pygraph is a Python-native graph class: graph structure (nodes,
tensors, op parameters) lives in Python with full introspection, and execution
dispatches over python DSL engines and the cuDNN C++ backend alike. The C++
graph builder is internal (cudnn._pybind_module.backend_graph) and is reached
exclusively through lowering.
Why: python-DSL engines (CuTe-DSL / cuTile style GEMM and attention fusions)
need to see the graph to decide whether and how to run it. Previously that
required monkey-patching the pybind class and recording calls; now the graph
is natively introspectable and an engine is one file implementing
BaseEngine.
“Pluggable” means the library ships a table, not that callers hand engines
over at runtime. There is no registration call: engines/manifest.py is the
only way a python engine exists. An engine handed over at runtime could not be
ranked anyway — it declares no Capabilities, so nothing could enumerate its
configs or place it against the backend.
The dispatch tree
Architecture
Graph IR
graph_types.Tensor,nodes.Node,_pygraph.pygraph— an engine-agnostic op DAG. Input/output port names equal the C++ pybind kwarg names, everywhere.- Three declarative op mechanisms cover 100% of the C++ op surface:
_POINTWISE_TENSOR_ARGS(54 uniform pointwise ops;mode== method name),_STRUCTURED_OPS(25 ops: norms, reduction, block-scale, MoE, conv, structural — one table entry each: ports, attrs, outputs, shape-infer),_CAPTURED_OPS(6 SDPA variants, ~130 kwargs: generic capture over an explicit per-op schema carrying positional order, output-direction kwargs, and conditional outputs).matmulis explicit for positional ergonomics.
Backend contract (engines/)
BaseEngine:check_support(graph)(accept, or decline by raising),build_plan(graph, plan, ctx) → CompiledPlan(the expensive JIT step, once per graph/plan, cached on the graph),CompiledPlan.execute(graph, operands, ExecutionContext)with explicit handle/stream/workspace. Dynamic-shape overrides are a backend-path feature: a python plan is compiled for the shapes the graph declared, soexecute()refuses them rather than silently running a different problem. Simple eager engines implementexecute()only.
The variant pack is normalized once
graph.execute() converts whatever the caller passed — a torch tensor, a
DeviceView, any __dlpack__ / __cuda_array_interface__ producer, or a bare
device address — into a VariantPack at the top, and everything below reads that.
Do not add a branch on what the caller’s object is. This exists because
there used to be two such branches: the backend path accepted a bare address
(_native_var_pack’s if type(d) is int: return d) while an engine got the
object untouched and frost.buffers.probe refused it. One public call, two
answers, decided by which plan the heuristics happened to pick — which the
caller does not control.
VariantPack carries the caller-filled uids ascending and the operands
themselves, held in a C type (pygraph/variant_pack.cpp) as one DLTensor
each. That type both consumes __dlpack_c_exchange_api__ — the C function
table a producer publishes on its type, which is how one crossing reads the
whole pack — and implements it, so a kernel reads an operand through the same
fast path it has for a framework tensor. address is the pointer array
_execute_with_raw_ptrs takes.
The pack’s vocabulary distinguishes the position from the thing at it:
pack.index_of(tensor_or_uid) gives an operand’s POSITION, and
pack.operands(indices) turns positions into OperandBuffers — one caller
buffer described (pointer, shape, stride, dtype), non-owning. Resolve positions
once; ask for buffers per call.
The declaration is the contract; the buffer supplies pointer, bytes and
alignment. That is how the cuDNN backend has always read a variant pack — it
takes the pointer and nothing else — and callers program against it: FlashInfer
binds a 2-D matrix to a [1, m, k] tensor, a flat quantizer blob to a
[b, bs_m, bs_k] F8_128x4-reordered scale tensor, a 0-d scalar to (1, 1, 1).
So _normalize describes a slot from the graph when the buffer is a DENSE run
of other extents that covers the declared bytes (the 2-D matrix, the flat blob,
the 0-d scalar) — exactly what a bare address gets — and records the slot in
VariantPack.graph_described, so an engine that reads the pack answers the way
the backend does and one call cannot get two answers by plan selection. Two
kinds of buffer keep their own description: one with the DECLARED extents under
its own strides (a padded or transposed view of this very tensor — the strides
carry information, and an engine that reads the pack honours them, which is how
the linear-attention engines serve strided inputs), and one that is SMALLER
than the declaration (a packed THD buffer under a padded declaration, say): the
engine decides, and an engine that needs the declared extent refuses it at
execute naming the operand (the backend would read past the allocation). A bare
address lends the declaration outright — it carries no extent, so there is
nothing to compare against: the caller guarantees the allocation covers the
declared bytes and meets the engine’s alignment. Rules that follow:
- The declaration supplies the dtype too. A buffer re-described from the
declaration takes its dtype with the extents (a byte blob covering a bf16
output becomes that output), and a buffer of the declared extents whose
slots are as wide as the declaration’s is read AS the declared dtype
(FlashInfer binds packed fp4 data and the e4m3 scale blob as
uint8; the backend read a pointer and never knew). A buffer too small for the declaration, or whose slots are not as wide as the declaration’s, keeps its own dtype with its own description. override_shapes/override_stridesspeak cuDNN element units in the graph’s axis order, like every declaration. They are written INTO the slot (in the buffer’s axis order,_in_axis_order_of), never carried around it, so an engine honours them without knowing the concept exists.- The pack speaks storage slots. fp4 packs two elements per slot along the
unit-stride axis, so a declared or overridden fp4 geometry is converted with
_storage_geometry(extent halves, the other strides halve) before it lands in a slot; an odd unit-stride extent has no slot spelling and is refused. An engine’s per-slot packing factor (frost_gemm’skpack) therefore applies to every slot alike — this is what stoppedfrost_gemmfrom doubling K on an overridden fp4 operand. - The pack still carries the shape about to RUN: read the IR port for the
shape the plan was built for; read the pack for the shape this call runs
(
frost_gemmreads its M/N/K there, so one plan serves many problem sizes).
The rule is on every execute’s critical path, backend plans included, so the
declaration side (storage_geometry of every slot) is computed once per graph
into a native DeclaredLayout and the comparison runs in one crossing per pack
(VariantPackNative.describe_from). Measured at 128×256×128 bf16, host enqueue
per execute on SM100: backend plan 11.7 → 12.0 µs, frost_gemm 21.9 → 22.0 µs,
and a FlashInfer-shaped 2-D binding costs what a declared one does (12.1). The
first, per-operand Python form of the same rule was +7.5 µs on both paths and
+10 µs more for the 2-D binding — two crossings and a storage_geometry per
operand per call.
Two rules that are easy to break by accident:
- The pointer array is per call. Two threads may execute one graph concurrently with different buffers. A shared array hands each thread the other’s pointers, and each pointer in it is individually valid, so the failure is a wrong number rather than a raise.
- The operand order has exactly one source, never a union. The lowered
graph’s variant-pack template when the graph has one — only C++ sees every
user slot, since a tensor’s
ragged_offsetis an operand but hangs off theTensorrather than off a node port, and the slots the graph fills itself (pass-by-value scalars, slice replacement destinations, workspace modifications) must be excluded. The python IR only for the python-only ops that never lower. The two sides do not have to agree: each indexes the layout it was handed.
is_virtual on a python-only graph does not mean “the caller supplies
nothing” — it is a statement about the backend’s lowering, and a gdn graph marks
its own O virtual while the caller passes a buffer for it. So the layout there
is every wired port, and an unfilled slot is an optional port the caller did not
request.
CompiledPlan.takes_variant_pack is the migration flag. An engine that has not
set it still receives the caller’s {uid: buffer} map and reaches ports through
resolve_node_buffers; the flag and that function both go once the last engine
has moved. execute() builds the Tensors only for a plan that sets the flag —
measured, normalizing for a plan that will not read the result costs more than
it saves.
What a per-execute path costs
An engine owns its internals, and this section does not change that. It exists
because the default outcome is expensive: an engine that re-derives its per-call
facts lands around 40 µs of host time per execute, and for a single-kernel
op that is most of what the caller pays. The same kernel with those facts read
once is 20. Both numbers are frost_gemm at 256×256×128 bf16, host
enqueue, min over 25 reps of a 64-call burst from a drained queue.
The budget it has to fit in, all measured on SM100:
Do not read a per-call cost out of an nsys trace: CUPTI adds ~2.2 µs per traced API call, which is more than the call.
Split the facts by when they are decided. Operand roles and majors, packing
factors, alignment requirements, output shape rules, which outputs need a seed —
all fixed when the kernel compiled. M/N/K, strides and pointers arrive per call.
Read the first set into a table at build (gemm/frost/recipe.py is the worked
example) and let the call read the table. That alone is 44 → 35.
Then lower the table into one closure per plan, with its constants captured and the operand structure flattened into the loop headers, so the call does no attribute lookup and takes no branch the build already settled. That is 35 → 20. Two rules make it safe:
- The lowered path never raises, and it is the only path that runs. What it refuses it hands to a checker that reads the same table, names the rule and raises — it launches nothing. A graph the closure cannot serve at all is declined when the engine is asked to support it, so it goes to the backend rather than to a second executor.
- So a refusal is the answer, not a slower route. The set of calls the closure refuses should equal the set of illegal calls; a legal call it will not serve is a bug. Keeping a reference executor instead would buy a differential that catches divergence but never a misconception the two share — which is exactly how an axis-order bug survived one here. The tests that matter are against intended semantics and against the BACKEND, at the shapes where two encodings coincide.
A loop over a flat table gets almost all of it, so do not hand-unroll per
flavor. Measured three ways on the same plan and buffers: interpreting the
table 35.8, looping over it flattened 19.7, a hand-written straight line with the
structure unrolled 17.5. The loop is worth 45%; unrolling adds 12% and costs one
closure body per operand shape — six flavors, six bodies to keep in agreement.
One loop over arg_plan (the launch argument order as data) serves aux, extra
outputs, multi-GEMM and block scale at 22 µs each, down from 39–50. Source
codegen off the same table is how to buy the last 12% back later, for every
flavor at once rather than for the one that was worth hand-writing.
This is a pattern to copy, not a framework to import. Sharing the code across engines would couple their kernels’ ABIs, which is the thing engine autonomy buys; sharing the shape of the solution costs nothing.
Costs that are easy to miss, each measured:
- A
from x import yinside a per-call function: 1.1 µs. It was 65% of what_check_plan_devicecost. torch.Tensor.permute(): 1.4 µs per call, per operand.- Rebuilding a
{id(tensor): buffer}map to look operands back up, when the operand order was settled at build and a list index would do. - Recomputing a pure function of values every call.
tensor_alignment’s layout half is 1.5 µs and memoizes on(shape, stride, elem_bytes)— values, so there is nothing to invalidate; only the pointer half is per call. - Reading an operand through the exchange vtable is 0.08 µs against 1.5 for the python attribute walk. Framework neutrality is not what costs.
Measure from a drained queue, and sweep the burst size. A number that is flat in the burst size is host-bound; one that climbs with it is the device rate, and back-to-back timing reads the device rate whenever host and device are close.
- An engine does not propose its own plans. Which configs to try, in what
order, and where the backend’s entries belong is one comparison across every
candidate, and no engine can make it from the inside — it sees neither its
siblings nor the backend. That decision lives in
engines/heuristics.pyand the per-family hook it dispatches to (see Ranking and the one plan list). - Every engine has a stable
engine_idin one flat id space (engines/engine_ids.py: backend[0, 10_000), C++ OSS[10_000, 20_000), python[20_000, …)with aFAMILY_BLOCK-wide block per family) — reproducible pinning/autotune. Engines do not DECLARE their id: the manifest holds every slot andinstantiate()hands each factory the ids its engines are to use, so an engine cannot claim a number it was not given, the whole space is readable in one file, and any id decodes back to a family and a slot (manifest.engine_for_id) with nothing registered first — which is what letscreate_execution_plan(engine_id, knobs)replay an autotune result. - An engine declines a graph ONLY via
NotImplementedError,cudnn.cudnnGraphNotSupportedError, orImportError; anything else is an engine bug and propagates.ImportErrorcounts because lowering imports are deferred pastcheck_support()(see Import boundaries), so a missing optional dependency can only surface at build time — without it, a host lacking thecutedslextra would lose graphs the backend could have served. - The contract is proven end to end without a GPU in
test/python/test_dispatch.py, with stand-in engines injected through the manifest — the same path production uses. Those engines do no arithmetic: what dispatch is responsible for is reaching the engine and resolving the caller’s buffers, and checking a result againsttorch.matmulwould put a torch reference implementation of matmul inside a dispatch test. GdnCuTileEngineexecutes the single-nodegdnandgdn_bwdops (Gated DeltaNet linear attention) via the cuTile chunked kernels. Both ops are THD-only: token-packed[total_T, heads, dim]tensors with a requiredcu_seqlens(a dense batch is[0, T, 2T, ...]).gdn_bwdtakes the forward inputs plusdO(and optionallyd_final_state) and producesdQ/dK/dV/dG/dBeta(+d_initial_stateiffinitial_stateis given,d_a_logiff the node carriessafe_gateand ana_loginput, andd_dt_biasiff it carriessafe_gateand adt_biasinput). Both gate parameters are optional undersafe_gate: an absenta_logis unit amplitude (exp(a_log) = 1), an absentdt_biasis zero bias, and no zero tensor is materialized for either.dG/dBetaare in raw-logit space undersafe_gate/use_beta_sigmoid; the cumulative gate and intra-chunk WY matrix are recomputed inside the engine, so the graph contract carries no forward intermediates. Both are python-engine-only ops: they have no cuDNN backend lowering, so routing them to the backend entry raisescudnnGraphNotSupportedErrorat lowering. The kernels live incudnn.linear_attention.cutile.kernels.gdn; the torch custom opcudnn.linear_attention.ops.gated_delta_netis a thin adapter that builds and executes cachedgdn/gdn_bwdgraphs (the SDPA op pattern), so it inherits whatever engine the planner selects. The optionaluse_qk_l2normattribute asks the engine to L2-normalize the q/k rows;GdnFrostEngine(the SM100-SM103 and SM107 default, serving bothgdnandgdn_bwdon the FROST chunked kernels) serves it through a workspace helper kernel (normalized q/k copies + saved inverse norms, with the backward Jacobian projection applied in place after the head-group fold), and likewise servessafe_gate(in-kernel raw-logit gate transform, withd_a_log/d_dt_biasfor the given parameters produced by a deterministic reduction helper) anduse_beta_sigmoid; the cuTile engine remains the fallback for non-128 head dims. Thegate_domainattribute ("log", the default, or"linear") selects whethergisln(alpha)oralphaitself; the FROST GDN / KDA / GDP / GDN-2 engines serve"linear"(forward and backward,dGwith respect toalpha); the cuTile engines are log-only.
KdaFrostEngine/KdaCuTileEnginedo the same for the single-nodekda/kda_bwdops (Kimi Delta Attention). KDA is GDN with a per-key-channel decay: itsgis the log-space vector gate[total_T, HV, K](GDN’s is the scalar[total_T, HV]);betastays scalar. The FROST engine (cudnn.linear_attention.frost.kda_engine) is the forward default on SM100-SM103 and SM107; the node’suse_qk_l2normattribute (in-kernel L2-normalization of q/k — the KDA model’s feature map) passes through to the kernel (without the in-kernel norm the caller owns the q/k conditioning). It serveskda_bwdon the FROST backward kernel, regenerating the per-chunk state checkpoints with a recompute pass when the graph does not provide them; the cuTile engine (cudnn.linear_attention.cutile.kernels.kda) is the fallback slot. The torch op iscudnn.linear_attention.ops.kimi_delta_attention.- Gated DeltaNet v2 (
gdn2/gdn2_bwd) has channel-wise gates —g/beta[total_T, HO, K]plus a NEW per-value write gatew[total_T, HO, V]. GDN-2 has no cuTile engine;Gdn2FrostEngine(cudnn.linear_attention.frost.gdn2_engine, SM100-SM103 and SM107) is its only engine, passes theuse_qk_l2normattribute through to the kernel (likeKdaFrostEngine), and servesgdn2_bwdthe same way (checkpoint recompute when the series is absent); the op iscudnn.linear_attention.ops.gated_delta_net_v2. - Gated DeltaProduct (
gdp/gdp_bwd) appliesnum_householderbeta-gated Householder updates per token with one scalar decay per token: the GDN recurrence on an expanded sub-token timeline (gate on sub-token 0, readout on sub-tokenn - 1). The node carries q/g/O/dO/dQ/dG at real-token rows and k/v/beta/dK/dV/dBeta attotal_T * num_householderrows;num_householder == 1is exactlygdn.GdpFrostEngine(cudnn.linear_attention.frost.gdp_engine, SM100/SM103) is its only engine and runs the shared GDN kernels, except thed_v == 64backward forkkernel/gdp_bprop_v64_f16.py, which reads q/dO and writes dQ in the token domain. The forward reads q compact: the prefill re-reads each 64-token q block into every chunk of the block and runs its q-side work (readout, rescale, O drain) once per block, so no expanded q copy exists. In thed_v == 128backward, q and dO are zero-scattered into an expanded workspace copy and dQ is gathered back (frost/common/expand.py); the gate is read compact with the sub-token rows derived in registers, O and dG are stored compact in-kernel, andcu_seqlensis scaled bynat every read site.checkpoint_every_n_tokenscounts expanded sub-tokens (64 = the bwd-reusable chunk cadence; a multiple oflcm(64, n)puts every checkpoint on a real-token boundary).safe_gate,use_beta_sigmoidandallow_neg_eigval(beta as2 * sigmoid(x)) all pass through. The op iscudnn.linear_attention.ops.gated_delta_product. - The FROST engines are pure pass-through:
check_supportrequires the kernel-native dtypes (fp32/bf16/fp16 gates — io-dtypebeta/wfor GDN-2 — int32 or int64cu_seqlens, fp32-or-bf16 state ports with matching initial/final dtypes, state gradients at the state’s own dtype, and io-dtypedBeta/dWfor GDN-2) and execute hands the caller’s buffers straight to the kernels, carving any scratch it needs out of the explicit workspace as DLPack views. The cuTile engines follow the same buffer contract: outputs are written in place (the caller’s output buffers, required in the kernel-native dtypes, are planted under the pipelines’ terminal workspace names), and their chunk-index tables are built on device fromcu_seqlens, so execution stays sync-free. Buffers only need__cuda_array_interface__or__dlpack__; the torch custom ops do the dtype normalization on their side. - Sequence splitting. The main kernels run one persistent CTA per (sequence,
head), so a long sequence on few tiles leaves SMs idle.
choose_mode(frost/common/piece_chain.py) fixes the scheme at build from the declared shapes and the device’s SM count, identically for the forward and the backward, and it is not a node attribute.chaincuts the batch into one wave ofB * Punit-aligned pieces,P = min(num_sm // (B * HO), 16, total_chunks // (4 * B * unit))withunit = lcm(expand_num, checkpoint cadence)chunks (one wave of tiles, at most 16 pieces, every piece at least 4 units long), and it is chosen only whenP >= 3for a forward plan orP >= 2for a backward plan; every piece spansceil(total_chunks / (B * P))chunks whichever sequence it belongs to, so a sequence fillsceil(len / span)slots and an uneven batch walks the same critical path as an even one (when the per-sequence ceilings would overflow the wave the span is recomputed againstB * P - (B - 1)slots).warmup(the decay-warmup split-K offrost/common/split_k.py) serves the band where the chain has no room, anduncutruns one item per (sequence, head). Underbatch_invariantthe geometry comes from the length rule alone,P = clamp(ceil(total / 8192), 1, 16)slots per sequence, of which each fillsclamp(ceil(len / 8192), 1, P)on device,uncutwhentotal <= 8192(except for the summary ops, which keep a one-piece chain there so that the emitting state chain below still composes their tail), so a sequence’s outputs are bitwise the same alone and in any batch. The chain is exact. The chain prologue writes the piece-wisecu_pieces(real tokens, the slots flat in sequence order: sequencebowns slotsmain_rows[b] // HO .. main_rows[b + 1] // HO) and the work-item tables (every filled slot per head, plus an empty sequence’s slot 0 as a passthrough; the summary table, which the main ops’ summary launch walks, only for sequences with two or more filled slots, a one-piece sequence’sXbeing its seed; the summary ops summarize every filled piece over the main table, below); a summary launch produces every multi-piece sequence’s per-piece state from a zero seed (H) and transition (M) in fp32 (GDN’s summaries read the T pass, the beta-folded chunk factor ofkernel/gdn_tinv_f16.py; a GDN prefill or bprop builds the factor itself when it is the tiles’ only consumer and, for the prefill, the plan fills at least half of the SMs, and reads the tiles otherwise); an fp32 FMA chain composesX_{j+1} = X_j M_j + H_jfrominitial_state(a one-piece sequence’sXis its seed, copied); the seeded main kernel runs over the piece work items as independent sequences and the last filled piece of each sequence writesfinal_statein the state dtype. The backward mirrors it with M (with H and X when no series is passed back), a G summary from the bprop-summary kernel, the reverse chain seeded byd_final_state, the bprop over the pieces seeded by row 0 of each piece’s series, and piece 0 writesd_initial_state. A one-piece chain is bitwise the uncut run. Workspace grows in chain mode by the piece table (piece_table_layout), fp32 H, M and X slots[B * P, HO, V or K, K], and the T pass tiles (tinv_rows * HO * 8 KB) plus their tensor-map slot. The summary ops (*_summary,*_summary_bwd) use the same rule; their chain summarizes every filled piece over the main work-item table (an empty sequence’s passthrough item yieldsH = 0,M = I) and one emitting state chain composes them, its tail beingfinal_state(in reverse,d_initial_state) and its running producttransition. - In-place state update. The main ops’
overwrite_initial_stateattribute (fwd and bwd nodes of GDN, KDA, GDN-2 and GDP) lets one buffer serve asinitial_stateandfinal_state(in the backward, asd_final_stateandd_initial_state): the planner keeps the chain and the uncut schedule and never takes the split-K cut, so the CTA that reads a sequence’s incoming state (or outgoing gradient) at its first chunk is the one that writes the outgoing state (incoming gradient) at its last, after the read. The chain is safe as well: the state chain consumes the seed (the reverse chain the outgoing gradient) before the seeded main kernel writes the final state (piece 0 the initial-state gradient). The attribute requires theinitial_stateinput and thefinal_stateoutput (bwd:d_final_stateandd_initial_state); binding distinct buffers stays allowed. The torch ops expose it on the forward only (cudnn::<op>_fwd_overwrite_state, a mutating op without an autograd formula, inference use); autograd never lets a backward mutate an incoming gradient, so the backward alias is a graph-API contract for callers that own their gradient buffers. - Pool-addressed state. The main ops’ forward nodes take an optional int32
state_indicesinput of[N]row ids:initial_stateis then a pool[N_pool, HO, V, K]whose rowstate_indices[i]seeds sequencei, andfinal_state(the pool itself underoverwrite_initial_state, or a distinct buffer of the pool’s shape) receives that sequence’s final state at the same row. The pool may pad its slot stride to any 16-byte multiple, as serving stacks do; each row stays a dense[HO, V, K]block. The chain seeds its state chain through the table and the prefill writes through it; the split-K cut stays available when the two buffers are distinct. Forward only, not combined withcheckpoint_every_n_tokens, and served by the FROST engines only: the cuTile engines, the KDA CAKE engine and the Hopper KDA engine decline a graph withstate_indices. The torch ops routestate_indicesthroughcudnn::<op>_fwd_overwrite_state, so the caller’s pool is updated in place and comes back asfinal_state. - Context parallelism across devices. The state ops are the per-span
summaries of a two-tier scheme: every rank summarizes its span at once, the
ranks exchange the summaries, and every rank runs its span’s main op seeded
by the composition. In the stored domain (state buffers V-major, the buffer
holds
S^T) the forward summary*_summaryreturnsH(the span’s final state from a zero seed) andM_buf(the transition, holdingM^T), composed forward from the incoming state asX_{j+1} = X_j @ M_buf_j + H_j; the span forward then runs withinitial_state = X_j. The backward summary*_summary_bwdreturnsG(the incoming state gradient from a zero outgoing gradient) and, withoutput_transition=True, the transition in the backward’s orientation (M_buf^T), composed in reverse from the outgoing state gradient asdX_j = dX_{j+1} @ transition_j + G_j; the span backward then runs withinitial_state = X_janddX_{j+1}as the gradient on itsfinal_state. Seeds, summaries and compositions are fp32, and every summary is itself the intra-device chain above.
The manifest: classify, then let the engine decide
engines/manifest.py answers only what an engine cannot answer about itself
without being imported. Everything else is the engine’s own check_support().
- A family is a KIND OF GRAPH — roughly the backend’s operation-graph mode, at a granularity of our choosing — not a group of engines that ship together. Every graph belongs to exactly one family or to none, so engines within a family compete and engines across families never do.
- Classification is a lookup.
_ANCHOR_NODE_TO_FAMILYmaps the node types that NAME a family;family_for(graph)is a function, so “two families claimed this graph” is not a case that can arise. A graph naming two (a matmul and an sdpa together) belongs to neither and goes to the backend. - The table holds anchors, not an envelope. Node types absent from it —
POINTWISE,REDUCTION, anything added tomorrow — are ignored when classifying, somatmul + pointwiseis a gemm graph. Whether a family can serve the WHOLE graph is its analyzer’s judgment. A coarser copy of that judgment here is whatclosed_underwas, and it promised RESHAPE support nothing implemented. family_for(graph)is a pure property of the graph — nosm, no environment. What kind of graph something is cannot depend on which machine is asking. Availability is separate (EngineFamily.offered_ids).- What the manifest does NOT decide: architecture. An engine’s
Capabilitiesdeclares an arch RANGE (sm_lo/sm_hi,major*10 + minor), because an sm100 kernel serves the sm100 LINE — enumerating the members that exist today silently declines the ones that ship later. - Maturity is per engine (
EngineSlot.opt_in, gated byCUDNN_FRONTEND_ENABLE_FROST_ENGINES), so one implementation can graduate while a sibling matures. It lives in the manifest rather than on the engine class because the gate must answer without importing the engine. - A family may name a
heuristicshook — likeanalyzer, a("module", "callable")pair kept as strings so the coarse key stays import-free. It is handed the facts and the family’s offered ids and returns its proposals; aBACKENDmarker in that list says where the backend’s own block goes inside each mode block (ours first when absent). The SDPA-forward family decides that per measured shard (sdpa/fwd/placement.py). - A family may name a
validatorhook — the same import-free pair,validate_graph(graph) -> bool. When the manifest offers a python engine for the graph,validate()runs it instead of the eager C++ lowering: it applies the family’s version- and arch-agnostic semantic rules with the classic error types, and returns False (classic lowering) for a graph holding a node it does not cover. It exists because the eager lowering coupled a graph a python engine fully serves to the installed backend’s version and per-arch gates (issue #704); the backend’s own verdict is not lost, only deferred to planning (_finalize_backend_layoutrecords the decline,plan()raises it if no python engine proposes a plan either). A family without one validates classically.cudnn/_sdpa_validate.pyandcudnn/_gemm_validate.pyare the two today; both import only the IR. - There is no registration call. The manifest is the only way a python
engine exists, and
_candidate_engines()is the graph’s family and nothing else. An engine handed over at runtime could never be ranked anyway: it declares noCapabilities, so nothing can enumerate its configs or place it against the backend — it was an entry point into the plan list, not into the decision. Tests inject their fakes as a manifest family, so they reach dispatch the way real engines do.
Facts: one description per graph, shared
A family may name an analyzer — a ("module", "callable") pair, kept as
strings so matching stays import-free. Planning resolves it and attaches the
record to the graph; engines read that record back rather than parsing again.
- Attached after the freeze. Planning runs
_finalize_backend_layout()→_freeze()→_attach_facts(). Analyzing a graph that can still change means chasing every mutation point — the layout the backend infers, a dtype set between twovalidate()calls — and missing one leaves facts describing a graph that is gone._facts_for()memoises only a frozen graph, so there is no invalidation rule to get wrong. - Keyed by the analyzer itself, so the ranking (which resolves it from
EngineFamily.analyzer) and the engine (which passes the callable it already imports) reach ONE record with no name to keep in sync. Two families MAY share an analyzer; SDPA forward and backward do. - Facts describe, capabilities judge.
SdpaGraphFactsrecordshas_bias=Trueas a fact, never an error; each engine’sCapabilitiesrow does the rejecting inmismatch(). A shared parser that starts rejecting becomes an if-ladder that must know every kernel. - Framework-neutral vocabulary:
cudnn.data_type, nottorch.dtype; device fromcudnn.create_device_properties(), the backend’s own descriptor. Facts are what every engine of a family reads, so expressing them in one framework’s types would make dispatch require that framework.
Handle, stream, and device
There is no ambient device state in dispatch. Three separate things:
- Handle —
execute(..., handle=h)if given, else the graph’s own. The handle is what carries the stream. - Stream —
_resolve_stream(handle)iscudnn.get_stream(handle). With no handle it isNone, and kernel wrappers then use the default stream. A failed query on a SUPPLIED handle raises rather than silently falling back to another stream: running on the wrong stream is a correctness bug, not a degradation. So it is a fallback for no handle, never a fallback for a handle whose stream could not be read. - Device — the analyzer reads compute capability and SM count from
cudnn.create_device_properties(), the same serialisable descriptor the C++ deviceless-AoT path uses, rather than fromtorch.cuda.current_device(). Engines re-check the arch incheck_support()regardless.
Import boundaries
Deciding whether an engine COULD serve a graph must not cost the machinery that
would serve it — importing the CuTe DSL is ~1 s and 357 modules, and paying it
only to decline is why closed_under existed.
- Package
__init__s undercudnn/sdpaare lazy (PEP 562) andEngineSpec.lowerresolves its DSL adapter at build time, soimport cudnn.sdpa.graph_analyzercosts 9 ms and 2 modules rather than 1059 ms and 381. - The graph API pulls no framework at all: describing and validating a graph imports neither torch nor cutlass — the family validators included (they see only the IR).
- A missing or too-old DSL is a DECLINE at
check_support(), probed without executing the module (importlib.util.find_spec,importlib.metadata).CUTEDSL_MIN_VERSIONinfrost/buffers.pyis the floor these engines want;pyproject’s required dependency deliberately sits below it (>=4.6.2), since pinning that high would make cudnn-frontend incompatible with anything holding the DSL back. test/python/test_import_boundaries.pyholds all of this, in a fresh interpreter, measuring the delta against an empty one.
Ranking and the one plan list
create_execution_plans()gathers the parsed facts, eligible engine ids, and the backend’s(engine_id, knobs)entries.engines/heuristics.py::rankcalls the family’s hook with(kind, facts, offered). The hook returns its own proposals plus an optional internalcudnn.engines.heuristics.BACKENDmarker._assemble()expands that marker into the mode’s backend block, then strips mode annotations and deduplicates to form the finalgraph.planslist.- The backend’s entries arrive tagged with the mode that produced them.
_create_backend_plans()asks C++ one heuristic mode at a time and recordsget_execution_plan_count()after each, so a family can say “the backend’s mode-A entries ahead of ours, its fallbacks behind”. C++ appends each query to the same plan list, which is exactly what makes the boundaries readable — no C++ change was needed. heur_mode.OPENSOURCEis mode A without the backend’s recommendation: the python engines ARE the open-source implementation. Combine it to measure coverage —[OPENSOURCE, A, FALLBACK]tries every python config first and still has the backend behind it, so a graph that runs on a backend plan is one no python engine covers.- The list IS
graph.plans, and the classic at-index APIs (get_execution_plan_count(),get_plan_name_at_index(),build_plan_at_index(),execute_plan_at_index(),get_workspace_size_plan_at_index()) address it, so code that loops over the plan count picks up python engines with no change. - A backend entry carries the
cpp_indexit holds in the lowered graph’s own plan list, so building it is onebuild_plan_at_index. Backend engine sets are still never statically enumerated: they are discovered per graph at plan time, andbackend_plan_entries()returns[]when the backend declines the graph or is not installed — the backend participates in the ranking, it is not a hard dependency. - A plan’s identity is
(engine_id, knobs), never itscpp_index. The index only says where one backend query happened to put it, so deduplicating on it lets[A, A], or one config both modes return, through twice — and an autotuner would build and time the same config twice. - Whether a mode SUCCEEDED is tracked per call, not read off the plan spans. An OPENSOURCE query registers a C++ OSS candidate without adding a plan, so it contributes no span; judging by spans would rethrow a later mode’s failure and discard the delegate that successful query earned.
- The family positions the backend block with its marker; without a marker,
its proposals precede the backend. Ranked indices address the assembled list,
never the marker or an unexpanded family proposal.
BACKEND_HEURISTIC_ENGINE_IDnames one thing only: the delegating entrybackend_plan_entries()appends, where the backend picks among OSS candidates it never exposes as plans and which therefore cannot be enumerated. It carries no mode. It leads the BACKEND’s entries, becauseGraph::build_planstries it before its own engine configs — but it does NOT lead the family’s, because it is not a pure OSS entry: if the C++ OSS engine declines, that same call falls through to the native configs already enqueued. Ahead of the family’s OPENSOURCE block it would answer an OSS-coverage question with a native kernel. build_plans()walks the list from the selected index and takes the first entry that builds; a decline advances to the next.select_plan(i)pins, and a pinned decline raises instead of running something else.- Ranking policy has ONE home: the family’s
heuristicshook, replaceable per family. Which side leads is meant to be a measurement — a cell timed slower than the backend follows it — not a default. Any cell that has not been timed keeps the historical order, and the code says so where it is written. - A recommendation always names a concrete config. A knob field is
Noneonly where the capability row declares no domain for that axis.Nonenever means “engine, pick for me” — that reading is what let the same choice be made in two places, once in the ranking and once in the DSL adapter, and drift.sdpa/fwd/heuristics.py::_sm120_tilesis the worked example of a rule: it reads facts only, names its choice first, and puts the rest of the domain behind it for a caller that autotunes. To add one — write the function, list the cell in_TILE_RULE_CELLS, and put the measurement in the commit.
Knobs: one public vocabulary, and every python plan lists the knobs it will build
KnobType_t(include/cudnn_frontend/knobs.h, mirrored ascudnn.knob_type) is the ONE vocabulary for backend and python plans. Backend values0..32are frozen; python-only axes live in the band fromFRONTEND_KNOB_TYPE_BASE(1000:SCHED_POLICY,PACK_GQA,SPLIT_KV,PIPELINE_ARCH,MMA_TILE_*,CTA_GROUP,WARPS_*). Both bands are append-only; a frontend-band value never reaches the backend (convert_to_backend_knob_typerefuses it). Reuse a backend type when the meaning matches (TILE_M,TILE_CGA_M,SPLIT_K_SLC,SWAP_AB); mint a frontend one only for an axis the backend has no word for.- A python engine keeps whatever native knob object it likes in
PlanConfig.knobs(SdpaFwdKnobs,SdpaBwdKnobs,GemmKnobs) and converts at the public boundary throughBaseEngine.knobs_to_public/knobs_from_public.get_engine_and_knobs_at_indexalways returns a dict ({}when the plan has no axes, neverNone);create_execution_plan( engine_id, {cudnn.knob_type: int})replays it;get_plan_name_at_indexprintsengine[SPLIT_KV=2, TILE_M=128, ...], sorted by knob name. - Every python plan is listed WITH the knobs it will build. A family’s
recommend(kind, facts, offered)proposes completePlanConfig(engine_id, knobs)entries from facts alone; the engine builds what was listed. Families with a real tuning axis — SDPA forward, SDPA backward (the SM120 row’s per-head-dim default tile, the MXFP8 row’s sole point), frost GEMM (the tile config as knobs) — all do this, so a recorded(engine_id, knobs)pins the kernel across releases even when a default table moves. A family whose kernels take no tuning decision (linear attention) lists{}, which IS its complete record. A knob-less plan whose engine picks insidebuild_planis a bug: the record would replay a different kernel after the pick changes. - Knobs are performance-only: a plan computes the same function under any knob
value, so an autotuner may pick freely. Anything numerics-changing
(
softmax_precision) is an op attribute declared in the op spec’spython_only_attrs: never forwarded to C++, a SET value makes the node backend-unlowerable (serialize()andkey()refuse it), and it surfaces as a graph fact the capability rows gate on.
One kernel per layout class, not per shape (SDPA THD)
A compile() of an SDPA template used to pin batch, head extents and every
stride in its fakes, so FlashInfer’s serving shapes minted one 2.6 s kernel per
(b, qh, kh, strides): 126 forward kernels in FlashInfer’s cuDNN attention
tests. The kernels never needed that — _host reads B / QH / KH from
problem_size at run time and the THD token totals were already sym_int.
Under compile(dynamic_bhk=True) (THD only) the template rebinds b, qh,
kh (and a padded LSE’s s_max) to cute.sym_int() right after the cache
key is taken, keeps plain ints for the problem_size fake, and gives the
metadata / O-descriptor fakes fresh symbols (SymInt has no __add__); a
packed declared stride is passed as None so the compact fake derives it from
the dynamic extents, a padded LSE in a compact dim order passes that order
(lse_padded_order), and anything else keeps its static key. The adapter
canonicalizes the key (b = qh = kh = 0, lse_padded_rows = 1) when the
module offers dynamic_bhk. What stays static: d, dtypes, masks, the GQA
ratio (CFG.QH_PER_KH), paged pool strides, dense (non-THD) shapes. Two DSL
facts shaped this: and is staged, so const_expr(CFG.PACK_GQA and q.shape[2] != ...) must nest its constant test outside; a const_expr on a
dynamic extent or stride is an error, which is why the padded-Stats store
selects on the fake’s RANK (rank-4) and not on shape[0] > 1.
Accept means run
For a python plan, check_support() accepted ⇒ build_plans() and
execute() succeed on the buffers a caller binds as the graph declares them,
and — when the graph has a backend plan — the result matches that plan; a
python-only graph (GDN/KDA/…) has no cuDNN reference and is held to its
engine’s own numerics tests instead. Nothing an engine can decide from the
DECLARATION (operand rank, blob size, scalar shape, layout) may refuse a call
after acceptance: it is decided at check_support, and the pack hands the
engine the declaration. Buffer capacity is not decidable before execute — the
buffers arrive with the call — so a buffer smaller than its declaration is a
caller error refused at execute, not a decline. test/python/gemm/frost/test_flashinfer_shaped_gemm.py
and test/python/sdpa/frost/test_flashinfer_shaped_sdpa.py re-declare
FlashInfer’s graphs and buffers byte for byte and assert exactly this; a decline
at check_support is reported as xfail with the row’s reason, a refusal after
acceptance fails. They are the acceptance gate for turning the opt-in rows on
by default.
The compiled-plan cache
cute.compile runs the whole backend once per process for every distinct
kernel — 0.7 s for a FROST GEMM, 2.6 s for an SDPA prefill on SM100 — and the
DSL’s own file cache stores only MLIR bytecode, so a process that warms tens
of plans pays minutes at start-up. cudnn.frost.compiled_cache keeps the
exported tvm-ffi object of every kernel compiled with --enable-tvm-ffi and
reloads it in milliseconds. compile_cached(fn, *args, cache_key=, symbol=, **kwargs) is the drop-in for cute.compile at a kernel’s compile site. Two
families route through it today: the GEMM templates, with the digest of their
generated source as the key (the full GEMM suites: 5655 tests, 20:53 cold,
1:35 warm), and the SDPA forward / backward templates, with
template_key(globals(), locals(), "<function>") as the FIRST statement of
each function that compiles — the file + params digest template_loader
records as FROST_SOURCE_DIGEST, joined with the function’s name and its
arguments, which pin shapes, strides and flags the way the params pin dtypes
and masks (SDPA + GEMM suites: 5867 tests, 53:06 cold, 2:53 warm). An
argument that is not a plain value (a tensor, a device, a stream) makes the
key None and the call compiles as before. Not cached: the linear-attention
templates (their compile() takes traced tensors) and the SDPA adapters’
helper kernels in api_dsl, which compile from real tensors at call time. A
hit is the same object a miss produced, so the reloaded kernel is called
exactly as the in-process one.
The rules, borrowed from FlashInfer’s autotune cache v2 so that a stale artifact can never be reused by accident:
- Identity is the whole manifest, hashed. Frontend version AND a digest
of every
.pyin thecudnnpackage (a kernel is compiled from the shared helpers it imports as much as from its template, and the version does not move in a checkout someone is editing — an uncommitted edit lands in another directory; a wheel hashes the same every process), cutlass-dsl and tvm-ffi versions, the CUDA driver, and the device’s name, compute capability, SM count and L2 size (FROST bakes the last two into kernels) name the directory<root>/v2/<env_hash>/. Any change lands elsewhere; an unreadable field is hashed as"unknown", never skipped. - An entry is reused only under its own embedded key.
entry.jsoncarries the full key, symbol and signature and is compared on load; the object is written first and the record after it (the commit marker), both via temp file +os.replace, so a crash or a concurrent writer never yields a loadable entry without its key. - Anything doubtful is a miss: missing, malformed, mismatched, or an
object
load_modulerefuses (an arch the device cannot run). Never an error. - A hit and a miss run the same thing. A reloaded tvm-ffi function is
positional-only, so it is wrapped with the kwargs wrapper the DSL itself
uses, rebuilt from the RUNTIME spec the record carries (
arg_names, defaults, keyword-only names) — the Python signature minuscutlass.Constexprand env-stream parameters, which do not exist at run time. A wrapper built from the Python signature would shift every argument after a constexpr one (the fp8 SDPA kernels have one). The miss path exports and then reloads, so a bad artifact fails at build time, not at the next start-up. Kernels whose in-process object converts raw pointer arguments, takes a dataclass argument, or has a default JSON cannot carry are not persisted. - Location:
CUDNN_FRONTEND_COMPILED_CACHE, else$XDG_CACHE_HOME/cudnn_frontend/compiled_plans;set_cache_dir()for a caller that owns a workspace (FlashInfer);CUDNN_FRONTEND_DISABLE_COMPILED_CACHE=1turns it off;stats()reports hits / misses / bypassed / invalid / pruned per process. Bump_SCHEMAon any incompatible change. - Dead environments are retired. The manifest hashes the package’s source,
so every edited checkout and every CI commit mints an environment directory
that will never be hit again — a few hundred MB per commit on a runner with a
persistent home. A process’s first write runs
prune(): whole environment directories go (never single entries), dead schema roots first, then oldest first, until the root is underCUDNN_FRONTEND_COMPILED_CACHE_MAX_BYTES(4 GiB by default; 0 disables); the process’s own environment is never a candidate. CI should still pointCUDNN_FRONTEND_COMPILED_CACHEat a job-local directory: nothing there is ever warm across commits, and the cap is a backstop, not a policy.
Not yet routed: kernels compiled from real tensors at call time (the
linear-attention chunk_* launchers, the SDPA adapters’ _dot_fn /
_reduce_fn / _setup_fn helpers): a key for them has to spell the traced
tensors’ dtype, rank and dynamic marks; same hook once it does.
Key invariants
- uid ownership: the Python IR owns the whole uid namespace; every uid is pushed explicitly to C++ and a post-build assertion fails loudly on violation (C++ auto-assignment never runs for Python-built graphs — its enumeration order is nondeterministic for multi-output ops). A user uid landing on an auto-assigned one steals it (the holder is renumbered); user-user collisions raise.
- Pure-python or pure-C++: a graph routed to a python engine never touches C++ on the execute path; mixed construction is unsupported. (Explicitly querying the backend plan space lowers the backend entry on demand — that is the caller asking for the backend.)
- One-shot planning: a second
create_execution_plans()raises (the classic C++ graph never supported re-planning — it appends engine configs by accident). Switch plans withselect_plan(); plan differently by building a new graph. - Whole-surface freeze: after lowering/planning, every public mutation
path raises — op builders and fluent setters, direct attribute writes on
Tensor/Node/GraphContext, dict writes on node ports/params (MappingProxy), in-place dim/stride edits (sealed to tuples). Inspection stays fully readable.validate()lowers and freezes any graph the backend CAN lower (classic error timing) unless the graph’s family validated it natively (see The manifest), so the mutable-after-validate window is the ops with no backend node (GDN/KDA/…) plus natively validated graphs; a mutation in it invalidates the validation and drops the cached candidate list. Planning freezes BEFORE it analyses, so facts never describe a graph that can still change. - Output layout contract: only USER-assigned output dim/stride are pushed to the lowered graph; IR-inferred strides are provisional (row-major) and the backend keeps its classic per-op layout inference (e.g. channels-last conv). A unified layout resolver across python/cuDNN candidates belongs to the heuristics follow-up.
- Classic parity: the public
cudnn.pygraphsurface behaves as before —cudnnGraphNotSupportedErroratvalidate()for a semantically invalid graph (natively or via the backend; with a python engine candidate, a rejection the backend alone would raise — version or arch gate — surfaces atplan()instead, and only if no python engine proposes a plan), conditional outputs returnNone, torch dtypes/torch.Sizeaccepted, ragged (THD) offsets and multipliers on outputs, serialize/deserialize passthrough, plan queries delegate to the lowered graph.
Naming
cudnn.pygraph— THE public graph class (Python IR), implemented incudnn/_pygraph.py.cudnn._pybind_module.backend_graph— the internal C++ builder the IR lowers to (renamed from its pre-flip public name to avoid two things calledpygraph).
Testing the backend path
The test_native_backend_lowering.py suite builds graphs natively, lowers,
executes on GPU, and checks numerics against torch references. Dispatch-level
assertions (selected_engine is None, backend plans created, lowered graph
present) prove the execution went through the cuDNN backend plan path rather
than a python engine; kernel identity below the backend API is deliberately
not asserted (kernel names are backend-internal and version-dependent).
What each machine actually covers
A dispatch change needs three runs, because the suites SKIP rather than fail on the wrong arch — a green sweep on one box says nothing about the others:
Enumerate the device with torch.cuda.get_device_properties(i).major and pin
it with CUDA_VISIBLE_DEVICES — CUDA’s device order is not nvidia-smi’s, and
defaulting to device 0 is how an SM100 suite silently skips in full.
Follow-ups (separate MRs)
- More per-family tuning rules on top of the ranking frame, each with the
measurements behind it.
sdpa/fwd/heuristics.py::_sm120_tilesand_pack_gqa_winsare the two today; every other cell falls back to its capability row’s sole point per axis, which is the honest answer while nobody has timed it. - FALLBACK is one config per cell today — the smallest tile the row admits, the
config that asks least of the device. Picking the handful that between them
cover the plane needs measurements; the TODO is in
_mode_fallback. sdpa/fwd/placement.pyplaces the SDPA-forward family using B200 / RTX PRO 6000 measurements; an untimed row keeps the order this dispatch has always had. Rankings are tuned and evaluated offline; unit tests check the planner’s marker contract independently of workload winners. A cost model that can compare a python config against a cuDNN engine on a common currency (predicted time) would replace the table with a number.- DSL engine integration (the cuTile matmul engine lives in this track).
- Structural cleanup: lifecycle state objects, a
CudnnBackendAdapterto removeselected_engine is Nonebranching, lowering extracted to its own module, op-identity dedup (NodeType vs registry keys), longer-term a typedOpSpecas the single per-op source for builder/validation/lowering. - SDPA forward THD, padded Stats: the
-infseed of the(b, s_max, h)buffer is a separate D32 memset ahead of the kernel; fold the tail-row write into the kernel’s persistent schedule to save the launch on that path. The dense padded-Q.item()read insdpa/fwd/engines.pyruns at execute and breaks CUDA-graph capture; decide it atcheck_supportor drop it.