Skip to content

Make NNlib's batched wrappers first-class operands - #12

Merged
AntonOresten merged 1 commit into
mainfrom
nnlib-batched-traits
Aug 7, 2026
Merged

AntonOresten merged 1 commit into
mainfrom
nnlib-batched-traits

Conversation

@AntonOresten

Copy link
Copy Markdown
Member

NNlib.batched_transpose / NNlib.batched_adjoint operands previously fell through unwrap_op's 'N' fallback and failed the unit-stride check with a misleading "needs unit stride" error (NNlib defines strides on the wrapper as the permuted parent strides).

Orientation wrappers touch the package at two seams, and both are now cleanly extensible:

  • unwrap_op (plan time) — reads the wrapper type into transA/transB for the descriptor. NNlibExt adds BatchedTranspose → 'T' and BatchedAdjoint → 'C', mirroring Base Transpose/Adjoint ('C' ≡ 'T' for real element types).
  • apply_operand (apply time) — strips the wrapper before taking the storage pointer (the transpose is the descriptor's job) and checks it against the plan's baked orientation, so a 'C' wrapper handed to a 'T' plan throws instead of silently reading sideways. The check body is hoisted into a shared checked_unwrap, making extension methods one-liners.

Everything downstream — layouts, leading dimensions, batch strides, the plan cache — derives from the unwrapped parent, exactly as for the existing PermutedDimsArray{<:Any,3,(2,1,3)} spelling; NNlib's strides definition on the wrapper never enters.

New testset mirrors the PermutedDimsArray one: planless and planned round-trips, batched_adjoint agreeing numerically with batched_transpose for Float32, and rejection of a wrapper that disagrees with the plan's orientation. Verified locally on RTX 6000 Ada: matmul + epilogue suites clean, new testset 8/8.

Deliberately not included: hooking NNlib.batched_mul! — routing decisions belong to the layer that owns the array types, not here.

🤖 Generated with Claude Code

@codecov

codecov Bot commented Aug 7, 2026

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.

📢 Thoughts on this report? Let us know!

@AntonOresten
AntonOresten merged commit 108c3ae into main Aug 7, 2026
4 checks passed
@AntonOresten
AntonOresten deleted the nnlib-batched-traits branch August 7, 2026 23:00
batched_transpose/batched_adjoint previously hit the unwrap_op fallback
and died on the unit-stride check with a misleading error (NNlib defines
strides on the wrapper as the permuted parent strides).

Orientation wrappers have two seams, and both are extensible: unwrap_op
reads the wrapper type at plan time (transA/transB into the descriptor),
and apply_operand strips the wrapper at apply time — cuBLASLt only ever
wants the storage pointer — while checking it against the plan's baked
orientation. The check body is hoisted into checked_unwrap so extension
methods are one-liners.

NNlibExt now defines both: BatchedTranspose → 'T', BatchedAdjoint → 'C'
(mirroring Base Adjoint; identical for real element types). Everything
downstream — layouts, strides, batching, the plan cache — derives from
the unwrapped parent, exactly as for PermutedDimsArray{<:Any,3,(2,1,3)}.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant