Python-native cudnn.pygraph and pluggable execution backends

View as Markdown

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.

cudnn.pygraph (Python IR) → create_execution_plans() → rank → ONE ranked plan list
nodes / tensors / params (freeze, analyze, (the graph's PlanConfig(engine_id,
fully introspectable query the backend) family) knobs) — python engines
AND backend entries

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

create_execution_plans([heur_mode.A, ...]) _pygraph.py
│
├─ validate() family declares a validator AND a python engine is
│ a candidate → native semantic validation (no lowering);
│ otherwise lowers + freezes any graph the backend CAN lower
├─ _finalize_backend_layout() backend layout inference lands, or records a decline
├─ _freeze() whole public surface sealed
├─ _attach_facts() family_for → resolve_analyzer → ONE parse, hung on the graph
│
└─ heuristics.rank(graph, _candidate_engines(), backend_plan_entries(), modes)
│ │ │
│ │ └─ _create_backend_plans(): one C++
│ │ create_execution_plans PER MODE, spans
│ │ recorded → every entry carries its mode
│ │ (+ the untagged delegating entry)
│ └─ manifest.engines_for(graph) — the family's offered slots, nothing else
│
├─ family_for(graph) → resolve_heuristics(family)
│ declares none → _unranked: accepting engines, then the backend
│
└─ _assemble(modes, <family>.recommend(kind, facts, offered), backend_plans)
│ e.g. sdpa/fwd/heuristics.py (propose)
├─ A → per eligible cell: a measured rule (_sm120_tiles) names the config,
│ runners-up behind it; a cell with one point per axis contributes one
│ entry. The backend's A block goes where the family's BACKEND marker
│ sits (sdpa/fwd/placement.py per shard; ours first without a marker).
├─ FALLBACK → the config expected to build, + the backend's FALLBACK block
├─ OPENSOURCE → our candidates, then the delegating entry (see below)
└─ dedup by (engine_id, knobs), first position wins
= graph.plans, position for position. build_plans() walks it.

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). matmul is 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, so execute() refuses them rather than silently running a different problem. Simple eager engines implement execute() 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_strides speak 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’s kpack) therefore applies to every slot alike — this is what stopped frost_gemm from 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_gemm reads 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_offset is an operand but hangs off the Tensor rather 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:

µs
cuLaunchKernelEx, untraced1.85
one CuTe-DSL entry~3.6
graph.execute() entry + _normalize~8
everything else is the engine’s

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 y inside a per-call function: 1.1 µs. It was 65% of what _check_plan_device cost.
  • 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.py and the per-family hook it dispatches to (see Ranking and the one plan list).
  • Every engine has a stable engine_id in one flat id space (engines/engine_ids.py: backend [0, 10_000), C++ OSS [10_000, 20_000), python [20_000, …) with a FAMILY_BLOCK-wide block per family) — reproducible pinning/autotune. Engines do not DECLARE their id: the manifest holds every slot and instantiate() 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 lets create_execution_plan(engine_id, knobs) replay an autotune result.
  • An engine declines a graph ONLY via NotImplementedError, cudnn.cudnnGraphNotSupportedError, or ImportError; anything else is an engine bug and propagates. ImportError counts because lowering imports are deferred past check_support() (see Import boundaries), so a missing optional dependency can only surface at build time — without it, a host lacking the cutedsl extra 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 against torch.matmul would put a torch reference implementation of matmul inside a dispatch test.
  • GdnCuTileEngine executes the single-node gdn and gdn_bwd ops (Gated DeltaNet linear attention) via the cuTile chunked kernels. Both ops are THD-only: token-packed [total_T, heads, dim] tensors with a required cu_seqlens (a dense batch is [0, T, 2T, ...]). gdn_bwd takes the forward inputs plus dO (and optionally d_final_state) and produces dQ/dK/dV/dG/dBeta (+ d_initial_state iff initial_state is given,
    • d_a_log iff the node carries safe_gate and an a_log input, and d_dt_bias iff it carries safe_gate and a dt_bias input). Both gate parameters are optional under safe_gate: an absent a_log is unit amplitude (exp(a_log) = 1), an absent dt_bias is zero bias, and no zero tensor is materialized for either. dG/dBeta are in raw-logit space under safe_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 raises cudnnGraphNotSupportedError at lowering. The kernels live in cudnn.linear_attention.cutile.kernels.gdn; the torch custom op cudnn.linear_attention.ops.gated_delta_net is a thin adapter that builds and executes cached gdn/gdn_bwd graphs (the SDPA op pattern), so it inherits whatever engine the planner selects. The optional use_qk_l2norm attribute asks the engine to L2-normalize the q/k rows; GdnFrostEngine (the SM100-SM103 and SM107 default, serving both gdn and gdn_bwd on 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 serves safe_gate (in-kernel raw-logit gate transform, with d_a_log/d_dt_bias for the given parameters produced by a deterministic reduction helper) and use_beta_sigmoid; the cuTile engine remains the fallback for non-128 head dims. The gate_domain attribute ("log", the default, or "linear") selects whether g is ln(alpha) or alpha itself; the FROST GDN / KDA / GDP / GDN-2 engines serve "linear" (forward and backward, dG with respect to alpha); the cuTile engines are log-only.
  • KdaFrostEngine / KdaCuTileEngine do the same for the single-node kda / kda_bwd ops (Kimi Delta Attention). KDA is GDN with a per-key-channel decay: its g is the log-space vector gate [total_T, HV, K] (GDN’s is the scalar [total_T, HV]); beta stays scalar. The FROST engine (cudnn.linear_attention.frost.kda_engine) is the forward default on SM100-SM103 and SM107; the node’s use_qk_l2norm attribute (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 serves kda_bwd on 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 is cudnn.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 gate w [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 the use_qk_l2norm attribute through to the kernel (like KdaFrostEngine), and serves gdn2_bwd the same way (checkpoint recompute when the series is absent); the op is cudnn.linear_attention.ops.gated_delta_net_v2.
  • Gated DeltaProduct (gdp / gdp_bwd) applies num_householder beta-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-token n - 1). The node carries q/g/O/dO/dQ/dG at real-token rows and k/v/beta/dK/dV/dBeta at total_T * num_householder rows; num_householder == 1 is exactly gdn. GdpFrostEngine (cudnn.linear_attention.frost.gdp_engine, SM100/SM103) is its only engine and runs the shared GDN kernels, except the d_v == 64 backward fork kernel/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 the d_v == 128 backward, 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, and cu_seqlens is scaled by n at every read site. checkpoint_every_n_tokens counts expanded sub-tokens (64 = the bwd-reusable chunk cadence; a multiple of lcm(64, n) puts every checkpoint on a real-token boundary). safe_gate, use_beta_sigmoid and allow_neg_eigval (beta as 2 * sigmoid(x)) all pass through. The op is cudnn.linear_attention.ops.gated_delta_product.
  • The FROST engines are pure pass-through: check_support requires the kernel-native dtypes (fp32/bf16/fp16 gates — io-dtype beta/w for GDN-2 — int32 or int64 cu_seqlens, fp32-or-bf16 state ports with matching initial/final dtypes, state gradients at the state’s own dtype, and io-dtype dBeta/dW for 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 from cu_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. chain cuts the batch into one wave of B * P unit-aligned pieces, P = min(num_sm // (B * HO), 16, total_chunks // (4 * B * unit)) with unit = 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 when P >= 3 for a forward plan or P >= 2 for a backward plan; every piece spans ceil(total_chunks / (B * P)) chunks whichever sequence it belongs to, so a sequence fills ceil(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 against B * P - (B - 1) slots). warmup (the decay-warmup split-K of frost/common/split_k.py) serves the band where the chain has no room, and uncut runs one item per (sequence, head). Under batch_invariant the geometry comes from the length rule alone, P = clamp(ceil(total / 8192), 1, 16) slots per sequence, of which each fills clamp(ceil(len / 8192), 1, P) on device, uncut when total <= 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-wise cu_pieces (real tokens, the slots flat in sequence order: sequence b owns slots main_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’s X being 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 of kernel/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 composes X_{j+1} = X_j M_j + H_j from initial_state (a one-piece sequence’s X is its seed, copied); the seeded main kernel runs over the piece work items as independent sequences and the last filled piece of each sequence writes final_state in 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 by d_final_state, the bprop over the pieces seeded by row 0 of each piece’s series, and piece 0 writes d_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 yields H = 0, M = I) and one emitting state chain composes them, its tail being final_state (in reverse, d_initial_state) and its running product transition.
  • In-place state update. The main ops’ overwrite_initial_state attribute (fwd and bwd nodes of GDN, KDA, GDN-2 and GDP) lets one buffer serve as initial_state and final_state (in the backward, as d_final_state and d_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 the initial_state input and the final_state output (bwd: d_final_state and d_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_indices input of [N] row ids: initial_state is then a pool [N_pool, HO, V, K] whose row state_indices[i] seeds sequence i, and final_state (the pool itself under overwrite_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 with checkpoint_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 with state_indices. The torch ops route state_indices through cudnn::<op>_fwd_overwrite_state, so the caller’s pool is updated in place and comes back as final_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 *_summary returns H (the span’s final state from a zero seed) and M_buf (the transition, holding M^T), composed forward from the incoming state as X_{j+1} = X_j @ M_buf_j + H_j; the span forward then runs with initial_state = X_j. The backward summary *_summary_bwd returns G (the incoming state gradient from a zero outgoing gradient) and, with output_transition=True, the transition in the backward’s orientation (M_buf^T), composed in reverse from the outgoing state gradient as dX_j = dX_{j+1} @ transition_j + G_j; the span backward then runs with initial_state = X_j and dX_{j+1} as the gradient on its final_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_FAMILY maps 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, so matmul + pointwise is a gemm graph. Whether a family can serve the WHOLE graph is its analyzer’s judgment. A coarser copy of that judgment here is what closed_under was, and it promised RESHAPE support nothing implemented.
  • family_for(graph) is a pure property of the graph — no sm, 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 Capabilities declares 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 by CUDNN_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 heuristics hook — like analyzer, 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; a BACKEND marker 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 validator hook — 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_layout records the decline, plan() raises it if no python engine proposes a plan either). A family without one validates classically. cudnn/_sdpa_validate.py and cudnn/_gemm_validate.py are 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 no Capabilities, 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 two validate() 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. SdpaGraphFacts records has_bias=True as a fact, never an error; each engine’s Capabilities row does the rejecting in mismatch(). A shared parser that starts rejecting becomes an if-ladder that must know every kernel.
  • Framework-neutral vocabulary: cudnn.data_type, not torch.dtype; device from cudnn.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) is cudnn.get_stream(handle). With no handle it is None, 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 from torch.cuda.current_device(). Engines re-check the arch in check_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 under cudnn/sdpa are lazy (PEP 562) and EngineSpec.lower resolves its DSL adapter at build time, so import cudnn.sdpa.graph_analyzer costs 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_VERSION in frost/buffers.py is 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.py holds 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::rank calls the family’s hook with (kind, facts, offered). The hook returns its own proposals plus an optional internal cudnn.engines.heuristics.BACKEND marker. _assemble() expands that marker into the mode’s backend block, then strips mode annotations and deduplicates to form the final graph.plans list.
  • The backend’s entries arrive tagged with the mode that produced them. _create_backend_plans() asks C++ one heuristic mode at a time and records get_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.OPENSOURCE is 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_index it holds in the lowered graph’s own plan list, so building it is one build_plan_at_index. Backend engine sets are still never statically enumerated: they are discovered per graph at plan time, and backend_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 its cpp_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_ID names one thing only: the delegating entry backend_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, because Graph::build_plans tries 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 heuristics hook, 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 None only where the capability row declares no domain for that axis. None never 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_tiles is 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 as cudnn.knob_type) is the ONE vocabulary for backend and python plans. Backend values 0..32 are frozen; python-only axes live in the band from FRONTEND_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_type refuses 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 through BaseEngine.knobs_to_public / knobs_from_public. get_engine_and_knobs_at_index always returns a dict ({} when the plan has no axes, never None); create_execution_plan( engine_id, {cudnn.knob_type: int}) replays it; get_plan_name_at_index prints engine[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 complete PlanConfig(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 inside build_plan is 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’s python_only_attrs: never forwarded to C++, a SET value makes the node backend-unlowerable (serialize() and key() 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 .py in the cudnn package (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.json carries 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_module refuses (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 minus cutlass.Constexpr and 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=1 turns it off; stats() reports hits / misses / bypassed / invalid / pruned per process. Bump _SCHEMA on 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 under CUDNN_FRONTEND_COMPILED_CACHE_MAX_BYTES (4 GiB by default; 0 disables); the process’s own environment is never a candidate. CI should still point CUDNN_FRONTEND_COMPILED_CACHE at 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 with select_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.pygraph surface behaves as before — cudnnGraphNotSupportedError at validate() 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 at plan() instead, and only if no python engine proposes a plan), conditional outputs return None, torch dtypes/torch.Size accepted, 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 in cudnn/_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 called pygraph).

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:

targetcovers
any GPU (CPU-only logic)test_dispatch.py, test_graph_native.py, test_import_boundaries.py — the ranking contract, one-shot planning, the at-index APIs, manifest classification
SM100every FROST SDPA-forward cell except sm120; the FROST GEMM family; linear_attention (GDN / KDA / GDN-2), which is where the cuTile-vs-FROST pin lives
SM120sdpa_fwd_prefill_sm120 and its tile rule; test_mhas_v2 routing tallies

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_tiles and _pack_gqa_wins are 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.py places 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 CudnnBackendAdapter to remove selected_engine is None branching, lowering extracted to its own module, op-identity dedup (NodeType vs registry keys), longer-term a typed OpSpec as the single per-op source for builder/validation/lowering.
  • SDPA forward THD, padded Stats: the -inf seed 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 in sdpa/fwd/engines.py runs at execute and breaks CUDA-graph capture; decide it at check_support or drop it.