Make NNlib's batched wrappers first-class operands - #12
Merged
Merged
Conversation
Codecov Report✅ All modified and coverable lines are covered by tests. 📢 Thoughts on this report? Let us know! |
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>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
NNlib.batched_transpose/NNlib.batched_adjointoperands previously fell throughunwrap_op's'N'fallback and failed the unit-stride check with a misleading "needs unit stride" error (NNlib definesstrideson 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 intotransA/transBfor the descriptor.NNlibExtaddsBatchedTranspose → 'T'andBatchedAdjoint → 'C', mirroring BaseTranspose/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 sharedchecked_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'sstridesdefinition on the wrapper never enters.New testset mirrors the PermutedDimsArray one: planless and planned round-trips,
batched_adjointagreeing numerically withbatched_transposeforFloat32, 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