From 5f771147a6150077e1ba6350f28c8b1031e5bf5d Mon Sep 17 00:00:00 2001 From: lkdvos Date: Thu, 24 Sep 2026 08:10:48 -0400 Subject: [PATCH] Add batched out-of-place tensor permutation (CPU + GPU) Adds Strided.batched_permutedims!/plan_batched_permutedims: permute a batch of differently-shaped tensors with one shared perm in a single CPU pass or a single GPU kernel launch, with selectable GPU execution strategies (elementwise, thread-tile, cooperative shared-memory tile) and optional per-tensor alpha/beta scaled accumulation. Co-Authored-By: Claude Sonnet 5 --- ext/StridedGPUArraysExt.jl | 4 + ext/StridedGPUArraysExt_batched.jl | 431 ++++++++++++ src/Strided.jl | 2 + src/batched_permutedims.jl | 466 +++++++++++++ src/batched_permutedims_cpu.jl | 144 ++++ test/batched_permutedims.jl | 669 ++++++++++++++++++ test/gpu.jl | 1018 ++++++++++++++++++++++++++++ test/runtests.jl | 2 + 8 files changed, 2736 insertions(+) create mode 100644 ext/StridedGPUArraysExt_batched.jl create mode 100644 src/batched_permutedims.jl create mode 100644 src/batched_permutedims_cpu.jl create mode 100644 test/batched_permutedims.jl diff --git a/ext/StridedGPUArraysExt.jl b/ext/StridedGPUArraysExt.jl index 7525b44..9d29774 100644 --- a/ext/StridedGPUArraysExt.jl +++ b/ext/StridedGPUArraysExt.jl @@ -163,4 +163,8 @@ function Strided.isblasmatrix(A::GPUStridedView{T, 2}) where {T <: LinearAlgebra end end +# ---------- GPU batched out-of-place permutation support ---------- + +include("StridedGPUArraysExt_batched.jl") + end diff --git a/ext/StridedGPUArraysExt_batched.jl b/ext/StridedGPUArraysExt_batched.jl new file mode 100644 index 0000000..d163568 --- /dev/null +++ b/ext/StridedGPUArraysExt_batched.jl @@ -0,0 +1,431 @@ +# GPU executor: `_batched_permute!` for `GPUStridedView`s, plus the planner hooks. ONE kernel +# launch per batch: descriptors, tile-id prefix sums and base addresses are uploaded as device +# arrays (cached per plan), and every work item locates its own tensor and tile on the device. + +import Strided: BatchedPermuteDesc, BatchedPermuteFamily, BatchedPermutePlan, BatchedPermuteStrategy +import Strided: BP_COPY, BP_PAYLOAD, BP_TRANSPOSE, BP_AUTO, BP_ELEMENTWISE, BP_THREADTILE, BP_GROUPTILE +import Strided: _batched_deviceid, _batched_tileshape, _batched_resolve_strategy, _tile_offsets, _batched_axpby +using GPUArrays.KernelAbstractions: @localmem, @synchronize + +# `BP_THREADTILE`: elements per thread and along which axis (`:dstfast`/`:srcfast`); `Ref`s so a +# benchmark can flip them at runtime without recompiling anything. +const _BP_THREADTILE_K = Ref(8) +const _BP_THREADTILE_ALONG = Ref(:dstfast) + +# `BP_GROUPTILE`: an EDGE x EDGE tile transposed through local memory by EDGE*ROWS threads +# (EDGE÷ROWS rows each). Each entry is compiled separately via `Val(edge)`/`Val(rows)` -- the ONE +# sanctioned exception to "no shape-like quantity enters a type domain", bounded by the menu size +# (<= 3, asserted). Nothing derived from a tensor's shape, stride, offset, batch length or tile +# count may join it or become a `Val` elsewhere. Ordered by decreasing edge; entry 1 is the +# default and wins ties. ROWS must divide EDGE, or the load/store loops silently skip tile rows. +const _BP_GROUPTILE_GEOMETRIES = ((edge = 32, rows = 4), (edge = 16, rows = 4)) +const _BP_GROUPTILE_PAD = 1 # +1 column against shared-memory bank conflicts +const _BP_GROUPTILE_META_LEN = 6 +@assert _BP_GROUPTILE_META_LEN == 6 # kernel uses literal meta[1..6]; lowering this silently +# leaves those indices out of bounds (a wrong result, not an error) +const _BP_GROUPTILE_SHMEM_BUDGET = 32768 # tile bytes only; well under CUDA's 48 KiB for Metal/ROCm +@assert _BP_GROUPTILE_SHMEM_BUDGET + _BP_GROUPTILE_META_LEN * sizeof(Int) <= 49152 +@assert 1 <= length(_BP_GROUPTILE_GEOMETRIES) <= 3 +@assert all(g -> g.edge >= 1 && g.rows >= 1 && g.edge % g.rows == 0, _BP_GROUPTILE_GEOMETRIES) +@assert all( + i -> _BP_GROUPTILE_GEOMETRIES[i].edge > _BP_GROUPTILE_GEOMETRIES[i + 1].edge, + 1:(length(_BP_GROUPTILE_GEOMETRIES) - 1) +) + +_batched_grouptile_tilebytes(::Type{T}, edge::Int) where {T} = + sizeof(T) * (edge + _BP_GROUPTILE_PAD) * edge +_batched_grouptile_fits(::Type{T}, edge::Int) where {T} = + _batched_grouptile_tilebytes(T, edge) <= _BP_GROUPTILE_SHMEM_BUDGET + +# Tile-slot utilization of `edge`: real elements / scheduled edge^2 thread slots over the batch. +# Active axes are 1 and `c`; the other extents multiply the tile count. +function _batched_grouptile_utilization(dims, c::Int, edge::Int) + elements = 0 + slots = 0 + for d in dims + ntiles = 1 + for g in eachindex(d) + ntiles *= (g == 1 || g == c) ? cld(d[g], edge) : d[g] + end + slots += ntiles * edge * edge + elements += prod(d) + end + return slots == 0 ? 1.0 : elements / slots +end + +# Menu index with the highest utilization among entries whose tile fits `T`; a tie keeps the +# earlier (larger) entry. Some entry always fits when this runs (the resolver admits +# `BP_GROUPTILE` only then), so the throw marks an invariant rather than a caught path. +function _batched_grouptile_geometry(dims, c::Int, ::Type{T}) where {T} + best = 0 + bestu = -1.0 + for (i, geom) in enumerate(_BP_GROUPTILE_GEOMETRIES) + _batched_grouptile_fits(T, geom.edge) || continue + u = _batched_grouptile_utilization(dims, c, geom.edge) + if u > bestu + best, bestu = i, u + end + end + best == 0 && throw(ArgumentError( + "batched_permutedims!: internal invariant violated: no BP_GROUPTILE geometry fits element type $T" + )) + return best +end + +# `BP_AUTO` -> the elementwise baseline (unchanged default). `BP_GROUPTILE` needs the transpose +# family (N >= 2, `qt[1] != 1`) and an element type some menu tile fits; otherwise it falls back +# to `BP_ELEMENTWISE` deliberately -- not an error, inspectable via `plan.strategy`. `any` fits +# <=> the smallest edge fits, and the tile-shape hook selects only among fitting entries. +function _batched_resolve_strategy( + ::KernelAbstractions.Backend, requested::BatchedPermuteStrategy, + family::BatchedPermuteFamily, ::Type{T} + ) where {T} + requested === BP_AUTO && return BP_ELEMENTWISE + requested === BP_GROUPTILE || return requested + fits = any(geom -> _batched_grouptile_fits(T, geom.edge), _BP_GROUPTILE_GEOMETRIES) + return (family === BP_TRANSPOSE && fits) ? BP_GROUPTILE : BP_ELEMENTWISE +end + +# Distinguishes backends only, not devices within one backend (KA backends carry no device index). +_batched_deviceid(a::GPUStridedView) = KernelAbstractions.get_backend(parent(a)) + +# ELEMENTWISE: tile 1 everywhere (one element per thread). THREADTILE: `k` elements along one +# axis -- `:srcfast` or BP_COPY: axis 1; `:dstfast`: `qt[1]` for transposes, axis 2 for payload +# (axis 1 is already the shared contiguous run) -- clamped to that axis's largest extent. +# GROUPTILE: a menu edge on axes 1 and `c = qt[1]` (distinct, since only the transpose family +# gets here), 1 elsewhere; the tile's edge is how the chosen geometry reaches the executor. +function _batched_tileshape( + ::KernelAbstractions.Backend, strategy::BatchedPermuteStrategy, family::BatchedPermuteFamily, + qt::NTuple{N, Int}, maxdims::NTuple{N, Int}, dims::Vector{NTuple{N, Int}}, ::Type{T}, + ::Int, ::Int + ) where {N, T} + if strategy === BP_THREADTILE + k = _BP_THREADTILE_K[] + along = _BP_THREADTILE_ALONG[] + along in (:dstfast, :srcfast) || + throw(ArgumentError("_BP_THREADTILE_ALONG[] must be :dstfast or :srcfast, got $along")) + axis = (along === :srcfast || family === BP_COPY) ? 1 : (family === BP_TRANSPOSE ? qt[1] : 2) + kk = clamp(k, 1, max(maxdims[axis], 1)) + return ntuple(g -> g == axis ? kk : 1, N) + elseif strategy === BP_GROUPTILE + c = qt[1] + edge = _BP_GROUPTILE_GEOMETRIES[_batched_grouptile_geometry(dims, c, T)].edge + return ntuple(g -> (g == 1 || g == c) ? edge : 1, N) + else # BP_ELEMENTWISE + return ntuple(Returns(1), N) + end +end + +# --- device upload --- + +# Descriptors hold plain integers (no pointers), so they upload like any array of numbers. +function _batched_upload(proto, x::Vector) + d = similar(proto, eltype(x), (length(x),)) + copyto!(d, x) + return d +end + +# Per-tensor base addresses: `B` is a runtime value, so the parents go into one uploaded array. +# Preferred element: `KernelAbstractions.argconvert(kernel!, parent)`, the backend's own +# indexable device-array struct. Where that is not `isbits` (the JLArrays backend) fall back to +# raw `UInt64` addresses, reinterpreted as `Ptr{T}` in the kernel. Detected by trying, not by name. +@inline function _batched_try_argconvert(kernel!, p) + local adapted + try + adapted = KernelAbstractions.argconvert(kernel!, p) + catch + return nothing + end + return isbitstype(typeof(adapted)) ? adapted : nothing +end + +function _batched_convert_bases(kernel!, parents::Vector) + first_base = _batched_try_argconvert(kernel!, parents[1]) + if first_base !== nothing + DT = typeof(first_base) + return DT[something(_batched_try_argconvert(kernel!, p))::DT for p in parents] + else + return UInt64[UInt64(UInt(pointer(p))) for p in parents] + end +end + +@inline _gpu_getbase(base::UInt64, ::Val{T}, i::Int) where {T} = + unsafe_load(reinterpret(Ptr{T}, base), i) +@inline _gpu_setbase!(base::UInt64, ::Val{T}, i::Int, v) where {T} = + unsafe_store!(reinterpret(Ptr{T}, base), v, i) +@inline _gpu_getbase(base, ::Val, i::Int) = (@inbounds base[i]) +@inline _gpu_setbase!(base, ::Val, i::Int, v) = (@inbounds base[i] = v; nothing) + +# Per-tensor coefficients inside a kernel: `nothing` when the call has none (the store below is +# then exactly the plain store), else `(alpha[b], beta[b])`. `beta == 0` must not read `dst`. +@inline _gpu_coef(::Nothing, ::Nothing, ::Int) = nothing +@inline _gpu_coef(alpha, beta, b::Int) = (@inbounds (alpha[b], beta[b])) +@inline _gpu_put!(dbase, valT::Val, i::Int, v, ::Nothing) = _gpu_setbase!(dbase, valT, i, v) +@inline function _gpu_put!(dbase, valT::Val, i::Int, v, (a, b)::Tuple) + w = iszero(b) ? a * v : _batched_axpby(a, v, b, _gpu_getbase(dbase, valT, i)) + return _gpu_setbase!(dbase, valT, i, w) +end + +# --- device binding cache (`plan.devcache`) --- + +# The upload is reused only if the device identity, EVERY parent's current `pointer(...)`, and +# the full host descriptors (dims, strides, offsets, tile counts -- plan reuse only requires +# matching shapes, so strides can differ at the same address) all still match. The per-call +# address read is mandatory, with no identity or length shortcut: `resize!` can move a buffer +# while keeping the same array object and length, and a stale binding would then read/write +# freed memory silently. Strong refs to the parents keep them alive while the binding is cached. +mutable struct _BatchedGPUBinding + deviceid::Any + srcaddrs::Vector{UInt} + dstaddrs::Vector{UInt} + hostdescs::Any # `plan.descs` at upload time. On unchanged reuse this IS the live + # vector, a valid snapshot only because plan descriptors are never mutated in place. + descs::Any + prefix::Any + srcbases::Any + dstbases::Any + srcparents::Vector{Any} + dstparents::Vector{Any} +end + +function _batched_addrs_match(addrs::Vector{UInt}, parents::Vector) + length(addrs) == length(parents) || return false + for i in eachindex(parents) + addrs[i] == UInt(pointer(parents[i])) || return false + end + return true +end + +# Stops at the first mismatch but never skips a check class; allocates nothing on a match. +# `===` only shortcuts the `==` that would follow (reflexive for a vector of `isbits` descriptors). +function _batched_binding_matches( + b::_BatchedGPUBinding, deviceid, srcparents::Vector, dstparents::Vector, hostdescs + ) + isequal(b.deviceid, deviceid) || return false + _batched_addrs_match(b.srcaddrs, srcparents) || return false + _batched_addrs_match(b.dstaddrs, dstparents) || return false + return b.hostdescs === hostdescs || b.hostdescs == hostdescs +end + +_batched_addresses(parents::Vector) = UInt[UInt(pointer(p)) for p in parents] + +# The binding is built for the kernel actually launched (`argconvert` takes the kernel); a plan's +# strategy is fixed, so a cached binding is only ever reused with the same kernel. +function _batched_get_binding!( + plan::BatchedPermutePlan, kernel!, srcparents::Vector, dstparents::Vector + ) + hostdescs = plan.descs + cached = plan.devcache[] + cached isa _BatchedGPUBinding && + _batched_binding_matches(cached, plan.deviceid, srcparents, dstparents, hostdescs) && + return cached + proto = srcparents[1] + binding = _BatchedGPUBinding( + plan.deviceid, _batched_addresses(srcparents), _batched_addresses(dstparents), hostdescs, + _batched_upload(proto, plan.descs), _batched_upload(proto, plan.prefix), + _batched_upload(proto, _batched_convert_bases(kernel!, srcparents)), + _batched_upload(proto, _batched_convert_bases(kernel!, dstparents)), + Any[srcparents...], Any[dstparents...] + ) + plan.devcache[] = binding + return binding +end + +# --- kernels --- + +# `searchsortedlast` by hand: device code cannot use Base's generic implementation. +@inline function _batched_dev_searchsortedlast(prefix, u::Int) + lo = 1 + hi = length(prefix) + while lo < hi + mid = (lo + hi + 1) >> 1 + @inbounds if prefix[mid] <= u + lo = mid + else + hi = mid - 1 + end + end + return lo +end + +# Tensor index, descriptor and local tile id of global tile `u`. +@inline function _batched_dev_locate(descs, prefix, u::Int) + b = _batched_dev_searchsortedlast(prefix, u) + @inbounds return b, descs[b], u - prefix[b] +end + +# One work item per global tile; shared by ELEMENTWISE and THREADTILE, which differ only in +# `tile`. `@index(Global, Cartesian)` with a 1-tuple ndrange: `Linear` fails on the JLArrays +# backend in this KA version (Strided's own GPU mapreduce kernel does the same). Work items past +# `ndrange` from workgroup padding are masked by KA's own `__validindex` guard. +@kernel function _batched_permute_gpu_kernel!( + descs, prefix, dstbases, srcbases, tile::NTuple{N, Int}, valT::Val{T}, alpha, beta + ) where {N, T} + Idx = @index(Global, Cartesian) + b, desc, l = _batched_dev_locate(descs, prefix, Idx[1] - 1) + soff0, doff0, d = _tile_offsets(desc, tile, l) + @inbounds sbase = srcbases[b] + @inbounds dbase = dstbases[b] + ab = _gpu_coef(alpha, beta, b) + for I in CartesianIndices(map(Base.OneTo, d)) + so = soff0 + do_ = doff0 + for g in 1:N + @inbounds so += (I[g] - 1) * desc.srcstrides[g] + @inbounds do_ += (I[g] - 1) * desc.dststrides[g] + end + v = _gpu_getbase(sbase, valT, so + 1) + _gpu_put!(dbase, valT, do_ + 1, v, ab) + end +end + +# `t[c]` for a runtime axis `c` WITHOUT indexing the tuple: a dynamic `getindex` in device code +# spills the whole tuple (and what it came from) to per-thread local memory -- measured as the +# single largest cost in the cooperative kernel. This unrolls to `N-1` branch-free selects. +# Defined for `1 <= c <= N`; out of range yields 0 (a wrong address, not an error), which is why +# the launch site checks the tile/`c` invariant before every launch. +@inline _batched_selaxis(t::Tuple, c::Int) = _batched_selaxis(t, c, 1) +@inline _batched_selaxis(t::Tuple, c::Int, g::Int) = + ifelse(g == c, t[1], _batched_selaxis(Base.tail(t), c, g + 1)) +@inline _batched_selaxis(::Tuple{}, ::Int, ::Int) = 0 + +# One workgroup of EDGE*ROWS threads per global tile, EDGE x EDGE on axes 1 (source-fastest) +# and `c` (destination-fastest). Reads coalesced along axis 1 into `lmem`, barrier, writes the +# tile transposed, coalesced along `c`. Partial tiles are guarded per load/store; the barrier +# never is. `EDGE`/`ROWS` come only from the menu entry matching `plan.tile` (launch site). +# +# Performance rules (both measured, not stylistic): +# * never index a tuple with a runtime index in this body -- use `_batched_selaxis`; +# * `g` is workgroup-uniform, so thread 1 publishes the six store-phase scalars into `meta` +# before the barrier: one lookup + `_tile_offsets` per workgroup instead of two per thread. +# KernelAbstractions CPU-backend (JLArrays) rules, both load-bearing: +# 1. every `@localmem`/`@index` is a bare top-level `lhs = ...` statement, `@localmem` first; +# nested in `if`/`let`/loops or larger expressions the CPU transform cannot re-splice it; +# 2. `@synchronize()` splits the body into separate work-item loops on the CPU, so no ordinary +# local survives it -- only `@localmem` arrays and re-spliced `@index` values do (hence +# `meta`, and the restated `l = @index(Local, Linear)` after the barrier). +@kernel function _batched_permute_gpu_grouptile_kernel!( + descs, prefix, dstbases, srcbases, tile::NTuple{N, Int}, c::Int, valT::Val{T}, + ::Val{EDGE}, ::Val{ROWS}, alpha, beta + ) where {N, T, EDGE, ROWS} + lmem = @localmem T (EDGE + _BP_GROUPTILE_PAD, EDGE) + meta = @localmem Int (_BP_GROUPTILE_META_LEN,) + g = @index(Group, Linear) + l = @index(Local, Linear) + b, desc, lt = _batched_dev_locate(descs, prefix, g - 1) + soff0, doff0, d = _tile_offsets(desc, tile, lt) + d1 = d[1] + dc = _batched_selaxis(d, c) + s1 = desc.srcstrides[1] + sc = _batched_selaxis(desc.srcstrides, c) + @inbounds sbase = srcbases[b] + if l == 1 + @inbounds meta[1] = b + @inbounds meta[2] = d1 + @inbounds meta[3] = dc + @inbounds meta[4] = desc.dststrides[1] + @inbounds meta[5] = _batched_selaxis(desc.dststrides, c) + @inbounds meta[6] = doff0 + end + li = (l - 1) % EDGE + 1 + lj0 = (l - 1) ÷ EDGE + 1 + for k in 0:(EDGE ÷ ROWS - 1) + lj = lj0 + k * ROWS + if li <= d1 && lj <= dc + @inbounds lmem[li, lj] = _gpu_getbase(sbase, valT, soff0 + (li - 1) * s1 + (lj - 1) * sc + 1) + end + end + @synchronize() + l = @index(Local, Linear) + @inbounds b2 = meta[1] + @inbounds d1b = meta[2] + @inbounds dcb = meta[3] + @inbounds t1 = meta[4] + @inbounds tc = meta[5] + @inbounds doff0 = meta[6] + @inbounds dbase = dstbases[b2] + ab = _gpu_coef(alpha, beta, b2) + li = (l - 1) % EDGE + 1 + lj0 = (l - 1) ÷ EDGE + 1 + for k in 0:(EDGE ÷ ROWS - 1) + lj = lj0 + k * ROWS + if lj <= d1b && li <= dcb + @inbounds v = lmem[lj, li] + _gpu_put!(dbase, valT, doff0 + (lj - 1) * t1 + (li - 1) * tc + 1, v, ab) + end + end +end + +# --- launch --- + +# Every call must have finished moving data before returning (it owns the lifetime of the +# device metadata). The JLArrays backend has no `synchronize` method in this KA version but runs +# launches synchronously, so exactly that `MethodError` means "nothing to wait for"; anything +# else propagates. +function _batched_gpu_synchronize(backend) + try + KernelAbstractions.synchronize(backend) + catch e + e isa MethodError || rethrow() + end + return nothing +end + +# `GC.@preserve` keeps the parents (and the per-call coefficient uploads in `args`) alive across +# launch + synchronize, on top of the binding's own strong references. +function _batched_launch!(plan::BatchedPermutePlan, kernel!, backend, srcparents, dstparents, args...; ndrange) + binding = _batched_get_binding!(plan, kernel!, srcparents, dstparents) + GC.@preserve srcparents dstparents args begin + kernel!(binding.descs, binding.prefix, binding.dstbases, binding.srcbases, plan.tile, args...; ndrange) + _batched_gpu_synchronize(backend) + end + return nothing +end + +# Coefficients are per call, so they are uploaded per call and never enter the binding cache. +_batched_upload_coeffs(::Any, ::Nothing) = (nothing, nothing) +_batched_upload_coeffs(proto, (a, b)::Tuple) = (_batched_upload(proto, a), _batched_upload(proto, b)) + +function Strided._batched_permute!( + plan::BatchedPermutePlan{N0, N, T}, dst::Vector{<:GPUStridedView}, src::Vector{<:GPUStridedView}, + coeffs + ) where {N0, N, T} + srcparents = parent.(src) + dstparents = parent.(dst) + backend = KernelAbstractions.get_backend(srcparents[1]) + alpha, beta = _batched_upload_coeffs(srcparents[1], coeffs) + strategy = plan.strategy + if strategy === BP_GROUPTILE + # The tile must be exactly a menu edge on axes 1 and `c` (with `c != 1`), 1 elsewhere, + # and that edge must fit `T`: `_batched_selaxis` would yield a wrong address, not an + # error, and an unfitting edge would request local memory over budget. Checked, not assumed. + c = plan.dstseq[1] + gi = findfirst(geom -> geom.edge == plan.tile[1], _BP_GROUPTILE_GEOMETRIES) + tileok = gi !== nothing && c != 1 && plan.tile[c] == plan.tile[1] && + all(g == 1 || g == c || plan.tile[g] == 1 for g in eachindex(plan.tile)) + tileok || throw(ArgumentError( + "batched_permutedims!: internal invariant violated: BP_GROUPTILE plan has tile=$(plan.tile), dstseq=$(plan.dstseq)" + )) + geom = _BP_GROUPTILE_GEOMETRIES[gi] + _batched_grouptile_fits(T, geom.edge) || throw(ArgumentError( + "batched_permutedims!: internal invariant violated: BP_GROUPTILE tile edge $(geom.edge) does not fit element type $T" + )) + wgsize = geom.edge * geom.rows + kernel! = _batched_permute_gpu_grouptile_kernel!(backend, (wgsize,)) + # `Val(geom.edge)`/`Val(geom.rows)`: `geom` is a menu entry, so at most one variant per entry. + _batched_launch!( + plan, kernel!, backend, srcparents, dstparents, c, Val(T), Val(geom.edge), Val(geom.rows), + alpha, beta; ndrange = (plan.totaltiles * wgsize,) + ) + elseif strategy === BP_ELEMENTWISE || strategy === BP_THREADTILE + kernel! = _batched_permute_gpu_kernel!(backend) + _batched_launch!( + plan, kernel!, backend, srcparents, dstparents, Val(T), alpha, beta; ndrange = (plan.totaltiles,) + ) + else + throw(ArgumentError( + "batched_permutedims!: internal invariant violated: unresolved GPU strategy $strategy at execution time" + )) + end + return dst +end diff --git a/src/Strided.jl b/src/Strided.jl index 5b5e319..d7bbc70 100644 --- a/src/Strided.jl +++ b/src/Strided.jl @@ -56,6 +56,8 @@ include("mapreduce.jl") include("broadcast.jl") include("macros.jl") include("convert.jl") +include("batched_permutedims.jl") +include("batched_permutedims_cpu.jl") include("precompile.jl") diff --git a/src/batched_permutedims.jl b/src/batched_permutedims.jl new file mode 100644 index 0000000..ddda753 --- /dev/null +++ b/src/batched_permutedims.jl @@ -0,0 +1,466 @@ +# Batched out-of-place permutation: `dsts[b] = permutedims(srcs[b], perm)` for every `b`, +# with one shared `perm`, one common rank and element type, and free per-tensor sizes. This +# file only plans; the executors are `batched_permutedims_cpu.jl` (CPU) and +# `ext/StridedGPUArraysExt_batched.jl` (GPU), both reached via `_batched_permute!(plan, dst, src)`. +# +# Addressing uses each view's own `strides(...)` as-is: no density check and NO ALIASING +# CHECK. Overlapping destinations, or a destination overlapping a source, is undefined +# behavior (a silently wrong result, not an error). Only source/source aliasing is safe. + +@enum BatchedPermuteFamily BP_COPY BP_PAYLOAD BP_TRANSPOSE + +# `BP_AUTO` lets the backend decide. The other values are GPU-only: on a CPU batch they are an +# error, never a silent fallback (a silent no-op would make strategy comparisons misleading). +@enum BatchedPermuteStrategy BP_AUTO BP_ELEMENTWISE BP_THREADTILE BP_GROUPTILE + +# Per-tensor addressing data in the batch's shared reduced-axis order: coordinate `c` lives at +# `srcoffset + sum(c .* srcstrides)` / `dstoffset + sum(c .* dststrides)`, `0 <= c[g] < dims[g]`. +# Strides are the tensor's own real strides (never derived from extents); `tilecounts[g] = +# cld(dims[g], tile[g])`. Plain integers only, so this is `isbits` and uploads to a GPU as-is. +struct BatchedPermuteDesc{N} + dims::NTuple{N, Int} + dststrides::NTuple{N, Int} + srcstrides::NTuple{N, Int} + tilecounts::NTuple{N, Int} + dstoffset::Int + srcoffset::Int +end + +# N0 = batch rank, N = reduced rank (>= 1); no tensor shape or batch length enters the type. +# `descs`/`cpublocks` are never mutated in place once stored (they are only filled while local +# to the function building them). `_rebuild_descriptors` and the GPU binding cache both rely +# on this to hand back the same vectors, or the same plan object, when nothing changed. +struct BatchedPermutePlan{N0, N, T, I} + perm::NTuple{N0, Int} + srcorder::NTuple{N0, Int} # heuristic source-axis order (`_infer_order`) + groups::NTuple{N0, Int} # reduced axis (1:N) of each srcorder position; 0 = dropped + dstseq::NTuple{N, Int} # heuristic destination order, in reduced labels + family::BatchedPermuteFamily + strategy::BatchedPermuteStrategy # resolved once, at plan time + tile::NTuple{N, Int} # batch-wide tile shape H + srclabels::NTuple{N, Int} # original source axis of each reduced axis (0 = dummy) + dstlabels::NTuple{N, Int} # original destination axis of each reduced axis (0 = dummy) + descs::Vector{BatchedPermuteDesc{N}} + cpublocks::Vector{NTuple{N, Int}} # CPU-only sub-tile blocking per tensor (empty on GPU) + prefix::Vector{Int} # exclusive prefix sum of per-tensor tile counts, length B+1 + totaltiles::Int + totalelements::Int + deviceid::I + devcache::Base.RefValue{Any} # extension-owned slot for cached device-side state +end + +# Hooks a GPU extension overrides for its own device-identity type; CPU is `deviceid === nothing`. +_batched_deviceid(::StridedView) = nothing + +# CPU has a single execution strategy, so only `BP_AUTO` is legal; anything else asks for a GPU +# strategy on non-GPU arrays and is rejected rather than ignored. +_batched_resolve_strategy(::Nothing, requested::BatchedPermuteStrategy, ::BatchedPermuteFamily, ::Type) = + requested === BP_AUTO ? BP_AUTO : + throw(ArgumentError("batched_permutedims!: strategy=$requested applies only to GPU-backed batches")) + +# CPU tile shape: about `totalelements / (nthreads * 8)` elements per tile (clamped), filled +# greedily from the axis fastest on both sides (else the source-fastest axis), then in axis +# order. `strategy`, `dims` and `T` are unused here; they exist for the GPU override. +function _fillorder(qt::NTuple{N, Int}) where {N} + order = Int[1] + qt[1] != 1 && push!(order, qt[1]) + for g in 1:N + g in order || push!(order, g) + end + return order +end + +function _batched_tileshape( + ::Nothing, ::BatchedPermuteStrategy, ::BatchedPermuteFamily, + qt::NTuple{N, Int}, maxdims::NTuple{N, Int}, ::Vector{NTuple{N, Int}}, ::Type, + totalelements::Int, nthreads::Int + ) where {N} + cap = clamp(cld(totalelements, nthreads * 8), 1 << 10, 1 << 16) + H = fill(1, N) + p = 1 + for g in _fillorder(qt) + H[g] = clamp(cap ÷ p, 1, max(maxdims[g], 1)) + p *= H[g] + end + return ntuple(i -> H[i], N) +end + +_svtype(::Type{<:StridedView{T}}) where {T} = T +_cprod(t::Tuple) = foldl(Base.checked_mul, t; init = 1) + +# Common storage order per side: a scheduling heuristic only, never used for addressing (that +# always uses each tensor's real strides). Takes the stride order of the first tensor with no +# singleton axis (unambiguous); otherwise `1:N0`, which is always safe, just unoptimized. +function _infer_order(views::Vector{<:StridedView}, N0::Int) + for v in views + sz = size(v) + (prod(sz) == 0 || any(==(1), sz)) && continue + st = strides(v) + return Tuple(sort!(collect(1:N0); by = k -> st[k])) + end + return ntuple(identity, N0) +end + +# Axis reduction: an axis with extent 1 in every tensor is dropped. Axes are never merged -- +# fusing adjacent axes needs `stride[i] == stride[i-1] * extent[i-1]`, which is neither assumed +# nor checked. `q[j]` = heuristic-order source position filling destination position j. +# Returns (groups, qt, N, keep): `groups[i]` = reduced axis of heuristic position i (0 = +# dropped), `qt` = `q` relabeled to reduced axes, `keep[g]` = position that became reduced axis g. +function _collapse(nn::Vector{NTuple{N0, Int}}, q::NTuple{N0, Int}) where {N0} + B = length(nn) + dropped = ntuple(i -> all(b -> nn[b][i] == 1, 1:B), N0) + keep = Int[i for i in 1:N0 if !dropped[i]] + Nsurv = length(keep) + relabel = Dict(i => k for (k, i) in enumerate(keep)) + N = max(Nsurv, 1) + groups = ntuple(i -> dropped[i] ? 0 : relabel[i], N0) + qt = Nsurv == 0 ? (1,) : Tuple(relabel[v] for v in q if !dropped[v]) + return groups, qt, N, keep +end + +# Reduced extents of one tensor under a fixed `groups` (overflow-checked products; 1 for the +# dummy axis). `N` is a runtime `Int` and the `@inline` is load-bearing: inlined where `N` is a +# static plan parameter the `ntuple` folds to a fixed tuple with no allocation, whereas out of +# line it would return an abstract tuple and box it on every call. +@inline function _groupdims(nn_b::NTuple{N0, Int}, groups::NTuple{N0, Int}, N::Int) where {N0} + return ntuple(N) do g + p = 1 + for i in 1:N0 + groups[i] == g && (p = Base.checked_mul(p, nn_b[i])) + end + p + end +end + +# Original source/destination axis feeding each reduced axis; 0 marks the dummy axis of an +# all-dropped batch. `y[d] = x[perm[d]]`, so source label `s` lands on destination axis `invp[s]`. +function _axislabels(sigma::NTuple{N0, Int}, invp::NTuple{N0, Int}, keep::Vector{Int}, N::Int) where {N0} + srclabels = ntuple(g -> g <= length(keep) ? sigma[keep[g]] : 0, N) + dstlabels = ntuple(g -> srclabels[g] == 0 ? 0 : invp[srclabels[g]], N) + return srclabels, dstlabels +end + +# Each tensor's own strides at the labeled axes, nothing computed from extents; a dummy axis +# (label 0) gets stride 0 since its only coordinate is 0. +function _realstrides( + sview::StridedView, dview::StridedView, + srclabels::NTuple{N, Int}, dstlabels::NTuple{N, Int} + ) where {N} + sst = strides(sview) + dst = strides(dview) + srcstrides = ntuple(g -> srclabels[g] == 0 ? 0 : sst[srclabels[g]], N) + dststrides = ntuple(g -> dstlabels[g] == 0 ? 0 : dst[dstlabels[g]], N) + return srcstrides, dststrides +end + +@inline function _make_desc( + dview::StridedView, sview::StridedView, dims_b::NTuple{N, Int}, + srclabels::NTuple{N, Int}, dstlabels::NTuple{N, Int}, H::NTuple{N, Int} + ) where {N} + srcstrides_b, dststrides_b = _realstrides(sview, dview, srclabels, dstlabels) + return BatchedPermuteDesc{N}( + dims_b, dststrides_b, srcstrides_b, map(cld, dims_b, H), offset(dview), offset(sview) + ) +end + +# CPU-only cache blocking of one tile, via Strided's existing `_computeblocks` heuristic; a pure +# function of the descriptor, `H` and `sizeof(T)`, so an unchanged descriptor implies an +# unchanged blocking. +function _cpuoblock(desc::BatchedPermuteDesc{N}, H::NTuple{N, Int}, sizeofT::Int) where {N} + bytestrides = (sizeofT .* desc.dststrides, sizeofT .* desc.srcstrides) + costs = _computecosts((desc.dststrides, desc.srcstrides)) + strideorders = (indexorder(desc.dststrides), indexorder(desc.srcstrides)) + return _computeblocks(min.(H, desc.dims), costs, bytestrides, strideorders) +end + +# --- validation shared by fresh planning and plan reuse --- + +function _check_deviceid(views, devid) + for v in views + isequal(_batched_deviceid(v), devid) || + throw(ArgumentError("batched_permutedims!: mixed backends/devices")) + end + return nothing +end + +function _validate_deviceid(dviews, sviews) + devid = !isempty(dviews) ? _batched_deviceid(dviews[1]) : + !isempty(sviews) ? _batched_deviceid(sviews[1]) : nothing + _check_deviceid(dviews, devid) + _check_deviceid(sviews, devid) + return devid +end + +function _validate_sizes(dviews, sviews, perm::NTuple{N0, Int}) where {N0} + for b in 1:length(dviews) + szd = size(dviews[b]) + szs = size(sviews[b]) + for j in 1:N0 + szd[j] == szs[perm[j]] || + throw(DimensionMismatch("batched_permutedims!: size mismatch between dsts[$b] and srcs[$b]")) + end + end + return nothing +end + +# Not `map(StridedView, xs)`: that infers as `Union{Vector{Any}, Vector{StridedView{...}}}` +# (collect's type widening) even for a concrete `eltype(xs)`, turning every downstream loop +# into dynamic dispatch. Fixing `V` up front via `promote_op` behind a function barrier gives a +# concrete `Vector{V}`; a non-concrete `eltype(xs)` keeps the `map` path (run-time narrowing). +_normview_typed(::Type{V}, xs) where {V} = V[StridedView(x) for x in xs] +function _normview(xs::AbstractVector) + V = Base.promote_op(StridedView, eltype(xs)) + isconcretetype(V) && return _normview_typed(V, xs) + isempty(xs) && return Vector{V}() + return map(StridedView, xs) +end + +function _check_rank(views, N0::Int) + for v in views + ndims(v) == N0 || throw(DimensionMismatch("batched_permutedims!: rank mismatch")) + end + return nothing +end + +# This API only moves raw bits, so a view carrying `conj`/`adjoint`/... must be rejected. +function _check_op(views) + for v in views + v.op === identity || throw(ArgumentError("batched_permutedims!: only identity ops are supported")) + end + return nothing +end + +function _normalize_and_check(dsts::AbstractVector, srcs::AbstractVector, N0::Int) + dviews = _normview(dsts) + sviews = _normview(srcs) + isconcretetype(eltype(dviews)) && + isconcretetype(eltype(sviews)) || + throw(ArgumentError("batched_permutedims!: mixed parent array types are not supported")) + _check_rank(dviews, N0) + _check_rank(sviews, N0) + T = _svtype(eltype(dviews)) + T === _svtype(eltype(sviews)) || + throw(ArgumentError("batched_permutedims!: dsts and srcs must share one element type")) + (isbitstype(T) && sizeof(T) > 0) || + throw(ArgumentError("batched_permutedims!: element type must be a nonzero-size isbits type")) + _check_op(dviews) + _check_op(sviews) + return dviews, sviews, T +end + +function _prepare(dsts::AbstractVector, srcs::AbstractVector, N0::Int) + Base.require_one_based_indexing(dsts, srcs) + length(dsts) == length(srcs) || + throw(DimensionMismatch("batched_permutedims!: dsts and srcs must have equal length")) + return _normalize_and_check(dsts, srcs, N0) +end + +_incompatible() = ArgumentError("batched_permutedims!: arrays are incompatible with this plan") + +# --- per-call coefficients --- + +# The one formula every executor uses for `beta != 0`; `beta == 0` never reads `y` at all (so a +# NaN-poisoned or uninitialized destination stays clean) and computes plain `a * x`. +@inline _batched_axpby(a, x, b, y) = a * x + b * y + +# Both omitted -> `nothing` (the unscaled path, unchanged); otherwise a `(Vector{T}, Vector{T})` +# pair. That is a 2-valued compile-time distinction on `Nothing` (as `_mapreduce_kernel!` does +# for `op`), deliberately not a `Val`. Coefficients are converted to `T` once here, so a complex +# coefficient with nonzero imaginary part on a real batch throws `InexactError` from `convert`. +_batched_coeffs(::Nothing, ::Nothing, ::Int, ::Type) = nothing +function _batched_coeffs(alpha, beta, B::Int, ::Type{T}) where {T} + T <: Number || throw(ArgumentError("batched_permutedims!: alpha/beta require a Number element type, got $T")) + a = alpha === nothing ? ones(T, B) : _batched_coeffvec(alpha, B, T) + b = beta === nothing ? zeros(T, B) : _batched_coeffvec(beta, B, T) + return (a, b) +end +function _batched_coeffvec(x::AbstractVector, B::Int, ::Type{T}) where {T} + length(x) == B || throw(DimensionMismatch("batched_permutedims!: alpha/beta must have length(dsts) entries")) + return convert(Vector{T}, x) +end + +# --- fresh plan --- + +function _plan_impl( + dsts::AbstractVector, srcs::AbstractVector, perm::NTuple{N0, Int}; + strategy::BatchedPermuteStrategy = BP_AUTO + ) where {N0} + isperm(perm) || throw(ArgumentError("batched_permutedims!: perm is not a valid permutation")) + dviews, sviews, T = _prepare(dsts, srcs, N0) + deviceid = _validate_deviceid(dviews, sviews) + _validate_sizes(dviews, sviews, perm) + sigma = _infer_order(sviews, N0) + tau = _infer_order(dviews, N0) + return _build_descriptors(dviews, sviews, perm, sigma, tau, deviceid, T, strategy) +end + +function _build_descriptors( + dviews::Vector{<:StridedView}, sviews::Vector{<:StridedView}, + perm::NTuple{N0, Int}, sigma::NTuple{N0, Int}, tau::NTuple{N0, Int}, + deviceid, ::Type{T}, strategy::BatchedPermuteStrategy + ) where {N0, T} + B = length(dviews) + invsigma = invperm(sigma) + invp = invperm(perm) + q = ntuple(j -> invsigma[perm[tau[j]]], N0) + nn = [ntuple(i -> size(sviews[b])[sigma[i]], N0) for b in 1:B] + groups, qt, N, keep = _collapse(nn, q) + srclabels, dstlabels = _axislabels(sigma, invp, keep, N) + dimsvec = Vector{NTuple{N, Int}}(undef, B) # typed explicitly: `N` is a runtime value here + for b in 1:B + dimsvec[b] = _groupdims(nn[b], groups, N) + end + maxdims = isempty(dimsvec) ? ntuple(_ -> 0, N) : reduce((a, c) -> map(max, a, c), dimsvec) + totalelements = foldl((s, d) -> Base.checked_add(s, _cprod(d)), dimsvec; init = 0) + family = N == 1 ? BP_COPY : (qt[1] == 1 ? BP_PAYLOAD : BP_TRANSPOSE) + resolved = _batched_resolve_strategy(deviceid, strategy, family, T) + H = _batched_tileshape( + deviceid, resolved, family, qt, maxdims, dimsvec, T, totalelements, get_num_threads() + ) + iscpu = deviceid === nothing + descs = Vector{BatchedPermuteDesc{N}}(undef, B) + cpublocks = iscpu ? Vector{NTuple{N, Int}}(undef, B) : NTuple{N, Int}[] + prefix = Vector{Int}(undef, B + 1) + prefix[1] = 0 + for b in 1:B + desc = descs[b] = _make_desc(dviews[b], sviews[b], dimsvec[b], srclabels, dstlabels, H) + prefix[b + 1] = Base.checked_add(prefix[b], _cprod(desc.tilecounts)) + iscpu && (cpublocks[b] = _cpuoblock(desc, H, sizeof(T))) + end + return BatchedPermutePlan{N0, N, T, typeof(deviceid)}( + perm, sigma, groups, qt, family, resolved, H, srclabels, dstlabels, + descs, cpublocks, prefix, prefix[B + 1], totalelements, + deviceid, Base.RefValue{Any}(nothing) + ) +end + +# --- plan reuse --- + +# Reuse requires matching shapes only (tile counts/prefix sums depend on them); strides and +# offsets are re-read from the arrays passed in. Order inference, collapsing, tile shape and +# strategy are trusted from the plan. Fast path: every descriptor is recomputed and compared +# bitwise (`BatchedPermuteDesc` is `isbits`); while all match, the cached vectors are reused and +# the same plan object is returned with no allocation. From the first differing tensor on, fresh +# `descs` (and, on CPU, `cpublocks`) vectors are built, copying the unchanged prefix. Sound only +# because `plan.descs`/`plan.cpublocks` are never mutated in place anywhere. +function _rebuild_descriptors( + dviews::Vector{<:StridedView}, sviews::Vector{<:StridedView}, + plan::BatchedPermutePlan{N0, N, T, I} + ) where {N0, N, T, I} + B = length(dviews) + olddescs = plan.descs + oldblocks = plan.cpublocks + B == length(olddescs) || throw(_incompatible()) + descs = olddescs # replaced by a fresh vector at the first differing tensor + cpublocks = oldblocks # likewise (CPU plans only) + iscpu = plan.deviceid === nothing + for b in 1:B + nn_b = ntuple(i -> size(sviews[b])[plan.srcorder[i]], N0) + dims_b = _groupdims(nn_b, plan.groups, N) + dims_b == olddescs[b].dims || throw(_incompatible()) + desc_b = _make_desc(dviews[b], sviews[b], dims_b, plan.srclabels, plan.dstlabels, plan.tile) + if descs === olddescs + desc_b == olddescs[b] && continue + descs = copyto!(Vector{BatchedPermuteDesc{N}}(undef, B), 1, olddescs, 1, b - 1) + iscpu && (cpublocks = copyto!(Vector{NTuple{N, Int}}(undef, B), 1, oldblocks, 1, b - 1)) + end + descs[b] = desc_b + iscpu && (cpublocks[b] = _cpuoblock(desc_b, plan.tile, sizeof(T))) + end + descs === olddescs && return plan + return BatchedPermutePlan{N0, N, T, I}( + plan.perm, plan.srcorder, plan.groups, plan.dstseq, plan.family, plan.strategy, + plan.tile, plan.srclabels, plan.dstlabels, descs, cpublocks, plan.prefix, + plan.totaltiles, plan.totalelements, plan.deviceid, plan.devcache + ) +end + +# --- public API --- + +""" + Strided.plan_batched_permutedims(dsts::AbstractVector, srcs::AbstractVector, perm; + strategy::BatchedPermuteStrategy = BP_AUTO) -> BatchedPermutePlan + +Plan a batched out-of-place permutation `dsts[b] = permutedims(srcs[b], perm)` for every `b`. +The plan can be passed to `batched_permutedims!` any number of times, including with different +arrays of the same shapes. + +`strategy` is resolved once here; the plan's `.strategy` field reports the resolved value (never +`BP_AUTO`, except on CPU-backed batches, which have a single strategy and only accept `BP_AUTO`; +anything else is an error there). See `BatchedPermuteStrategy`. + +Addressing uses each array's own strides, so non-dense and negative-stride views are supported. +There is no aliasing check: overlapping destinations, or a destination overlapping a source, is +undefined behavior. Sources may alias each other. +""" +function plan_batched_permutedims( + dsts::AbstractVector, srcs::AbstractVector, perm::NTuple{N0, Int}; + strategy::BatchedPermuteStrategy = BP_AUTO + ) where {N0} + return _plan_impl(dsts, srcs, perm; strategy) +end +function plan_batched_permutedims( + dsts::AbstractVector, srcs::AbstractVector, perm::AbstractVector{<:Integer}; + strategy::BatchedPermuteStrategy = BP_AUTO + ) + return _plan_impl(dsts, srcs, ntuple(i -> Int(perm[i]), length(perm)); strategy) +end + +""" + Strided.batched_permutedims!(dsts::AbstractVector, srcs::AbstractVector, perm; + strategy::BatchedPermuteStrategy = BP_AUTO, + alpha = nothing, beta = nothing) -> dsts + Strided.batched_permutedims!(dsts::AbstractVector, srcs::AbstractVector, plan::BatchedPermutePlan; + alpha = nothing, beta = nothing) -> dsts + +Batched out-of-place permutation `dsts[b] .= permutedims(srcs[b], perm)` for every `b`, or, with +coefficients, `dsts[b] .= alpha[b] .* permutedims(srcs[b], perm) .+ beta[b] .* dsts[b]`. The +`perm` form is `plan_batched_permutedims` (forwarding `strategy`) followed by the `plan` form, +which takes no `strategy` keyword since the plan carries its resolved one. + +`alpha` and `beta` are optional per-tensor coefficient vectors of length `length(dsts)`. An +omitted `alpha` means all ones, an omitted `beta` all zeros, and omitting both is the plain copy +with no arithmetic at all. Coefficients are converted to the batch's element type `T` (which +must then be a `Number`): a real coefficient on a complex batch becomes `a + 0im`, and a complex +coefficient with a nonzero imaginary part on a real batch throws an `InexactError`. Wherever +`beta[b] == 0` the previous contents of `dsts[b]` are never read, so an uninitialized destination +is safe there; `alpha[b] == 1` gets no special treatment. The scaled result is ordinary +floating-point arithmetic in `T` and is not guaranteed to agree bitwise between the CPU and a +GPU backend (a GPU may fuse the multiply-add), unlike the coefficient-free copy, which moves bits. + +On a GPU with `beta[b] != 0` for a `BP_TRANSPOSE` batch, `strategy = BP_GROUPTILE` reads the old +destination through shared memory and is substantially faster than the `BP_AUTO`/`BP_ELEMENTWISE` +default, which reads it through the same strided access pattern as the (already slower) transpose +store. + +Not safe for aliased inputs: destinations must be pairwise disjoint and disjoint from every +source, or the result is unspecified; this is never checked. Sharing one plan across concurrent +tasks with *different* arrays on a GPU backend is not thread-safe (the plan's device cache is +mutated in place on reuse); concurrent reuse with the same arrays is safe. +""" +function batched_permutedims!( + dsts::AbstractVector, srcs::AbstractVector, perm; strategy::BatchedPermuteStrategy = BP_AUTO, + alpha::Union{Nothing, AbstractVector} = nothing, beta::Union{Nothing, AbstractVector} = nothing + ) + plan = plan_batched_permutedims(dsts, srcs, perm; strategy) + return batched_permutedims!(dsts, srcs, plan; alpha, beta) +end + +function batched_permutedims!( + dsts::AbstractVector, srcs::AbstractVector, plan::BatchedPermutePlan{N0, N, T, I}; + alpha::Union{Nothing, AbstractVector} = nothing, beta::Union{Nothing, AbstractVector} = nothing + ) where {N0, N, T, I} + dviews, sviews, Tact = _prepare(dsts, srcs, N0) + Tact === T || throw(_incompatible()) + coeffs = _batched_coeffs(alpha, beta, length(dviews), T) + deviceid = _validate_deviceid(dviews, sviews) + isequal(deviceid, plan.deviceid) || throw(_incompatible()) + _validate_sizes(dviews, sviews, plan.perm) + newplan = _rebuild_descriptors(dviews, sviews, plan) + newplan.totaltiles == 0 || _batched_permute!(newplan, dviews, sviews, coeffs) + return dsts +end + +# `_batched_permute!(plan, dst::Vector{<:StridedView}, src::Vector{<:StridedView}, coeffs)` is +# defined in `batched_permutedims_cpu.jl`; the GPU extension adds a method for its own view +# subtype. `coeffs` is `nothing` or the `(alpha, beta)` pair from `_batched_coeffs`. diff --git a/src/batched_permutedims_cpu.jl b/src/batched_permutedims_cpu.jl new file mode 100644 index 0000000..b08b714 --- /dev/null +++ b/src/batched_permutedims_cpu.jl @@ -0,0 +1,144 @@ +# CPU executor: `_batched_permute!` for plain `StridedView`s (`plan.deviceid === nothing`). +# Global tile ids `0:totaltiles-1` are consecutive per tensor (tensor b owns +# `prefix[b]:prefix[b+1]-1`); each is located, decoded and copied as one tile. + +# Decode local tile id `l` (mixed radix over `tilecounts`, axis 1 fastest) into stride-multiplied +# source/destination offsets and the tile's extents (`H[g]`, or the remainder at the tensor's +# edge). Recursion over `Base.tail` unrolls fully for a static `N`: one specialization per +# reduced rank, none per shape or batch. +@inline function _tile_geom( + dims::NTuple{N, Int}, tilecounts::NTuple{N, Int}, srcstrides::NTuple{N, Int}, + dststrides::NTuple{N, Int}, H::NTuple{N, Int}, l::Int + ) where {N} + tc = tilecounts[1] + o = (l % tc) * H[1] + r = l ÷ tc + d1 = min(H[1], dims[1] - o) + soffrest, doffrest, drest = _tile_geom( + Base.tail(dims), Base.tail(tilecounts), Base.tail(srcstrides), + Base.tail(dststrides), Base.tail(H), r + ) + return o * srcstrides[1] + soffrest, o * dststrides[1] + doffrest, (d1, drest...) +end +@inline _tile_geom(::Tuple{}, ::Tuple{}, ::Tuple{}, ::Tuple{}, ::Tuple{}, l::Int) = (0, 0, ()) + +# Absolute source/destination base offsets and extents of tile `l` of one tensor. Shared by the +# CPU runner and both GPU kernels, so every executor agrees on tile geometry by construction. +@inline function _tile_offsets(desc::BatchedPermuteDesc{N}, H::NTuple{N, Int}, l::Int) where {N} + dsoff, ddoff, d = _tile_geom(desc.dims, desc.tilecounts, desc.srcstrides, desc.dststrides, H, l) + return desc.srcoffset + dsoff, desc.dstoffset + ddoff, d +end + +# One tile through Strided's serial tiled kernel (no reduction, so each destination element is +# written once). Unscaled: `f = identity` is exactly a strided copy. Scaled with `beta == 0`: +# `a * x`, still never reading `dst`. Otherwise `dst` is also the third input, read and +# rewritten at the same index in one call, as Strided's own `axpby!` broadcast does. +@inline function _batched_tile_kernel!( + ::Nothing, ::Int, d, blocks, dst::StridedView, src::StridedView, dstr, sstr, doff::Int, soff::Int + ) + _mapreduce_kernel!(identity, nothing, nothing, d, blocks, (dst, src), (dstr, sstr), (doff, soff)) + return nothing +end +@inline function _batched_tile_kernel!( + coeffs::Tuple, b::Int, d, blocks, dst::StridedView, src::StridedView, dstr, sstr, doff::Int, soff::Int + ) + a = coeffs[1][b] + bb = coeffs[2][b] + if iszero(bb) + _mapreduce_kernel!(Base.Fix1(*, a), nothing, nothing, d, blocks, (dst, src), (dstr, sstr), (doff, soff)) + else + _mapreduce_kernel!( + (x, y) -> _batched_axpby(a, x, bb, y), nothing, nothing, d, blocks, + (dst, src, dst), (dstr, sstr, dstr), (doff, soff, doff) + ) + end + return nothing +end + +@inline function _batched_run_tile!( + dst::Vector{<:StridedView}, src::Vector{<:StridedView}, + plan::BatchedPermutePlan, b::Int, l::Int, coeffs + ) + desc = plan.descs[b] + soff, doff, d = _tile_offsets(desc, plan.tile, l) + _batched_tile_kernel!( + coeffs, b, d, plan.cpublocks[b], dst[b], src[b], desc.dststrides, desc.srcstrides, doff, soff + ) + return nothing +end + +# Run global tile ids `[ustart, uend)`, which may span tensors: locate the first tensor once, +# then advance `b` across prefix boundaries instead of re-searching per tile. +function _batched_chunk!( + dst::Vector{<:StridedView}, src::Vector{<:StridedView}, + plan::BatchedPermutePlan, ustart::Int, uend::Int, coeffs + ) + ustart >= uend && return nothing + prefix = plan.prefix + b = searchsortedlast(prefix, ustart) + for u in ustart:(uend - 1) + while u >= prefix[b + 1] + b += 1 + end + _batched_run_tile!(dst, src, plan, b, u - prefix[b], coeffs) + end + return nothing +end + +# Worker boundaries balancing *elements*, not tile counts: tiles differ in size across tensors +# and are contiguous per tensor, so equal id ranges can hand nearly all the work to one worker. +# A tensor straddling a worker's element budget is split proportionally (assuming roughly +# equal tiles, which holds except at its edge; this only picks boundaries, never addresses). +function _batched_worker_bounds(plan::BatchedPermutePlan, nw::Int) + B = length(plan.descs) + totaltiles = plan.totaltiles + total = plan.totalelements + bounds = Vector{Int}(undef, nw + 1) + bounds[1] = 0 + bounds[end] = totaltiles + b = 1 + cum = 0 # elements in tensors fully before tensor b + for i in 1:(nw - 1) + target = (total * i) ÷ nw + while true + elemsb = _cprod(plan.descs[b].dims) + Tb = plan.prefix[b + 1] - plan.prefix[b] + if b < B && cum + elemsb <= target + cum += elemsb + b += 1 + continue + end + pertile = Tb == 0 ? 0.0 : elemsb / Tb + localtiles = pertile <= 0 ? 0 : round(Int, (target - cum) / pertile) + bounds[i + 1] = plan.prefix[b] + clamp(localtiles, 0, Tb) + break + end + end + for i in 2:nw # rounding could produce a tiny local decrease; disallow it + bounds[i] = clamp(bounds[i], bounds[i - 1], totaltiles) + end + return bounds +end + +# Below `MINTHREADLENGTH` or with one worker, run every tile on the calling task; otherwise split +# into `min(nthreads, totaltiles)` element-balanced ranges, one task each (the caller runs the +# first). No per-worker scratch and no state keyed by `threadid()`, so task migration is harmless. +function _batched_permute!( + plan::BatchedPermutePlan, dst::Vector{<:StridedView}, src::Vector{<:StridedView}, coeffs + ) + totaltiles = plan.totaltiles + nw = get_num_threads() + if nw == 1 || plan.totalelements <= MINTHREADLENGTH + _batched_chunk!(dst, src, plan, 0, totaltiles, coeffs) + else + nw = min(nw, totaltiles) + bounds = _batched_worker_bounds(plan, nw) + @sync begin + for i in 2:nw + Threads.@spawn _batched_chunk!(dst, src, plan, bounds[i], bounds[i + 1], coeffs) + end + _batched_chunk!(dst, src, plan, bounds[1], bounds[2], coeffs) + end + end + return dst +end diff --git a/test/batched_permutedims.jl b/test/batched_permutedims.jl new file mode 100644 index 0000000..7278bd5 --- /dev/null +++ b/test/batched_permutedims.jl @@ -0,0 +1,669 @@ +# Independent reference/oracle and adversarial fixtures for +# Strided.batched_permutedims!/plan_batched_permutedims. This file is self-contained (run +# via `include`), does not depend on runtests.jl, and must never consult planner +# internals: only `perm` (Base's `permutedims` convention) and plain array indexing are +# used to build the oracle. +# +# Note on scope: this planner does not check that inputs are densely packed, and does not +# check aliasing at all (see the module docstring on `plan_batched_permutedims`). So +# non-dense and negative-stride views are exercised below as *correctness* fixtures (they +# are genuinely, correctly supported, not merely tolerated), while aliasing/overlap and +# malformed-view cases are deliberately NOT exercised here as error-throwing tests, since +# they no longer throw and actually running them would either produce an unspecified +# (silently wrong) result or, for a view that doesn't fit inside its own parent, write out +# of bounds -- there is nothing safe to assert about either outcome. + +using Test +using Random +using LinearAlgebra +using Strided +using Strided: StridedView + +Random.seed!(20240915) + +# ---------------------------------------------------------------------------- +# Reference oracle: two independent computations of the same +# thing, cross-checked against each other and never against planner code. +# ---------------------------------------------------------------------------- + +# explicit loop +function ref_loop(src::AbstractArray, perm) + N0 = length(perm) + out = Array{eltype(src)}(undef, ntuple(j -> size(src, perm[j]), N0)) + for I in CartesianIndices(size(src)) + out[ntuple(j -> I[perm[j]], N0)...] = src[I] + end + return out +end + +# second, independent cross-check via Base.permutedims +function ref_base(src::AbstractArray, perm) + return permutedims(Array(src), perm) +end + +# D11: bit-preserving comparison, never isapprox. Reuses only the `compare` +# *structure* of test/gpu.jl (not its isapprox tolerance semantics). +refequal(a::AbstractArray, b::AbstractArray) = isequal(Array(a), Array(b)) + +# builds the oracle output for one batch and self-checks the two independent +# routes agree, so a bug in `ref_loop` cannot silently pass as ground truth +function refbatch(srcs, perm) + refs = map(s -> ref_loop(s, perm), srcs) + for (s, r) in zip(srcs, refs) + @test refequal(ref_base(s, perm), r) + end + return refs +end + +# full round-trip harness for a *valid* fixture: compute the oracle, run the +# public API, check identity of the return value, check dst == oracle, and +# check srcs were not mutated (read-only sources, D5/section 9). +function runcase!(dsts, srcs, perm) + srcsnap = map(deepcopy, srcs) + refs = refbatch(srcs, perm) + out = Strided.batched_permutedims!(dsts, srcs, perm) + @test out === dsts + for (d, r) in zip(dsts, refs) + @test refequal(d, r) + end + for (s, s0) in zip(srcs, srcsnap) + @test refequal(s, s0) + end + return nothing +end + +@testset "Strided.batched_permutedims! (reference oracle self-check)" begin + # sanity: the oracle itself is correct before it judges anyone else's code. + # (also duplicated standalone in the scratchpad, run against plain Base.) + A = rand(3, 4, 5) + p = (3, 1, 2) + @test refequal(ref_loop(A, p), permutedims(A, p)) + @test refequal(ref_loop(A, p), ref_base(A, p)) + B = rand(ComplexF64, 2, 0, 4) # empty array, still a valid oracle input + @test refequal(ref_loop(B, (2, 3, 1)), permutedims(B, (2, 3, 1))) +end + +@testset "Strided.batched_permutedims! (correctness fixtures)" begin + @testset "empty batch (B == 0)" begin + dsts = Matrix{Float64}[] + srcs = Matrix{Float64}[] + out = Strided.batched_permutedims!(dsts, srcs, (2, 1)) + @test out === dsts + end + + @testset "all tensors empty (totaltiles == 0)" begin + srcs = [zeros(3, 0, 4), zeros(0, 2, 4), zeros(3, 2, 0)] + dsts = [zeros(size(s, 2), size(s, 1), size(s, 3)) for s in srcs] + runcase!(dsts, srcs, (2, 1, 3)) + end + + @testset "empty tensors interleaved (repeated prefix values)" begin + srcs = Any[rand(3, 4), zeros(0, 4), zeros(3, 0), rand(2, 5), zeros(0, 0)] + perm = (2, 1) + dsts = Any[zeros(size(s, 2), size(s, 1)) for s in srcs] + runcase!(dsts, srcs, perm) + end + + @testset "rank 0 (perm == ())" begin + srcs = [fill(1.5), fill(-2.0), fill(NaN)] + dsts = [fill(0.0), fill(0.0), fill(0.0)] + runcase!(dsts, srcs, ()) + end + + @testset "singleton axis: present in some tensors only (must not be dropped)" begin + s1 = rand(3, 1, 4) # axis 2 singleton + s2 = rand(3, 2, 4) # axis 2 not singleton + perm = (3, 1, 2) + srcs = [s1, s2] + dsts = [zeros(ntuple(j -> size(s, perm[j]), 3)) for s in srcs] + runcase!(dsts, srcs, perm) + end + + @testset "singleton axis: present in every tensor (dropped internally)" begin + s1 = rand(3, 1, 4) + s2 = rand(5, 1, 4) + perm = (3, 1, 2) + srcs = [s1, s2] + dsts = [zeros(ntuple(j -> size(s, perm[j]), 3)) for s in srcs] + runcase!(dsts, srcs, perm) + end + + @testset "singleton axis with a fabricated stride still addresses correctly" begin + # StridedViews rewrites a singleton axis's stride to whatever is convenient + # internally; here permuting turns axis 2 (originally stride 4) into a singleton + # axis with a fabricated stride of 16. Since that axis's only coordinate is 0, + # its stride value (fabricated or not) never actually gets multiplied by + # anything nonzero, so this is harmless -- demonstrated directly rather than + # asserted as an invariant. + A = rand(4, 1, 4) # strides(StridedView(A)) == (1, 4, 4) + sv = StridedView(A) + wv = permutedims(sv, (3, 2, 1)) + @test strides(wv) == (4, 16, 1) + dst = zeros(size(wv)) + runcase!([dst], [wv], (1, 2, 3)) + end + + @testset "perm given as AbstractVector{<:Integer}, not NTuple" begin + s = rand(2, 3, 4) + perm = [2, 3, 1] + dst = zeros(3, 4, 2) + runcase!([dst], [s], perm) + end + + @testset "real transpose/adjoint normalize to identity op (accepted)" begin + A = rand(4, 5) + srcT = transpose(A) + dst = zeros(5, 4) + runcase!([dst], [srcT], (1, 2)) # identity perm on the transposed view + B = rand(4, 5) + srcA = adjoint(B) # real adjoint == transpose, still identity op + dst2 = zeros(5, 4) + runcase!([dst2], [srcA], (1, 2)) + end + + @testset "nonzero offsets, disjoint views of one parent" begin + parent = collect(1.0:100.0) + src = reshape(view(parent, 1:24), 4, 6) + dst = reshape(view(parent, 25:48), 6, 4) + runcase!([dst], [src], (2, 1)) + end + + @testset "srcs[b] === srcs[b'] accepted (source/source overlap allowed)" begin + s = rand(3, 4) + srcs = [s, s] + dsts = [zeros(4, 3), zeros(4, 3)] + runcase!(dsts, srcs, (2, 1)) + end + + @testset "order-inference heuristic falls back cleanly when no tensor is unambiguous" begin + # Every tensor here has a singleton axis, so `_infer_order` never finds a + # fully-unambiguous tensor to copy the order from and falls back to the plain + # 1:N0 axis order. Correctness must not depend on that choice: addressing always + # uses each tensor's own real strides, regardless of which heuristic order the + # planner picked for scheduling. + A = permutedims(StridedView(rand(5, 3, 1)), (1, 3, 2)) # size (5,1,3), dense + B = permutedims(StridedView(rand(4, 1, 3)), (2, 1, 3)) # size (1,4,3), dense + perm = (1, 2, 3) + srcs = [A, B] + dsts = [zeros(size(s)) for s in srcs] + runcase!(dsts, srcs, perm) + end + + @testset "non-dense (strided) view: correctly supported, not merely tolerated" begin + s = view(rand(10), 1:2:9) # stride 2, not unit-stride + dst = zeros(5) + runcase!([dst], [s], (1,)) + end + + @testset "negative-stride view: correctly supported" begin + A = rand(6) + s = @view A[6:-1:1] + dst = zeros(6) + runcase!([dst], [s], (1,)) + end + + @testset "dominant tensor plus many tiny ones, spread of sizes" begin + srcs = Any[rand(97, 131)] + for n in (1, 2, 3, 7, 8, 9, 15, 16, 17, 63, 64, 65) + push!(srcs, rand(n, n + 1)) + end + perm = (2, 1) + dsts = Any[zeros(size(s, 2), size(s, 1)) for s in srcs] + runcase!(dsts, srcs, perm) + end + + @testset "tile-edge spread, rank 3, transpose-family-shaped batch" begin + srcs = Any[] + for a in (1, 2, 8, 9, 63, 64, 65), b in (1, 3, 16, 17) + push!(srcs, rand(a, b, 2)) + end + perm = (3, 1, 2) + dsts = Any[zeros(ntuple(j -> size(s, perm[j]), 3)) for s in srcs] + runcase!(dsts, srcs, perm) + end +end + +@testset "Strided.batched_permutedims! (plan reuse)" begin + s1 = rand(6, 7) + d1 = zeros(7, 6) + perm = (2, 1) + plan = Strided.plan_batched_permutedims([d1], [s1], perm) + refs = refbatch([s1], perm) + dsts1 = [d1] + out = Strided.batched_permutedims!(dsts1, [s1], plan) + @test out === dsts1 + @test refequal(d1, refs[1]) + dsts1b = [d1] + out2 = Strided.batched_permutedims!(dsts1b, [s1], perm) + @test out2 === dsts1b + @test refequal(d1, refs[1]) + + # plan reused with a *different* but shape/stride-compatible array pair + s2 = rand(6, 7) + d2 = zeros(7, 6) + refs2 = refbatch([s2], perm) + Strided.batched_permutedims!([d2], [s2], plan) + @test refequal(d2, refs2[1]) +end + +@testset "Strided.batched_permutedims! (validation: DimensionMismatch)" begin + s = rand(3, 4) + d = zeros(4, 3) + + @testset "length(dsts) != length(srcs)" begin + @test_throws DimensionMismatch Strided.batched_permutedims!([d, d], [s], (2, 1)) + end + + @testset "ndims != N0" begin + s3 = rand(3, 4, 1) + @test_throws DimensionMismatch Strided.batched_permutedims!([d], [s3], (2, 1)) + end + + @testset "size(dst)[j] != size(src)[perm[j]]" begin + badd = zeros(3, 4) # should be (4,3) for perm (2,1) + @test_throws DimensionMismatch Strided.batched_permutedims!([badd], [s], (2, 1)) + end +end + +@testset "Strided.batched_permutedims! (validation: ArgumentError)" begin + @testset "invalid perm (not isperm)" begin + s = rand(3, 4) + d = zeros(4, 3) + @test_throws ArgumentError Strided.batched_permutedims!([d], [s], (1, 1)) + end + + @testset "mixed element type within one side (non-concrete eltype)" begin + s1 = rand(Float64, 3, 4) + s2 = rand(Float32, 3, 4) + d1 = zeros(4, 3) + d2 = zeros(4, 3) + @test_throws ArgumentError Strided.batched_permutedims!([d1, d2], Any[s1, s2], (2, 1)) + end + + @testset "matching eltype within each side, mismatched across sides" begin + s = rand(Float64, 3, 4) + d = zeros(Float32, 4, 3) + @test_throws ArgumentError Strided.batched_permutedims!([d], [s], (2, 1)) + end + + @testset "non-identity op: adjoint/conj of complex array rejected" begin + A = rand(ComplexF64, 4, 5) + d = zeros(ComplexF64, 5, 4) + @test_throws ArgumentError Strided.batched_permutedims!([d], [adjoint(A)], (1, 2)) + @test_throws ArgumentError Strided.batched_permutedims!([d], Any[conj(StridedView(A))], (1, 2)) + end + + # Aliasing (dsts[b] === srcs[b], overlapping destinations, overlapping src/dst into + # one parent), non-dense views, negative strides, and a view that doesn't fit inside + # its own parent are all UNCHECKED here on purpose (see the file header and + # `plan_batched_permutedims`'s docstring). Non-dense/negative-stride views are + # exercised as *correctness* fixtures above, not here; aliasing and malformed-parent + # cases are not exercised at all, since there is nothing safe to assert about their + # outcome once the checks are gone (see also: a genuinely cyclic/incompatible common + # order between two tensors used to be rejected here too, but order inference is now + # only a scheduling heuristic that falls back silently -- see "order-inference + # heuristic falls back cleanly..." above, which exercises exactly that fixture as a + # correctness case instead of an error case). + + @testset "Tuple collections get no method" begin + s = rand(3, 4) + d = zeros(4, 3) + @test_throws Union{MethodError, UndefVarError} Strided.batched_permutedims!((d,), (s,), (2, 1)) + end +end + +@testset "Strided.batched_permutedims! (metadata-only OverflowError, no allocation)" begin + # Real backing parents are tiny (1 element); declared size/strides deliberately + # overflow Int64 in a way that wraps to a value <= the (tiny) real parent length, + # so V11 containment passes on the wrapped product and the genuine overflow is + # only caught later by V13's checked_* arithmetic. No large array is allocated. + n1 = 2^62 + n2 = 4 # n1 * n2 == 2^64, wraps to 0 in Int64 arithmetic + srcparent = zeros(1) + dstparent = zeros(1) + src = StridedView(srcparent, (n1, n2), (1, n1), 0) + dst = StridedView(dstparent, (n1, n2), (1, n1), 0) + @test_throws OverflowError Strided.plan_batched_permutedims([dst], [src], (1, 2)) +end + +@testset "Strided.plan_batched_permutedims (strategy keyword, CPU-backed batches)" begin + @testset "default strategy resolves to BP_AUTO" begin + s = rand(3, 4) + d = zeros(4, 3) + plan = Strided.plan_batched_permutedims([d], [s], (2, 1)) + @test plan.strategy == Strided.BP_AUTO + end + + @testset "explicit BP_AUTO on CPU arrays succeeds and resolves to BP_AUTO" begin + s = rand(3, 4) + d = zeros(4, 3) + plan = Strided.plan_batched_permutedims([d], [s], (2, 1); strategy = Strided.BP_AUTO) + @test plan.strategy == Strided.BP_AUTO + refs = refbatch([s], (2, 1)) + Strided.batched_permutedims!([d], [s], plan) + @test refequal(d, refs[1]) + end + + @testset "non-AUTO strategy on CPU arrays is rejected before any write" begin + for strategy in (Strided.BP_ELEMENTWISE, Strided.BP_THREADTILE, Strided.BP_GROUPTILE) + s = rand(3, 4) + d = zeros(4, 3) + @test_throws ArgumentError Strided.plan_batched_permutedims([d], [s], (2, 1); strategy) + @test_throws ArgumentError Strided.batched_permutedims!([d], [s], (2, 1); strategy) + @test all(iszero, d) # rejected before any write reached dst + end + end + + @testset "plan reuse preserves the resolved strategy unchanged" begin + s1 = rand(6, 7) + d1 = zeros(7, 6) + perm = (2, 1) + plan = Strided.plan_batched_permutedims([d1], [s1], perm; strategy = Strided.BP_AUTO) + @test plan.strategy == Strided.BP_AUTO + + # rebuild directly against a different-but-compatible array pair, and check the + # rebuilt plan's strategy field matches the original plan's, unchanged + s2 = rand(6, 7) + d2 = zeros(7, 6) + dviews = map(StridedView, [d2]) + sviews = map(StridedView, [s2]) + newplan = Strided._rebuild_descriptors(dviews, sviews, plan) + @test newplan.strategy == plan.strategy == Strided.BP_AUTO + + # and via the public reuse path, executing still works and the plan is unaffected + refs2 = refbatch([s2], perm) + Strided.batched_permutedims!([d2], [s2], plan) + @test refequal(d2, refs2[1]) + @test plan.strategy == Strided.BP_AUTO + end +end + +@testset "Strided._rebuild_descriptors (plan reuse: same-descriptor fast path)" begin + perm = (3, 1, 2) + N0 = length(perm) + srcs = [rand(4, 5, 6), rand(2, 3, 7), rand(1, 8, 2), rand(6, 6, 6)] + dsts = [zeros(ntuple(j -> size(s, perm[j]), N0)) for s in srcs] + plan = Strided.plan_batched_permutedims(dsts, srcs, perm) + # a uniformly-typed view vector, so one entry can be swapped for a view with a + # different offset/strides without changing the vector's element type + sviews0 = map(StridedView, srcs) + dviews0 = map(StridedView, dsts) + + @testset "same arrays again: the identical plan object comes back" begin + dviews, sviews, _ = Strided._normalize_and_check(dsts, srcs, N0) + same = Strided._rebuild_descriptors(dviews, sviews, plan) + @test same === plan + @test same.descs === plan.descs && same.cpublocks === plan.cpublocks + # and it agrees with a freshly built reference plan, field for field + fresh = Strided.plan_batched_permutedims(dsts, srcs, perm) + @test same.descs == fresh.descs && same.cpublocks == fresh.cpublocks + @test same.prefix == fresh.prefix && same.totaltiles == fresh.totaltiles + # fully inferred and allocation-free on this path + @test (@inferred Strided._rebuild_descriptors(dviews, sviews, plan)) === plan + rebuild(dv, sv, p) = Strided._rebuild_descriptors(dv, sv, p) + rebuild(dviews, sviews, plan) + @test (@allocated rebuild(dviews, sviews, plan)) == 0 + end + + @testset "one source swapped for a view with a different offset: new plan, prefix copied" begin + big = rand(20, 3, 7) + sviews2 = copy(sviews0) + sviews2[2] = view(StridedView(big), 9:10, :, :) # same shape as srcs[2]; offset 8, strides (1,20,60) + @test size(sviews2[2]) == size(srcs[2]) && Strided.offset(sviews2[2]) != 0 + changed = Strided._rebuild_descriptors(dviews0, sviews2, plan) + @test changed !== plan + @test changed.descs !== plan.descs && changed.cpublocks !== plan.cpublocks + @test changed.descs[2] != plan.descs[2] + @test changed.descs[[1, 3, 4]] == plan.descs[[1, 3, 4]] + @test changed.cpublocks[[1, 3, 4]] == plan.cpublocks[[1, 3, 4]] + ref2 = Strided.plan_batched_permutedims(dviews0, sviews2, perm) + @test changed.descs == ref2.descs && changed.cpublocks == ref2.cpublocks + # the original plan is untouched + @test plan.descs == Strided.plan_batched_permutedims(dsts, srcs, perm).descs + # and the public reuse path with the swapped arrays computes the right thing + refs = refbatch(sviews2, perm) + Strided.batched_permutedims!(dviews0, sviews2, plan) + for (d, r) in zip(dviews0, refs) + @test refequal(d, r) + end + end + + @testset "first tensor differs: prefix copy is a no-op, nothing before it to reuse" begin + big = rand(20, 5, 6) + sviews6 = copy(sviews0) + sviews6[1] = view(StridedView(big), 9:12, :, :) # same shape as srcs[1] (4,5,6), different offset/strides + changed = Strided._rebuild_descriptors(dviews0, sviews6, plan) + @test changed !== plan + @test changed.descs[1] != plan.descs[1] + @test changed.descs[[2, 3, 4]] == plan.descs[[2, 3, 4]] + @test changed.cpublocks[[2, 3, 4]] == plan.cpublocks[[2, 3, 4]] + ref6 = Strided.plan_batched_permutedims(dviews0, sviews6, perm) + @test changed.descs == ref6.descs && changed.cpublocks == ref6.cpublocks + @test plan.descs == Strided.plan_batched_permutedims(dsts, srcs, perm).descs + end + + @testset "last tensor differs: the whole prefix is copied, nothing after it to write" begin + big = rand(20, 6, 6) + sviews7 = copy(sviews0) + sviews7[4] = view(StridedView(big), 9:14, :, :) # same shape as srcs[4] (6,6,6), different offset/strides + changed = Strided._rebuild_descriptors(dviews0, sviews7, plan) + @test changed !== plan + @test changed.descs[4] != plan.descs[4] + @test changed.descs[[1, 2, 3]] == plan.descs[[1, 2, 3]] + @test changed.cpublocks[[1, 2, 3]] == plan.cpublocks[[1, 2, 3]] + ref7 = Strided.plan_batched_permutedims(dviews0, sviews7, perm) + @test changed.descs == ref7.descs && changed.cpublocks == ref7.cpublocks + @test plan.descs == Strided.plan_batched_permutedims(dsts, srcs, perm).descs + end + + @testset "one source swapped for a fresh same-shape Array: descriptors equal, same plan" begin + # A different array object with identical shape, strides and offset yields a + # bitwise-identical descriptor, so the plan object is reused; that carries no + # stale-array risk because the executor is handed the new views separately. + srcs3 = copy(srcs) + srcs3[3] = rand(size(srcs[3])...) + dviews3, sviews3, _ = Strided._normalize_and_check(dsts, srcs3, N0) + @test Strided._rebuild_descriptors(dviews3, sviews3, plan) === plan + refs = refbatch(srcs3, perm) + Strided.batched_permutedims!(dsts, srcs3, plan) + for (d, r) in zip(dsts, refs) + @test refequal(d, r) + end + end + + @testset "shape mismatch still throws, before and after the first differing tensor" begin + # mismatch on the very first tensor (nothing has been rebuilt yet) + srcs4 = copy(srcs) + srcs4[1] = rand(4, 5, 7) + dsts4 = copy(dsts) + dsts4[1] = zeros(7, 4, 5) + dviews4, sviews4, _ = Strided._normalize_and_check(dsts4, srcs4, N0) + @test_throws ArgumentError Strided._rebuild_descriptors(dviews4, sviews4, plan) + # mismatch on a later tensor, after an earlier tensor's descriptor already differed + big = rand(20, 3, 7) + sviews5 = copy(sviews0) + sviews5[2] = view(StridedView(big), 9:10, :, :) + sviews5[4] = StridedView(rand(6, 6, 5)) + dviews5 = copy(dviews0) + dviews5[4] = StridedView(zeros(5, 6, 6)) + @test_throws ArgumentError Strided._rebuild_descriptors(dviews5, sviews5, plan) + @test plan.descs == Strided.plan_batched_permutedims(dsts, srcs, perm).descs + end + + @testset "_normview infers a concrete view vector for concretely-typed inputs" begin + v = @inferred Strided._normview(srcs) + @test isconcretetype(eltype(v)) && v == map(StridedView, srcs) + @inferred Strided._normalize_and_check(dsts, srcs, N0) + # `Any[...]` inputs still narrow to the actual common view type, or get rejected + anyv = Strided._normview(Any[rand(3, 4), zeros(0, 4)]) + @test isconcretetype(eltype(anyv)) + @test !isconcretetype(eltype(Strided._normview(Any[rand(Float64, 3, 4), rand(Float32, 3, 4)]))) + e = Strided._normview(Matrix{Float64}[]) + @test isempty(e) && isconcretetype(eltype(e)) + end +end + +# ---------------------------------------------------------------------------- +# Scaled path: dsts[b] = alpha[b] * permutedims(srcs[b], perm) + beta[b] * dsts[b]. +# Unlike the plain copy, this computes, so `isequal` is only a sound oracle on data whose +# every intermediate is exactly representable: small integer-valued entries and small +# integer / half-integer coefficients, for which IEEE arithmetic gives bit-identical results +# whatever the evaluation order or multiply-add fusion. A mismatch here is a real finding, +# never something to relax to `isapprox`. +# ---------------------------------------------------------------------------- + +exactdata(::Type{T}, dims) where {T <: Real} = T.(rand(-8:8, dims)) +exactdata(::Type{T}, dims) where {T <: Complex} = T.(complex.(rand(-8:8, dims), rand(-8:8, dims))) +exactcoeffs(::Type{T}) where {T <: Real} = T[-3, -1, 0, 0.5, 1, 2, 5] +exactcoeffs(::Type{T}) where {T <: Complex} = T[-3, -1, 0, 0.5, 1, 2, 5, 1 - 2im] + +# reference in plain Array arithmetic in T; `beta == 0` must never read the destination +scaledref(a, s, b, d0, perm) = iszero(b) ? a .* ref_base(s, perm) : a .* ref_base(s, perm) .+ b .* d0 + +# alpha/beta cycling through the exact set with a shift, so tensor 1 gets beta == 0, tensor 3 +# gets alpha == 0, and every later tensor gets some nonzero beta +function scaledcoeffs(::Type{T}, B::Int) where {T} + cs = exactcoeffs(T) + n = length(cs) + return [cs[mod1(b, n)] for b in 1:B], [cs[mod1(b + 2, n)] for b in 1:B] +end + +dstshape(s, perm) = ntuple(j -> size(s, perm[j]), length(perm)) + +# destinations: NaN-poisoned wherever beta == 0 (they must come out clean), exact data elsewhere +function scaleddsts(::Type{T}, srcs, perm, beta) where {T} + return [iszero(beta[b]) ? fill(T(NaN), dstshape(srcs[b], perm)) : exactdata(T, dstshape(srcs[b], perm)) + for b in eachindex(srcs)] +end + +function scaledcase!(dsts, srcs, perm, alpha, beta; plan = nothing) + T = eltype(dsts[1]) + d0 = map(copy, dsts) + srcsnap = map(deepcopy, srcs) + out = plan === nothing ? Strided.batched_permutedims!(dsts, srcs, perm; alpha, beta) : + Strided.batched_permutedims!(dsts, srcs, plan; alpha, beta) + @test out === dsts + for b in eachindex(dsts) + @test refequal(dsts[b], scaledref(T(alpha[b]), srcs[b], T(beta[b]), d0[b], perm)) + end + for (s, s0) in zip(srcs, srcsnap) + @test refequal(s, s0) + end + return dsts +end + +@testset "Strided.batched_permutedims! (alpha/beta: exact-arithmetic fixtures)" begin + @testset "mixed batch, T=$T" for T in (Float32, Float64, ComplexF64) + # rank-2 transposes including one tensor above MINTHREADLENGTH elements (the + # multi-threaded path when threads are enabled), then rank-3 transpose, payload and + # identity permutations; more tensors than coefficients so the whole set cycles + for (shapes, perm) in ( + ([(3, 4), (65, 33), (1, 7), (200, 200), (2, 2), (31, 1), (64, 64), (9, 8), (16, 16)], (2, 1)), + ([(5, 6, 7), (33, 2, 31), (1, 1, 1), (2, 40, 3), (17, 17, 17), (1, 64, 2), (3, 3, 3), (8, 1, 8)], (3, 1, 2)), + ([(5, 6, 7), (2, 33, 31), (1, 1, 32), (32, 2, 1)], (1, 3, 2)), + ([(5, 6, 7), (33, 1, 2), (4, 4, 4)], (1, 2, 3)), + ) + perm == (2, 1) && @test maximum(prod, shapes) > Strided.MINTHREADLENGTH + srcs = [exactdata(T, s) for s in shapes] + alpha, beta = scaledcoeffs(T, length(shapes)) + @test iszero(beta[1]) && iszero(alpha[3]) && any(!iszero, beta) + scaledcase!(scaleddsts(T, srcs, perm, beta), srcs, perm, alpha, beta) + end + end + + @testset "omitted alpha/beta is exactly ones/zeros, T=$T" for T in (Float64, ComplexF64) + shapes = [(7, 9), (65, 33), (1, 5), (200, 200)] + perm = (2, 1) + srcs = [exactdata(T, s) for s in shapes] + B = length(srcs) + alpha, beta = scaledcoeffs(T, B) + beta = T[2, -1, 0.5, 1] # all nonzero: the destination is read in every tensor + d0 = [exactdata(T, dstshape(s, perm)) for s in srcs] + runwith(a, b) = Strided.batched_permutedims!(map(copy, d0), srcs, perm; alpha = a, beta = b) + # both omitted: the plain copy, and the same bits as explicit (ones, zeros) + plain = runwith(nothing, nothing) + for (d, s) in zip(plain, srcs) + @test refequal(d, ref_base(s, perm)) + end + @test all(map(refequal, plain, runwith(ones(T, B), zeros(T, B)))) + @test all(map(refequal, runwith(alpha, nothing), runwith(alpha, zeros(T, B)))) + @test all(map(refequal, runwith(nothing, beta), runwith(ones(T, B), beta))) + # integer coefficient vectors convert to T + @test all(map(refequal, runwith([2, -1, 0, 5], [1, 1, 0, 2]), runwith(T[2, -1, 0, 5], T[1, 1, 0, 2]))) + # a NaN-poisoned destination with beta omitted (== 0) comes out clean + poisoned = [fill(T(NaN), dstshape(s, perm)) for s in srcs] + Strided.batched_permutedims!(poisoned, srcs, perm; alpha) + for (d, s, a) in zip(poisoned, srcs, alpha) + @test refequal(d, a .* ref_base(s, perm)) + end + end + + @testset "plan reuse with varying coefficients" begin + T = Float64 + shapes = [(6, 7), (33, 65), (1, 1), (200, 200)] + perm = (2, 1) + srcs = [exactdata(T, s) for s in shapes] + alpha, beta = scaledcoeffs(T, length(srcs)) + dsts = scaleddsts(T, srcs, perm, beta) + plan = Strided.plan_batched_permutedims(dsts, srcs, perm) + scaledcase!(dsts, srcs, perm, alpha, beta; plan) + # the plan carries no coefficient state: other coefficients, and none at all, on reuse + scaledcase!(dsts, srcs, perm, reverse(alpha), T[1, 0, -3, 0.5]; plan) + scaledcase!(dsts, srcs, perm, T[0.5, 0.5, 0.5, 0.5], zeros(T, 4); plan) + Strided.batched_permutedims!(dsts, srcs, plan) + for (d, s) in zip(dsts, srcs) + @test refequal(d, ref_base(s, perm)) + end + end + + @testset "offset / strided / negative-stride views, destination read through the view" begin + T = Float64 + perm = (2, 1) + sph = [exactdata(T, (60, 60)) for _ in 1:4] + dph = [exactdata(T, (60, 60)) for _ in 1:4] + cases = ( + ((3:12, 5:20), (2:17, 7:16)), # offset sub-blocks, dense + ((3:2:21, 5:20), (2:3:47, 7:16)), # strided + ((12:-1:3, 5:20), (17:-1:2, 7:16)), # reversed axis 1 on both sides + ((21:-2:3, 20:-1:5), (2:17, 16:-1:7)), # reversed non-fastest axes too + ) + svs = [view(StridedView(p), sr...) for (p, (_, sr)) in zip(sph, cases)] + dvs = [view(StridedView(p), dr...) for (p, (dr, _)) in zip(dph, cases)] + alpha = T[2, -1, 0.5, 5] + beta = T[0, 1, -3, 0] + expected = map(copy, dph) + for (e, (dr, sr), s, a, b) in zip(expected, cases, sph, alpha, beta) + view(e, dr...) .= scaledref(a, view(s, sr...), b, view(e, dr...), perm) + end + out = Strided.batched_permutedims!(dvs, svs, perm; alpha, beta) + @test out === dvs + for (p, e) in zip(dph, expected) + @test refequal(p, e) # scaled bits inside the view, untouched outside it + end + end + + @testset "validation happens before any write" begin + s = [exactdata(Float64, (3, 4)), exactdata(Float64, (5, 2))] + poison() = [fill(NaN, 4, 3), fill(NaN, 2, 5)] + plan = Strided.plan_batched_permutedims(poison(), s, (2, 1)) + for (alpha, beta) in ((ones(3), nothing), (nothing, zeros(1)), (ones(2), Float64[]), ([1.0], [0.0])) + d = poison() + @test_throws DimensionMismatch Strided.batched_permutedims!(d, s, (2, 1); alpha, beta) + @test_throws DimensionMismatch Strided.batched_permutedims!(d, s, plan; alpha, beta) + @test all(x -> all(isnan, x), d) + end + # a complex coefficient with nonzero imaginary part on a real batch is inexact + d = poison() + @test_throws InexactError Strided.batched_permutedims!(d, s, (2, 1); alpha = [1 + 2im, 1]) + @test_throws InexactError Strided.batched_permutedims!(d, s, plan; beta = [0, 0.5im]) + @test all(x -> all(isnan, x), d) + # ... but a complex coefficient with zero imaginary part is fine + scaledcase!(poison(), s, (2, 1), [2 + 0im, -1 + 0im], [0, 0]) + # coefficients need a Number element type; the plain copy does not + st = [rand(NTuple{2, Float64}, 3, 4)] + dt = [Array{NTuple{2, Float64}}(undef, 4, 3)] + @test_throws ArgumentError Strided.batched_permutedims!(dt, st, (2, 1); alpha = [1.0]) + @test_throws ArgumentError Strided.batched_permutedims!(dt, st, (2, 1); beta = [0.0]) + runcase!(dt, st, (2, 1)) + end +end diff --git a/test/gpu.jl b/test/gpu.jl index cf08ee5..63b10ce 100644 --- a/test/gpu.jl +++ b/test/gpu.jl @@ -245,3 +245,1021 @@ end end end end + +# ---------- batched out-of-place permutation: GPU execution strategies ---------- +# +# Correctness of every `BatchedPermuteStrategy` on every available GPU array type, each +# checked bit-exactly (`isequal`, never `isapprox`: this operation moves bits and performs +# no arithmetic) against `permutedims` on the CPU copy of the same data, plus the strategy +# resolution table (`BP_AUTO` -> elementwise baseline; `BP_GROUPTILE` only for the +# transpose family, otherwise a silent, inspectable fallback to `BP_ELEMENTWISE`). +# +# Note on families: the planner never merges axes, so the copy family (`BP_COPY`, reduced +# rank 1) only arises for rank-1 batches; an identity permutation on a rank >= 2 batch is +# the payload family (`BP_PAYLOAD`). Both are exercised below. + +const BP_STRATEGIES = (Strided.BP_AUTO, Strided.BP_ELEMENTWISE, Strided.BP_THREADTILE, Strided.BP_GROUPTILE) + +bp_refequal(a::AbstractArray, b::AbstractArray) = isequal(Array(a), Array(b)) + +# what the resolution table says a request must resolve to, given the family +# Mirrors the resolver's own contract independently (family AND shared-memory-budget +# clauses both restated), so this stays a check on the resolution table rather than a +# restatement of only part of it -- a lower `_BP_GROUPTILE_SHMEM_BUDGET` must correctly +# flip this function's answer too, or a budget-related regression would go unnoticed. The +# budget clause is "some menu geometry's tile fits", not "the largest one does": an element +# type whose tile only fits at a smaller edge still runs the cooperative kernel, on that edge. +function bp_expected_strategy(requested, family, ::Type{T}) where {T} + requested === Strided.BP_AUTO && return Strided.BP_ELEMENTWISE + if requested === Strided.BP_GROUPTILE + ext = Base.get_extension(Strided, :StridedGPUArraysExt) + pad, budget = ext._BP_GROUPTILE_PAD, ext._BP_GROUPTILE_SHMEM_BUDGET + fits = any(g -> sizeof(T) * (g.edge + pad) * g.edge <= budget, ext._BP_GROUPTILE_GEOMETRIES) + (family !== Strided.BP_TRANSPOSE || !fits) && return Strided.BP_ELEMENTWISE + end + return requested +end + +# The tile a `BP_GROUPTILE` plan must carry, restated independently of the extension's own +# selection code: over the plan's per-tensor reduced extents, each menu edge `e` schedules +# `cld(d1, e) * cld(dc, e) * e^2` thread slots per tensor (times the product of the other +# extents) for `d1 * dc * ...` real elements; the edge with the highest ratio wins among the +# edges whose tile fits the element type's local-memory budget, a tie going to the larger +# edge. That edge sits on the source-fastest axis 1 and the destination-fastest axis +# `c = plan.dstseq[1]`, 1 elsewhere. +function bp_expected_grouptile_tile(plan::Strided.BatchedPermutePlan{N0, N, T}) where {N0, N, T} + ext = Base.get_extension(Strided, :StridedGPUArraysExt) + pad, budget = ext._BP_GROUPTILE_PAD, ext._BP_GROUPTILE_SHMEM_BUDGET + c = plan.dstseq[1] + dims = [d.dims for d in plan.descs] + elements = sum(prod, dims; init = 0) + slots(e) = sum(prod(g == 1 || g == c ? cld(d[g], e) : d[g] for g in 1:N) * e * e for d in dims; init = 0) + best, bestu = 0, -1.0 + for g in ext._BP_GROUPTILE_GEOMETRIES + sizeof(T) * (g.edge + pad) * g.edge <= budget || continue + s = slots(g.edge) + u = s == 0 ? 1.0 : elements / s + u > bestu && ((best, bestu) = (g.edge, u)) + end + return ntuple(g -> (g == 1 || g == c) ? best : 1, N) +end + +function bp_fixture(AT, T, shapes, perm) + srcs_cpu = [rand(T, s...) for s in shapes] + srcs = [AT(s) for s in srcs_cpu] + dsts = [AT(zeros(T, ntuple(j -> size(s, perm[j]), length(perm)))) for s in srcs_cpu] + refs = [permutedims(s, perm) for s in srcs_cpu] + return srcs_cpu, srcs, dsts, refs +end + +# plan with `strategy`, run, check the return value, bit-exactness, unmodified sources, the +# reported family, and that the resolved strategy is exactly what the table prescribes. +function bp_runcase(AT, T, shapes, perm, strategy, family) + srcs_cpu, srcs, dsts, refs = bp_fixture(AT, T, shapes, perm) + plan = Strided.plan_batched_permutedims(dsts, srcs, perm; strategy) + @test plan.family === family + @test plan.strategy === bp_expected_strategy(strategy, family, T) + out = Strided.batched_permutedims!(dsts, srcs, plan) + @test out === dsts + for (d, r) in zip(dsts, refs) + @test bp_refequal(d, r) + end + for (s, s0) in zip(srcs, srcs_cpu) + @test bp_refequal(s, s0) + end + return plan +end + +@testset "batched_permutedims! GPU strategies ($AT)" for AT in ATs + Ext = Base.get_extension(Strided, :StridedGPUArraysExt) + EDGE = Ext._BP_GROUPTILE_GEOMETRIES[1].edge + + @testset "strategy=$strategy" for strategy in BP_STRATEGIES + # transpose family, extents deliberately not multiples of 32 (edge tiles), both a + # 4-byte and a 16-byte element type + for T in (Float32, ComplexF64) + plan = bp_runcase(AT, T, [(67, 51), (33, 32), (1, 64), (31, 2), (32, 32)], (2, 1), + strategy, Strided.BP_TRANSPOSE) + if plan.strategy === Strided.BP_ELEMENTWISE + @test all(==(1), plan.tile) # BP_AUTO must be the unchanged baseline + elseif plan.strategy === Strided.BP_THREADTILE + @test count(!=(1), plan.tile) == 1 # exactly one tiled axis + @test maximum(plan.tile) == Ext._BP_THREADTILE_K[] + elseif plan.strategy === Strided.BP_GROUPTILE + # which menu edge: the utilization rule on this batch's own extents (the + # edges themselves are pinned by hand in the geometry-menu testset) + @test plan.tile == bp_expected_grouptile_tile(plan) + @test count(!=(1), plan.tile) == 2 + end + end + # payload family (rank 3, shared contiguous axis 1): GROUPTILE must fall back + bp_runcase(AT, Float32, [(5, 6, 7), (2, 33, 31), (1, 1, 32), (32, 2, 1)], (1, 3, 2), + strategy, Strided.BP_PAYLOAD) + # identity permutation on a rank-3 batch: also payload family + bp_runcase(AT, Float32, [(5, 6, 7), (33, 1, 2)], (1, 2, 3), strategy, Strided.BP_PAYLOAD) + # copy family (rank 1): GROUPTILE must fall back + bp_runcase(AT, Float32, [(5,), (33,), (1,), (64,), (31,)], (1,), strategy, Strided.BP_COPY) + # rank 3 and rank 4 transposes with edge-adjacent extents 1, 2, 31, 32, 33 mixed in + plan3 = bp_runcase(AT, Float32, [(33, 2, 31), (1, 32, 33), (2, 1, 1), (31, 33, 2)], (3, 1, 2), + strategy, Strided.BP_TRANSPOSE) + plan4 = bp_runcase(AT, Float32, [(32, 1, 33, 2), (31, 2, 1, 32), (1, 33, 2, 31), (2, 2, 2, 2)], + (4, 2, 1, 3), strategy, Strided.BP_TRANSPOSE) + if strategy === Strided.BP_GROUPTILE + for p in (plan3, plan4) + @test p.tile == bp_expected_grouptile_tile(p) + @test p.tile[1] in map(g -> g.edge, Ext._BP_GROUPTILE_GEOMETRIES) + @test p.tile[p.dstseq[1]] == p.tile[1] && count(==(p.tile[1]), p.tile) == 2 + end + end + # one dominant tensor plus many tiny ones (tile ids straddle many prefix boundaries) + shapes = Tuple{Int, Int}[(129, 97)] + append!(shapes, [(3, 2) for _ in 1:20]) + append!(shapes, [(1, 1) for _ in 1:5]) + bp_runcase(AT, Float32, shapes, (2, 1), strategy, Strided.BP_TRANSPOSE) + end + + @testset "plan reuse keeps its resolved strategy" begin + for strategy in (Strided.BP_THREADTILE, Strided.BP_GROUPTILE) + shapes = [(67, 51), (33, 32), (1, 5)] + _, srcs1, dsts1, _ = bp_fixture(AT, Float32, shapes, (2, 1)) + plan = Strided.plan_batched_permutedims(dsts1, srcs1, (2, 1); strategy) + @test plan.strategy === strategy + for _ in 1:2 + srcs_cpu, srcs, dsts, refs = bp_fixture(AT, Float32, shapes, (2, 1)) + Strided.batched_permutedims!(dsts, srcs, plan) + for (d, r) in zip(dsts, refs) + @test bp_refequal(d, r) + end + for (s, s0) in zip(srcs, srcs_cpu) + @test bp_refequal(s, s0) + end + end + @test plan.strategy === strategy + end + end + + @testset "plan reuse re-checks strides, not just addresses/offsets" begin + # Regression: a device-binding cache keyed only on base address + offset would + # wrongly reuse stale, previously-uploaded strides for a second call that shares + # the same parent, same offset, but a genuinely different view (here: every-other + # column) into it -- producing a silently wrong result with no error raised. + for strategy in (Strided.BP_ELEMENTWISE, Strided.BP_GROUPTILE) + Aparent = AT(Float32.(reshape(1:64, 8, 8))) + s1 = view(StridedView(Aparent), 1:4, 1:4) # dense 4x4 block, strides (1,8) + s2 = view(StridedView(Aparent), 1:4, 1:2:8) # same base & offset, strides (1,16) + ref1 = permutedims(Array(Aparent)[1:4, 1:4], (2, 1)) + ref2 = permutedims(Array(Aparent)[1:4, 1:2:8], (2, 1)) + dst = AT(zeros(Float32, 4, 4)) + plan = Strided.plan_batched_permutedims([dst], [s1], (2, 1); strategy) + Strided.batched_permutedims!([dst], [s1], plan) + @test bp_refequal(dst, ref1) + fill!(dst, 0) + Strided.batched_permutedims!([dst], [s2], plan) + @test bp_refequal(dst, ref2) + end + end + + @testset "BP_THREADTILE orientation knob" begin + along0 = Ext._BP_THREADTILE_ALONG[] + try + Ext._BP_THREADTILE_ALONG[] = :srcfast + plan = bp_runcase(AT, Float32, [(67, 51), (3, 70)], (2, 1), Strided.BP_THREADTILE, Strided.BP_TRANSPOSE) + @test plan.tile == (Ext._BP_THREADTILE_K[], 1) + # axis 1's largest extent here (5) is below K, so the tile is clamped to it + plan = bp_runcase(AT, Float32, [(5, 6, 7), (2, 33, 31)], (1, 3, 2), Strided.BP_THREADTILE, Strided.BP_PAYLOAD) + @test plan.tile == (min(Ext._BP_THREADTILE_K[], 5), 1, 1) + finally + Ext._BP_THREADTILE_ALONG[] = along0 + end + plan = bp_runcase(AT, Float32, [(5, 6, 7), (2, 33, 31)], (1, 3, 2), Strided.BP_THREADTILE, Strided.BP_PAYLOAD) + @test plan.tile == (1, Ext._BP_THREADTILE_K[], 1) + end + + @testset "resolution table" begin + backend = GPUArrays.KernelAbstractions.get_backend(AT(zeros(Float32, 1))) + resolve(req, fam, T) = Strided._batched_resolve_strategy(backend, req, fam, T) + for fam in (Strided.BP_COPY, Strided.BP_PAYLOAD, Strided.BP_TRANSPOSE), T in (Float32, ComplexF64) + @test resolve(Strided.BP_AUTO, fam, T) === Strided.BP_ELEMENTWISE + @test resolve(Strided.BP_ELEMENTWISE, fam, T) === Strided.BP_ELEMENTWISE + @test resolve(Strided.BP_THREADTILE, fam, T) === Strided.BP_THREADTILE + @test resolve(Strided.BP_GROUPTILE, fam, T) === + (fam === Strided.BP_TRANSPOSE ? Strided.BP_GROUPTILE : Strided.BP_ELEMENTWISE) + end + # an element type whose 32x33 tile exceeds the local-memory budget stays eligible as + # long as the menu's smallest tile fits it (it then runs on that geometry); only an + # element type for which no menu tile fits falls back + small = minimum(g -> g.edge, Ext._BP_GROUPTILE_GEOMETRIES) + Tbig = NTuple{8, Float64} + Tnone = NTuple{16, Float64} + @test sizeof(Tbig) * (EDGE + Ext._BP_GROUPTILE_PAD) * EDGE > Ext._BP_GROUPTILE_SHMEM_BUDGET + @test sizeof(Tbig) * (small + Ext._BP_GROUPTILE_PAD) * small <= Ext._BP_GROUPTILE_SHMEM_BUDGET + @test sizeof(Tnone) * (small + Ext._BP_GROUPTILE_PAD) * small > Ext._BP_GROUPTILE_SHMEM_BUDGET + @test resolve(Strided.BP_GROUPTILE, Strided.BP_TRANSPOSE, Tbig) === Strided.BP_GROUPTILE + @test resolve(Strided.BP_GROUPTILE, Strided.BP_TRANSPOSE, Tnone) === Strided.BP_ELEMENTWISE + @test resolve(Strided.BP_GROUPTILE, Strided.BP_TRANSPOSE, ComplexF64) === Strided.BP_GROUPTILE + end + + @testset "BP_GROUPTILE tile-shape invariant is enforced at launch" begin + _, srcs, dsts, _ = bp_fixture(AT, Float32, [(67, 51)], (2, 1)) + plan = Strided.plan_batched_permutedims(dsts, srcs, (2, 1); strategy = Strided.BP_GROUPTILE) + @test plan.strategy === Strided.BP_GROUPTILE + # same plan, but with a tile shape the cooperative kernel does not implement + P = typeof(plan) + bad = P(plan.perm, plan.srcorder, plan.groups, plan.dstseq, plan.family, plan.strategy, + ntuple(Returns(1), length(plan.tile)), plan.srclabels, plan.dstlabels, plan.descs, + plan.cpublocks, plan.prefix, plan.totaltiles, plan.totalelements, plan.deviceid, + Base.RefValue{Any}(nothing)) + @test_throws ArgumentError Strided.batched_permutedims!(dsts, srcs, bad) + end +end + +# ---------- batched out-of-place permutation: the device-binding cache ---------- +# +# The GPU executor caches its per-plan device upload (descriptors, prefix sums, base +# addresses) in `plan.devcache` and reuses it only if the device identity, EVERY parent +# array's current device address, and every host-side descriptor still match. These tests +# pin both halves of that contract: a genuine repeat call reuses the very same cached +# binding (and is bit-exact), and anything that moves data to a different address -- a +# swapped source, a swapped destination, or a `resize!` that relocates a buffer behind an +# unchanged array object -- invalidates it, so the result reflects the data now at the +# arrays actually passed in, never a stale upload. The swapped-array cases are built so +# that the planner's own reuse check cannot be what catches them: a fresh array of the same +# shape has identical strides and offset, hence a bitwise-identical descriptor, so the plan +# object itself is reused as-is and only the address re-check stands between the call and +# a silently wrong answer. +@testset "device binding cache: reuse and invalidation ($AT)" for AT in ATs + Ext = Base.get_extension(Strided, :StridedGPUArraysExt) + shapes = [(33, 17), (5, 40), (16, 16), (2, 65)] + perm = (2, 1) + for strategy in (Strided.BP_ELEMENTWISE, Strided.BP_GROUPTILE) + @testset "strategy=$strategy" begin + srcs_cpu, srcs, dsts, refs = bp_fixture(AT, Float32, shapes, perm) + plan = Strided.plan_batched_permutedims(dsts, srcs, perm; strategy) + @test plan.devcache[] === nothing + Strided.batched_permutedims!(dsts, srcs, plan) + b1 = plan.devcache[] + @test b1 isa Ext._BatchedGPUBinding + for (d, r) in zip(dsts, refs) + @test bp_refequal(d, r) + end + + # (a) the same arrays again: the same cached binding object, bit-exact results + for d in dsts + fill!(d, NaN32) + end + Strided.batched_permutedims!(dsts, srcs, plan) + @test plan.devcache[] === b1 + for (d, r) in zip(dsts, refs) + @test bp_refequal(d, r) + end + # the planner's reuse check indeed hands back the same plan object here, so the + # binding cache's own address check is the only thing guarding the cases below + dviews, sviews, _ = Strided._normalize_and_check(dsts, srcs, 2) + @test Strided._rebuild_descriptors(dviews, sviews, plan) === plan + + # (b) one source replaced by a fresh same-shape array, the old one overwritten + # with a sentinel: the result must show the new array's data + oldsrc = srcs[2] + newsrc_cpu = rand(Float32, shapes[2]...) + srcs[2] = AT(newsrc_cpu) + fill!(oldsrc, -1.0f0) + for d in dsts + fill!(d, NaN32) + end + dviews, sviews, _ = Strided._normalize_and_check(dsts, srcs, 2) + @test Strided._rebuild_descriptors(dviews, sviews, plan) === plan # same descriptors + Strided.batched_permutedims!(dsts, srcs, plan) + @test plan.devcache[] !== b1 + @test bp_refequal(dsts[2], permutedims(newsrc_cpu, perm)) + @test !any(==(-1.0f0), Array(dsts[2])) + for i in (1, 3, 4) + @test bp_refequal(dsts[i], refs[i]) + end + @test all(==(-1.0f0), Array(oldsrc)) # sources are never written + + # (b') the same for a destination: the new one gets written, the old one (now + # holding a sentinel) is left alone + b2 = plan.devcache[] + olddst = dsts[3] + dsts[3] = AT(fill(NaN32, size(olddst))) + fill!(olddst, -2.0f0) + Strided.batched_permutedims!(dsts, srcs, plan) + @test plan.devcache[] !== b2 + @test bp_refequal(dsts[3], refs[3]) + @test all(==(-2.0f0), Array(olddst)) + + # (c) a repeat call after the swaps is a cache hit again + b3 = plan.devcache[] + Strided.batched_permutedims!(dsts, srcs, plan) + @test plan.devcache[] === b3 + end + end + + if AT === CuArray + @testset "resize! relocating a CuVector behind the same array object" begin + # `resize!` on a CuVector allocates a new buffer and frees the old one, even when + # growing and then shrinking back to the original length -- the array object, + # its length and its shape are all unchanged afterwards, but its data lives at + # a new address (observed to relocate on every trial in this environment; the + # address is still checked below rather than assumed). An identity- or + # length-based cache shortcut would keep using the freed old buffer here. + a_cpu = rand(Float32, 1000) + a = AT(a_cpu) + d = AT(zeros(Float32, 1000)) + plan = Strided.plan_batched_permutedims([d], [a], (1,)) + Strided.batched_permutedims!([d], [a], plan) + @test bp_refequal(d, a_cpu) + b1 = plan.devcache[] + addr0 = UInt(pointer(a)) + resize!(a, 2000) + resize!(a, 1000) + addr1 = UInt(pointer(a)) + @test length(a) == 1000 + new_cpu = rand(Float32, 1000) + copyto!(a, new_cpu) + fill!(d, NaN32) + Strided.batched_permutedims!([d], [a], plan) + @test bp_refequal(d, new_cpu) + # only meaningful as a cache-invalidation check if the buffer actually moved; + # the correctness assertion above holds either way + if addr1 != addr0 + @test plan.devcache[] !== b1 + end + end + end +end + +# ---------- batched out-of-place permutation: addressing under views and deep lookups ---------- +# +# The cooperative `BP_GROUPTILE` kernel derives every address from per-tensor descriptors +# (offset, strides along the two active axes, tile extents) that travel from its load phase +# to its store phase through workgroup-local memory, and finds its tensor by a binary search +# over the batch's tile-id prefix sums. The checks below exercise exactly those paths, for +# `BP_GROUPTILE` and (as the reference implementation of the same addressing) the +# elementwise baseline: a batch deep enough that the search has real depth, extents that sit +# on, just inside and just outside a tile edge along both active axes, and source *and* +# destination views with offsets, non-unit strides and negative strides -- each with a guard +# band, i.e. the whole destination parent compared against a sentinel-filled reference so +# that a write landing anywhere outside the view is caught, not just a wrong value inside it. + +const BP_EDGE_EXTENTS = (1, 2, 31, 32, 33, 63, 64, 65) + +# One batch of views: `cases` is a vector of `(dranges, sranges)` pairs, each carving a +# destination view and a source view out of its own freshly allocated parent pair. Returns +# the plan so the caller can assert what strategy the request actually resolved to. +function bp_viewcase(AT, T, perm, dparentsize, sparentsize, cases, strategy) + dphs = [fill(T(NaN), dparentsize...) for _ in cases] + sphs = [rand(T, sparentsize...) for _ in cases] + dps = [AT(p) for p in dphs] + sps = [AT(p) for p in sphs] + dvs = [view(StridedView(dp), dr...) for (dp, (dr, _)) in zip(dps, cases)] + svs = [view(StridedView(sp), sr...) for (sp, (_, sr)) in zip(sps, cases)] + for (dv, sv) in zip(dvs, svs) + @test size(dv) == ntuple(j -> size(sv, perm[j]), length(perm)) + end + plan = Strided.plan_batched_permutedims(dvs, svs, perm; strategy) + out = Strided.batched_permutedims!(dvs, svs, plan) + @test out === dvs + for (dp, sp, dph, sph, (dr, sr)) in zip(dps, sps, dphs, sphs, cases) + expected = copy(dph) + view(expected, dr...) .= permutedims(view(sph, sr...), perm) + @test isequal(Array(dp), expected) # right bits inside the view, sentinel outside + @test isequal(Array(sp), sph) # source parent untouched + end + return plan +end + +@testset "batched_permutedims! GPU addressing: views, edge extents, deep lookup ($AT)" for AT in ATs + Ext = Base.get_extension(Strided, :StridedGPUArraysExt) + EDGE = Ext._BP_GROUPTILE_GEOMETRIES[1].edge + + @testset "deep tile->tensor lookup, strategy=$strategy" for strategy in BP_STRATEGIES + # >= 300 tensors of mixed extents: the prefix array has > 2^8 entries, so the + # device-side binary search runs to depth 9, and consecutive tile ids constantly + # cross from one tensor to the next (most tensors here are a single, partial tile). + cycle = [(5, 7), (33, 31), (1, 64), (64, 1), (2, 2), (31, 33), (65, 3), (3, 65), + (32, 32), (17, 9), (1, 1), (2, 33), (33, 2), (64, 64)] + shapes = [cycle[mod1(i, length(cycle))] for i in 1:320] + plan = bp_runcase(AT, Float32, shapes, (2, 1), strategy, Strided.BP_TRANSPOSE) + @test length(plan.prefix) == 321 + @test plan.totaltiles == sum(prod(cld.(s, plan.tile)) for s in shapes) + end + + @testset "edge extents on both active axes, strategy=$strategy" for strategy in + (Strided.BP_ELEMENTWISE, Strided.BP_GROUPTILE) + # every pair of extents from {1, 2, 31, 32, 33, 63, 64, 65} on the source-fastest and + # destination-fastest axes, i.e. tiles that are exactly full, one element short, one + # element over, and multi-tile in each direction, for rank 2 and for rank 3 with the + # destination-fastest axis being source axis 3 (perm (3,1,2)) or 2 (perm (2,3,1)) + grid = [(i, j) for i in BP_EDGE_EXTENTS for j in BP_EDGE_EXTENTS] + plan2 = bp_runcase(AT, Float32, grid, (2, 1), strategy, Strided.BP_TRANSPOSE) + m = Iterators.cycle((1, 2, 3)) + grid3a = [(i, mm, j) for ((i, j), mm) in zip(grid, m)] + plan3a = bp_runcase(AT, Float32, grid3a, (3, 1, 2), strategy, Strided.BP_TRANSPOSE) + grid3b = [(i, j, mm) for ((i, j), mm) in zip(grid, m)] + plan3b = bp_runcase(AT, Float32, grid3b, (2, 3, 1), strategy, Strided.BP_TRANSPOSE) + if strategy === Strided.BP_GROUPTILE + @test plan2.tile == bp_expected_grouptile_tile(plan2) + @test plan3a.tile == bp_expected_grouptile_tile(plan3a) && plan3a.dstseq[1] == 3 + @test plan3b.tile == bp_expected_grouptile_tile(plan3b) && plan3b.dstseq[1] == 2 + end + end + + @testset "offset / strided / negative-stride views with guard band, strategy=$strategy" for + strategy in (Strided.BP_ELEMENTWISE, Strided.BP_GROUPTILE) + T = Float32 + P = (210, 210) + # Rank 2, perm (2,1): destination views are 40 x 67, source views 67 x 40. In every + # case here the view's fastest axis is still its axis 1 (positive or negative unit + # or non-unit stride), so the batch is the transpose family and a `BP_GROUPTILE` + # request must actually run the cooperative kernel, not fall back. + for (dr, sr) in ( + ((3:42, 5:71), (2:68, 7:46)), # offset sub-blocks, dense + ((3:2:81, 5:71), (2:68, 7:46)), # stride-2 destination axis 1 + ((3:42, 5:71), (2:3:200, 7:46)), # stride-3 source axis 1 + ((42:-1:3, 5:71), (2:68, 7:46)), # reversed destination axis 1 + ((3:42, 5:71), (68:-1:2, 7:46)), # reversed source axis 1 + ((81:-2:3, 5:71), (200:-3:2, 7:46)), # both reversed and strided + ((3:42, 7:2:139), (2:68, 7:46)), # stride-2 destination axis 2 + ) + plan = bp_viewcase(AT, T, (2, 1), P, P, [(dr, sr)], strategy) + @test plan.family === Strided.BP_TRANSPOSE + @test plan.strategy === strategy + end + # Several views in one batch, each from its own parent pair, mixing the variants + # above so that per-tensor offsets and strides genuinely differ across the batch. + plan = bp_viewcase(AT, T, (2, 1), P, P, [ + ((3:42, 5:71), (2:68, 7:46)), + ((81:-2:3, 5:71), (2:3:200, 7:46)), + ((42:-1:3, 5:71), (68:-1:2, 7:46)), + ((100:139, 7:73), (50:116, 150:189)), + ], strategy) + @test plan.family === Strided.BP_TRANSPOSE + @test plan.strategy === strategy + # A 16-byte element type through the same path (one case, to bound compile time). + plan = bp_viewcase(AT, ComplexF64, (2, 1), P, P, [((81:-2:3, 5:71), (200:-3:2, 7:46))], strategy) + @test plan.strategy === strategy + # Rank 3, perm (3,1,2): destination 34 x 31 x 3, source 31 x 3 x 34, with an offset + # source and a reversed destination axis 1. + P3 = (40, 40, 40) + plan = bp_viewcase(AT, T, (3, 1, 2), P3, P3, + [((37:-1:4, 5:35, 2:4), (2:32, 5:7, 3:36)), ((4:37, 5:35, 2:4), (32:-1:2, 5:7, 3:36))], + strategy) + @test plan.family === Strided.BP_TRANSPOSE + @test plan.strategy === strategy + # A negative stride on a view's *non*-fastest axis flips the planner's heuristic + # axis order (it sorts by signed stride), which reclassifies the batch as the + # payload family and makes a `BP_GROUPTILE` request fall back to the elementwise + # kernel. The results must still be bit-exact either way; which strategy actually + # ran is deliberately not asserted here, so a later change to that heuristic does + # not turn this into a false failure. + for (dr, sr) in ( + ((3:42, 5:71), (2:68, 46:-1:7)), # reversed source axis 2 + ((3:42, 71:-1:5), (2:68, 7:46)), # reversed destination axis 2 + ) + plan = bp_viewcase(AT, T, (2, 1), P, P, [(dr, sr)], strategy) + @test plan.strategy in (Strided.BP_ELEMENTWISE, Strided.BP_GROUPTILE) + end + end +end + +# ---------- batched out-of-place permutation: BP_GROUPTILE tile-geometry menu ---------- +# +# `BP_GROUPTILE` picks its tile edge from a small fixed menu (`_BP_GROUPTILE_GEOMETRIES`), +# so that a batch of tensors much smaller than a 32x32 tile can run on a 16x16 one instead of +# idling three quarters of every workgroup. The edge is chosen per batch by the exact +# tile-slot utilization rule over every tensor's own extents (so a batch mixing a few large +# tensors with many tiny ones is weighed tensor by tensor), restricted to the edges whose +# tile fits the element type's local-memory budget. The checks below pin down: the menu's +# shape (bounded, ordered, divisibility), the utilization rule on hand-computable synthetic +# extents, the resolver's "some edge fits" eligibility, what the tile-shape hook actually +# selects (asserted on `plan.tile` against hand-computed slot counts AND against +# `bp_expected_grouptile_tile`) for small-only, mixed, tie and large batches, bit-exact runs +# through both geometries on full and partial tiles (Float32, rank 2 and 3 only, to bound +# the number of compiled kernel variants; one 64-byte element type on the small geometry), +# and the executor's rejection of any tile edge that is not a menu entry or does not fit the +# element type. + +@testset "batched_permutedims! GPU BP_GROUPTILE geometry menu ($AT)" for AT in ATs + Ext = Base.get_extension(Strided, :StridedGPUArraysExt) + menu = Ext._BP_GROUPTILE_GEOMETRIES + EDGE = Ext._BP_GROUPTILE_GEOMETRIES[1].edge + edges = map(g -> g.edge, menu) + budget = Ext._BP_GROUPTILE_SHMEM_BUDGET + tilebytes = Ext._batched_grouptile_tilebytes + # a 64-byte element type: its 32x33 tile is over the budget, its 16x17 tile is not; + # and a 128-byte one for which no menu tile fits at all + Tbig = NTuple{8, Float64} + Tnone = NTuple{16, Float64} + @test tilebytes(Tbig, 16) <= budget < tilebytes(Tbig, EDGE) + @test tilebytes(Tnone, 16) > budget + bigpoison = ntuple(Returns(NaN), 8) # `zeros`/`T(NaN)` do not exist for tuple types + + @testset "menu invariants" begin + @test 1 <= length(menu) <= 3 + @test menu[1] == (edge = 32, rows = 4) + for (i, geom) in enumerate(menu) + @test geom.edge >= 1 && geom.rows >= 1 + @test geom.edge % geom.rows == 0 # full (li, lj) coverage by EDGE÷ROWS rows per thread + i > 1 && @test geom.edge < menu[i - 1].edge # strictly decreasing: entry 1 is the largest + # a smaller edge never needs more local memory than the largest one + for T in (Float32, Float64, ComplexF64) + @test Ext._batched_grouptile_tilebytes(T, geom.edge) <= + Ext._batched_grouptile_tilebytes(T, EDGE) + end + end + # the checks below assume the shipped menu: 32 first, 16 present + @test EDGE == 32 + @test 16 in edges + end + + @testset "utilization rule on synthetic extents" begin + util = Ext._batched_grouptile_utilization + pick(dims, c, T = Float32) = menu[Ext._batched_grouptile_geometry(dims, c, T)].edge + # one 16x16 tensor: a 32x32 tile is a quarter full, a 16x16 tile is exactly full + @test util([(16, 16)], 2, 32) == 0.25 + @test util([(16, 16)], 2, 16) == 1.0 + # rank 3, active axes 1 and 3: the middle axis multiplies the tile count, not the waste + @test util([(16, 16, 16)], 3, 32) == 0.25 + @test util([(16, 16, 16)], 3, 16) == 1.0 + # 33x33: 4 tiles of 32x32 (1089/4096) vs 9 tiles of 16x16 (1089/2304) + @test util([(33, 33)], 2, 32) == 1089 / 4096 + @test util([(33, 33)], 2, 16) == 1089 / 2304 + # no tiles at all (empty batch, or every tensor empty): defined, and a tie + @test util(NTuple{2, Int}[], 2, 32) == 1.0 + @test util([(0, 16), (16, 0)], 2, 16) == 1.0 + # which geometry wins, by hand + @test pick([(16, 16) for _ in 1:2000], 2) == 16 + @test pick([(16, 16, 16) for _ in 1:2000], 3) == 16 + @test pick([(8, 8), (16, 1), (1, 16), (2, 3)], 2) == 16 + @test pick([(33, 33)], 2) == 16 # finer edge idles fewer slots + @test pick([(64, 64) for _ in 1:10], 2) == 32 # full tiles either way: tie -> largest + @test pick([(4096, 4096)], 2) == 32 + @test pick([(20, 20)], 2) == 32 # 1x1 tiles of 32 vs 2x2 of 16: same slots, tie + @test pick(NTuple{2, Int}[], 2) == 32 + @test pick([(0, 16)], 2) == 32 + # mixed: one full-tile tensor plus many small ones -- the small ones dominate the slots + @test pick(vcat([(64, 64)], [(16, 16) for _ in 1:1000]), 2) == 16 + # the other active axis is `c`, not always 2 + @test pick([(16, 100, 16)], 3) == 16 + @test pick([(16, 16, 100)], 2) == 16 + @test pick([(32, 100, 32)], 3) == 32 + # an element type whose tile only fits at the small edge gets that edge whatever the + # extents (even on full 32x32 tiles, and on the no-tile tie), and one for which no + # menu tile fits is refused outright rather than silently given a geometry + @test pick([(4096, 4096)], 2, Tbig) == 16 + @test pick([(64, 64) for _ in 1:10], 2, Tbig) == 16 + @test pick(NTuple{2, Int}[], 2, Tbig) == 16 + @test_throws ArgumentError pick([(16, 16)], 2, Tnone) + end + + @testset "resolver: eligible as soon as some menu geometry fits the local-memory budget" begin + backend = GPUArrays.KernelAbstractions.get_backend(AT(zeros(Float32, 1))) + resolve(req, fam, T) = Strided._batched_resolve_strategy(backend, req, fam, T) + anyfit(T) = any(g -> tilebytes(T, g.edge) <= budget, menu) + for T in (Float32, Float64, ComplexF32, ComplexF64, Tbig, Tnone) + @test resolve(Strided.BP_GROUPTILE, Strided.BP_TRANSPOSE, T) === + (anyfit(T) ? Strided.BP_GROUPTILE : Strided.BP_ELEMENTWISE) + @test resolve(Strided.BP_GROUPTILE, Strided.BP_TRANSPOSE, T) === bp_expected_strategy(Strided.BP_GROUPTILE, Strided.BP_TRANSPOSE, T) + # the smallest edge alone decides, because tile bytes are monotone in the edge + @test anyfit(T) == (tilebytes(T, minimum(edges)) <= budget) + # never for the other families, whatever fits + @test resolve(Strided.BP_GROUPTILE, Strided.BP_PAYLOAD, T) === Strided.BP_ELEMENTWISE + @test resolve(Strided.BP_GROUPTILE, Strided.BP_COPY, T) === Strided.BP_ELEMENTWISE + end + # the two element types the byte math above singled out + @test resolve(Strided.BP_GROUPTILE, Strided.BP_TRANSPOSE, Tbig) === Strided.BP_GROUPTILE + @test resolve(Strided.BP_GROUPTILE, Strided.BP_TRANSPOSE, Tnone) === Strided.BP_ELEMENTWISE + # End to end for the eligible one: on extents that tie between the two edges (every + # extent a multiple of 32 or in the upper half of its last 32-block) a 4-byte element + # type gets the larger edge, while this element type gets the 16 tile because the 32 + # one is over budget -- and the cooperative kernel runs bit-exactly on 64-byte + # elements through it (NaN-poisoned destinations). + shapes = [(64, 64), (20, 20), (50, 60), (96, 17)] + plan = bp_runcase(AT, Float32, shapes, (2, 1), Strided.BP_GROUPTILE, Strided.BP_TRANSPOSE) + @test plan.tile == (32, 32) == bp_expected_grouptile_tile(plan) + srcs_cpu = [rand(Tbig, s...) for s in shapes] + srcs = [AT(s) for s in srcs_cpu] + dsts = [AT(fill(bigpoison, size(s, 2), size(s, 1))) for s in srcs_cpu] + plan = Strided.plan_batched_permutedims(dsts, srcs, (2, 1); strategy = Strided.BP_GROUPTILE) + @test plan.strategy === Strided.BP_GROUPTILE + @test plan.family === Strided.BP_TRANSPOSE + @test plan.tile == (16, 16) == bp_expected_grouptile_tile(plan) + @test plan.totaltiles == 16 + 4 + 16 + 12 + out = Strided.batched_permutedims!(dsts, srcs, plan) + @test out === dsts + for (d, s) in zip(dsts, srcs_cpu) + @test isequal(Array(d), permutedims(s, (2, 1))) + end + for (s, s0) in zip(srcs, srcs_cpu) + @test isequal(Array(s), s0) + end + # the ineligible one falls back and still runs correctly + srcs_cpu = [rand(Tnone, 9, 5)] + srcs = [AT(s) for s in srcs_cpu] + dsts = [AT(fill(ntuple(Returns(NaN), 16), 5, 9))] + plan = Strided.plan_batched_permutedims(dsts, srcs, (2, 1); strategy = Strided.BP_GROUPTILE) + @test plan.strategy === Strided.BP_ELEMENTWISE + Strided.batched_permutedims!(dsts, srcs, plan) + @test isequal(Array(dsts[1]), permutedims(srcs_cpu[1], (2, 1))) + end + + @testset "plan.tile follows the exact per-tensor utilization rule" begin + S = Strided.BP_GROUPTILE + # Every case is asserted twice: against the rule restated in `bp_expected_grouptile_tile` + # on the plan's own descriptors, and against a hand-computed edge. Hand arithmetic + # (rank 2): a tensor `d1 x dc` schedules `cld(d1, e) * cld(dc, e) * e^2` slots on edge + # `e`; the batch's edge is the one with fewer slots in total, 32 on a tie. + function tileof(shapes, perm) + plan = bp_runcase(AT, Float32, shapes, perm, S, Strided.BP_TRANSPOSE) + @test plan.tile == bp_expected_grouptile_tile(plan) + return plan + end + # all small (within 16 on both active axes), rank 2 and rank 3 with the + # destination-fastest axis being source axis 3 (perm (3,1,2)) or 2 (perm (2,3,1)) + plan = tileof([(16, 16), (3, 16), (16, 1), (1, 1), (15, 9)], (2, 1)) + @test plan.tile == (16, 16) + @test plan.totaltiles == 5 + plan = tileof([(16, 5, 16), (2, 3, 4), (16, 16, 16), (1, 7, 15)], (3, 1, 2)) + @test plan.tile == (16, 1, 16) && plan.dstseq[1] == 3 + plan = tileof([(16, 16, 5), (4, 2, 3), (16, 16, 16), (15, 1, 7)], (2, 3, 1)) + @test plan.tile == (16, 16, 1) && plan.dstseq[1] == 2 + # mixed: one 40x40 (32: 2x2 tiles = 4096 slots; 16: 3x3 = 2304) plus fifty 16x16 + # (32: 51200; 16: 12800) -- 55296 vs 15104 slots for 14400 elements: 16. A rule + # working from the batch's per-axis maxima (40, 40) alone could not see this. + shapes = Tuple{Int, Int}[(40, 40)] + append!(shapes, [(16, 16) for _ in 1:50]) + plan = tileof(shapes, (2, 1)) + @test plan.tile == (16, 16) + @test plan.totaltiles == 9 + 50 + # the same mix at rank 3, either placement of the destination-fastest axis + plan = tileof(vcat([(40, 3, 40)], [(16, 2, 16) for _ in 1:50]), (3, 1, 2)) + @test plan.tile == (16, 1, 16) && plan.dstseq[1] == 3 + plan = tileof(vcat([(40, 40, 3)], [(16, 16, 2) for _ in 1:50]), (2, 3, 1)) + @test plan.tile == (16, 16, 1) && plan.dstseq[1] == 2 + # varied moderate-to-large: 64x64 (4096 / 4096), 65x33 (6144 / 3840), 16x16 + # (1024 / 256), 512x3 (16384 / 8192) -- 27648 vs 16384: 16 + plan = tileof([(64, 64), (65, 33), (16, 16), (512, 3)], (2, 1)) + @test plan.tile == (16, 16) + # a single extent of 17 on an active axis: 1 tile of 32 (1024) vs 2x1 of 16 (512), + # on top of small tensors -- 16 (a per-axis-maxima rule would have kept 32 here) + plan = tileof([(17, 16), (16, 16), (1, 1)], (2, 1)) + @test plan.tile == (16, 16) + plan = tileof([(16, 17), (16, 16)], (2, 1)) + @test plan.tile == (16, 16) + # tie: every active extent is a multiple of 32 or lands in the upper half of its last + # 32-block (17, 20, 50, 60, 96: `cld(d, 32) * 32 == cld(d, 16) * 16`), so both edges + # idle exactly the same slots -- the tie goes to the larger edge + plan = tileof([(20, 20), (50, 60), (64, 64), (96, 17)], (2, 1)) + @test plan.tile == (32, 32) + # ... and a single lower-half extent anywhere breaks that tie toward 16: ten 64x64 + # (40960 slots either way) plus one 33x33 (4096 vs 2304) + plan = tileof(vcat([(64, 64) for _ in 1:10], [(33, 33)]), (2, 1)) + @test plan.tile == (16, 16) + # uniformly large, full tiles only: tie -> 32 + plan = tileof([(64, 64), (128, 96), (256, 32)], (2, 1)) + @test plan.tile == (32, 32) + plan = tileof([(64, 5, 96), (128, 2, 32)], (3, 1, 2)) + @test plan.tile == (32, 1, 32) && plan.dstseq[1] == 3 + # empty tensors contribute neither tiles nor elements, so they never sway the choice + plan = tileof([(0, 16), (16, 16), (16, 0), (40, 40)], (2, 1)) + @test plan.tile == (16, 16) + plan = tileof([(0, 5), (64, 64), (7, 0)], (2, 1)) + @test plan.tile == (32, 32) + # all empty: no tiles at all, the tie-break default + plan = tileof([(0, 16), (16, 0)], (2, 1)) + @test plan.tile == (32, 32) + @test plan.totaltiles == 0 + # the non-active axis multiplies both edges' slot counts alike: no effect + plan = tileof([(16, 300, 16), (2, 1, 4)], (3, 1, 2)) + @test plan.tile == (16, 1, 16) + plan = tileof([(64, 300, 64)], (3, 1, 2)) + @test plan.tile == (32, 1, 32) + end + + @testset "bit-exact on a mixed batch through the 16x16 geometry (Float32, rank 2 and 3)" begin + S = Strided.BP_GROUPTILE + # Full 16x16 tiles, tiles partial on one or both axes, multi-tile tensors and + # single-element ones, in one batch; the distinct shapes alone total 22528 slots on + # 32 vs 13568 on 16, and the appended 16x16 tensors only widen that -- so this runs + # on the 16 edge with both full and partial tiles. + base = [(40, 40), (33, 17), (16, 16), (1, 16), (31, 2), (48, 20), (17, 33), (65, 1), (1, 1), (32, 32), (47, 49)] + shapes = vcat(base, [(16, 16) for _ in 1:30]) + plan2 = bp_runcase(AT, Float32, shapes, (2, 1), S, Strided.BP_TRANSPOSE) + @test plan2.tile == (16, 16) == bp_expected_grouptile_tile(plan2) + @test plan2.totaltiles == sum(prod(cld.(s, 16)) for s in shapes) + m = Iterators.cycle((1, 2, 3)) + shapes3a = [(i, mm, j) for ((i, j), mm) in zip(shapes, m)] + plan3a = bp_runcase(AT, Float32, shapes3a, (3, 1, 2), S, Strided.BP_TRANSPOSE) + @test plan3a.tile == (16, 1, 16) == bp_expected_grouptile_tile(plan3a) && plan3a.dstseq[1] == 3 + shapes3b = [(i, j, mm) for ((i, j), mm) in zip(shapes, m)] + plan3b = bp_runcase(AT, Float32, shapes3b, (2, 3, 1), S, Strided.BP_TRANSPOSE) + @test plan3b.tile == (16, 16, 1) == bp_expected_grouptile_tile(plan3b) && plan3b.dstseq[1] == 2 + # views with a guard band: one large offset block next to reversed/strided/tiny + # views, each carved from its own parent, so per-tensor offsets and strides differ + P = (90, 90) + plan = bp_viewcase(AT, Float32, (2, 1), P, P, [ + ((3:42, 5:44), (2:41, 7:46)), # 40 x 40 dense offset block + ((81:-2:3, 5:21), (2:3:50, 7:46)), # 40 x 17: reversed dst, stride-3 src + ((3:18, 5:16), (13:-1:2, 7:22)), # 16 x 12: reversed source axis 1 + ((60:75, 60:75), (60:75, 60:75)), # 16 x 16, one full tile + ((5:5, 1:33), (1:33, 5:5)), # 1 x 33 + ], S) + @test plan.strategy === S + @test plan.tile == (16, 16) == bp_expected_grouptile_tile(plan) + P3 = (50, 50, 50) + plan = bp_viewcase(AT, Float32, (3, 1, 2), P3, P3, [ + ((44:-1:5, 5:37, 2:4), (2:34, 5:7, 3:42)), # 40 x 33 x 3, reversed dst axis 1 + ((5:20, 5:19, 2:4), (16:-1:2, 5:7, 3:18)), # 16 x 15 x 3, reversed src axis 1 + ((1:1, 1:1, 3:19), (1:1, 3:19, 4:4)), # 1 x 1 x 17 + ], S) + @test plan.strategy === S + @test plan.tile == (16, 1, 16) == bp_expected_grouptile_tile(plan) + end + + @testset "bit-exact on partial tiles through the 32x32 geometry (Float32, rank 2 and 3)" begin + S = Strided.BP_GROUPTILE + # The 32 edge is only ever chosen on a tie, i.e. when every tensor's two active + # extents are multiples of 32 or land in the upper half of a 32-block (a remainder + # of 17..31) -- those are exactly the partial tiles the 32x32 kernel meets through + # the planner, so its edge guards are exercised on them here, on every pair of such + # extents. + exts = (17, 20, 31, 32, 50, 63, 64, 96) + grid = [(i, j) for i in exts for j in exts] + plan2 = bp_runcase(AT, Float32, grid, (2, 1), S, Strided.BP_TRANSPOSE) + @test plan2.tile == (32, 32) == bp_expected_grouptile_tile(plan2) + @test plan2.totaltiles == sum(prod(cld.(s, 32)) for s in grid) + m = Iterators.cycle((1, 2, 3)) + grid3a = [(i, mm, j) for ((i, j), mm) in zip(grid, m)] + plan3a = bp_runcase(AT, Float32, grid3a, (3, 1, 2), S, Strided.BP_TRANSPOSE) + @test plan3a.tile == (32, 1, 32) == bp_expected_grouptile_tile(plan3a) && plan3a.dstseq[1] == 3 + grid3b = [(i, j, mm) for ((i, j), mm) in zip(grid, m)] + plan3b = bp_runcase(AT, Float32, grid3b, (2, 3, 1), S, Strided.BP_TRANSPOSE) + @test plan3b.tile == (32, 32, 1) == bp_expected_grouptile_tile(plan3b) && plan3b.dstseq[1] == 2 + # views with a guard band: 50 x 60 destinations (partial 32-tiles on both axes), + # offset, strided and reversed + P = (200, 200) + for (dr, sr) in ( + ((3:52, 5:64), (2:61, 7:56)), # offset, dense + ((101:-2:3, 5:64), (2:3:179, 7:56)), # reversed stride-2 dst, stride-3 src + ((52:-1:3, 5:64), (61:-1:2, 7:56)), # both axis-1 reversed + ) + plan = bp_viewcase(AT, Float32, (2, 1), P, P, [(dr, sr)], S) + @test plan.strategy === S + @test plan.tile == (32, 32) == bp_expected_grouptile_tile(plan) + end + end + + @testset "bit-exact through the 16x16 geometry (Float32, rank 2 and 3)" begin + S = Strided.BP_GROUPTILE + exts = (1, 2, 15, 16) + # every extent pair around the small edge on both active axes + grid = [(i, j) for i in exts for j in exts] + plan2 = bp_runcase(AT, Float32, grid, (2, 1), S, Strided.BP_TRANSPOSE) + @test plan2.tile == (16, 16) + m = Iterators.cycle((1, 2, 3)) + grid3a = [(i, mm, j) for ((i, j), mm) in zip(grid, m)] + plan3a = bp_runcase(AT, Float32, grid3a, (3, 1, 2), S, Strided.BP_TRANSPOSE) + @test plan3a.tile == (16, 1, 16) + grid3b = [(i, j, mm) for ((i, j), mm) in zip(grid, m)] + plan3b = bp_runcase(AT, Float32, grid3b, (2, 3, 1), S, Strided.BP_TRANSPOSE) + @test plan3b.tile == (16, 16, 1) + # deep tile->tensor lookup with the small geometry: 320 single-tile tensors + cycle = [(5, 7), (16, 16), (1, 16), (16, 1), (2, 2), (15, 16), (16, 3), (3, 15), (9, 9), (1, 1)] + shapes = [cycle[mod1(i, length(cycle))] for i in 1:320] + plan = bp_runcase(AT, Float32, shapes, (2, 1), S, Strided.BP_TRANSPOSE) + @test plan.tile == (16, 16) + @test plan.totaltiles == 320 && length(plan.prefix) == 321 + # offset / strided / negative-stride views with a guard band, all within 16 extents + P = (60, 60) + for (dr, sr) in ( + ((3:18, 5:16), (2:13, 7:22)), # offset sub-blocks, dense: dst 16x12, src 12x16 + ((3:2:33, 5:16), (2:13, 7:22)), # stride-2 destination axis 1 + ((3:18, 5:16), (2:3:35, 7:22)), # stride-3 source axis 1 + ((18:-1:3, 5:16), (13:-1:2, 7:22)), # both axis-1 reversed + ((33:-2:3, 5:16), (35:-3:2, 7:22)), # reversed and strided + ) + plan = bp_viewcase(AT, Float32, (2, 1), P, P, [(dr, sr)], S) + @test plan.strategy === S + @test plan.tile == (16, 16) + end + P3 = (30, 30, 30) + plan = bp_viewcase(AT, Float32, (3, 1, 2), P3, P3, + [((20:-1:5, 5:19, 2:4), (2:16, 5:7, 3:18)), ((5:20, 5:19, 2:4), (16:-1:2, 5:7, 3:18))], S) + @test plan.strategy === S + @test plan.tile == (16, 1, 16) + # plan reuse keeps the small geometry across calls with fresh arrays + shapes = [(16, 16), (3, 16), (15, 9)] + _, srcs1, dsts1, _ = bp_fixture(AT, Float32, shapes, (2, 1)) + plan = Strided.plan_batched_permutedims(dsts1, srcs1, (2, 1); strategy = S) + @test plan.tile == (16, 16) + for _ in 1:2 + srcs_cpu, srcs, dsts, refs = bp_fixture(AT, Float32, shapes, (2, 1)) + Strided.batched_permutedims!(dsts, srcs, plan) + for (d, r) in zip(dsts, refs) + @test bp_refequal(d, r) + end + end + end + + @testset "executor rejects a tile edge that is not a menu entry or does not fit the element type" begin + _, srcs, dsts, _ = bp_fixture(AT, Float32, [(16, 16), (3, 16)], (2, 1)) + plan = Strided.plan_batched_permutedims(dsts, srcs, (2, 1); strategy = Strided.BP_GROUPTILE) + @test plan.strategy === Strided.BP_GROUPTILE + P = typeof(plan) + rebuild(tile) = P(plan.perm, plan.srcorder, plan.groups, plan.dstseq, plan.family, plan.strategy, + tile, plan.srclabels, plan.dstlabels, plan.descs, plan.cpublocks, plan.prefix, + plan.totaltiles, plan.totalelements, plan.deviceid, Base.RefValue{Any}(nothing)) + for edge in (8, 24, 64, 1) + edge in edges && continue + @test_throws ArgumentError Strided.batched_permutedims!(dsts, srcs, rebuild((edge, edge))) + end + # a menu edge on axis 1 but a different one on axis c is not a valid geometry either + @test_throws ArgumentError Strided.batched_permutedims!(dsts, srcs, rebuild((16, 32))) + @test_throws ArgumentError Strided.batched_permutedims!(dsts, srcs, rebuild((32, 16))) + # A menu edge whose tile does not fit the plan's element type is refused before any + # launch: the planner gives this 64-byte element type the 16 tile; forcing the 32 + # tile onto the same plan must not reach the kernel (its tile would be over budget). + srcsb = [AT(rand(Tbig, 16, 16)), AT(rand(Tbig, 3, 16))] + dstsb = [AT(fill(bigpoison, 16, 16)), AT(fill(bigpoison, 16, 3))] + planb = Strided.plan_batched_permutedims(dstsb, srcsb, (2, 1); strategy = Strided.BP_GROUPTILE) + @test planb.strategy === Strided.BP_GROUPTILE + @test planb.tile == (16, 16) + Pb = typeof(planb) + badb = Pb(planb.perm, planb.srcorder, planb.groups, planb.dstseq, planb.family, planb.strategy, + (EDGE, EDGE), planb.srclabels, planb.dstlabels, planb.descs, planb.cpublocks, planb.prefix, + planb.totaltiles, planb.totalelements, planb.deviceid, Base.RefValue{Any}(nothing)) + @test_throws ArgumentError Strided.batched_permutedims!(dstsb, srcsb, badb) + @test all(x -> isequal(x, bigpoison), Array(dstsb[1])) # nothing was written + # the untouched plan still runs, bit-exactly + Strided.batched_permutedims!(dstsb, srcsb, planb) + @test isequal(Array(dstsb[1]), permutedims(Array(srcsb[1]), (2, 1))) + @test isequal(Array(dstsb[2]), permutedims(Array(srcsb[2]), (2, 1))) + end +end + +# ---------- batched out-of-place permutation: per-tensor alpha/beta on the GPU ---------- +# +# `dsts[b] = alpha[b] * permutedims(srcs[b], perm) + beta[b] * dsts[b]` on every strategy. +# Unlike the plain copy this computes, and a GPU may fuse the multiply-add where the CPU +# reference does not, so `isequal` is only a sound oracle on exactly-representable data: +# small integer-valued entries and small integer / half-integer coefficients, for which IEEE +# arithmetic is bit-identical whatever the evaluation order or fusion. A mismatch here is a +# real finding, never something to relax to `isapprox`. Kept to a few strategy x type x rank +# combinations on purpose: each scaled (strategy, T, rank, edge) is one more compiled kernel. + +bp_exactdata(::Type{T}, dims) where {T <: Real} = T.(rand(-8:8, dims...)) +bp_exactdata(::Type{T}, dims) where {T <: Complex} = T.(complex.(rand(-8:8, dims...), rand(-8:8, dims...))) +bp_exactcoeffs(::Type{T}) where {T <: Real} = T[-3, -1, 0, 0.5, 1, 2, 5] +bp_exactcoeffs(::Type{T}) where {T <: Complex} = T[-3, -1, 0, 0.5, 1, 2, 5, 1 - 2im] +bp_scaledref(a, s, b, d0, perm) = iszero(b) ? a .* permutedims(s, perm) : a .* permutedims(s, perm) .+ b .* d0 +bp_dstshape(s, perm) = ntuple(j -> size(s, perm[j]), length(perm)) + +# alpha/beta cycling through the exact set with a shift: tensor 1 has beta == 0 (NaN-poisoned +# destination, must come out clean), tensor 3 has alpha == 0, later tensors have beta != 0 +function bp_scaledcoeffs(::Type{T}, B::Int) where {T} + cs = bp_exactcoeffs(T) + n = length(cs) + return [cs[mod1(b, n)] for b in 1:B], [cs[mod1(b + 2, n)] for b in 1:B] +end + +# plan with `strategy`, run with coefficients, check the return value, bit-exactness against +# the CPU reference, and unmodified sources; returns the plan for strategy/tile assertions +function bp_scaledcase(AT, T, shapes, perm, strategy; coeffs = bp_scaledcoeffs(T, length(shapes))) + alpha, beta = coeffs + srcs_cpu = [bp_exactdata(T, s) for s in shapes] + dsts_cpu = [iszero(beta[b]) ? fill(T(NaN), bp_dstshape(s, perm)) : bp_exactdata(T, bp_dstshape(s, perm)) + for (b, s) in enumerate(srcs_cpu)] + srcs = [AT(s) for s in srcs_cpu] + dsts = [AT(d) for d in dsts_cpu] + plan = Strided.plan_batched_permutedims(dsts, srcs, perm; strategy) + out = Strided.batched_permutedims!(dsts, srcs, plan; alpha, beta) + @test out === dsts + for b in eachindex(dsts) + @test bp_refequal(dsts[b], bp_scaledref(T(alpha[b]), srcs_cpu[b], T(beta[b]), dsts_cpu[b], perm)) + end + for (s, s0) in zip(srcs, srcs_cpu) + @test bp_refequal(s, s0) + end + return plan +end + +# `bp_viewcase` with coefficients: each view carved from its own parent pair, the destination +# parent NaN-poisoned where beta == 0 and exact data elsewhere, and the whole parent compared +# (scaled bits inside the view, the original bits outside it) +function bp_scaled_viewcase(AT, T, perm, dparentsize, sparentsize, cases, strategy, alpha, beta) + dphs = [iszero(b) ? fill(T(NaN), dparentsize...) : bp_exactdata(T, dparentsize) for b in beta] + sphs = [bp_exactdata(T, sparentsize) for _ in cases] + dps = [AT(p) for p in dphs] + sps = [AT(p) for p in sphs] + dvs = [view(StridedView(dp), dr...) for (dp, (dr, _)) in zip(dps, cases)] + svs = [view(StridedView(sp), sr...) for (sp, (_, sr)) in zip(sps, cases)] + plan = Strided.plan_batched_permutedims(dvs, svs, perm; strategy) + out = Strided.batched_permutedims!(dvs, svs, plan; alpha, beta) + @test out === dvs + for (dp, sp, dph, sph, (dr, sr), a, b) in zip(dps, sps, dphs, sphs, cases, alpha, beta) + expected = copy(dph) + view(expected, dr...) .= bp_scaledref(T(a), view(sph, sr...), T(b), view(dph, dr...), perm) + @test isequal(Array(dp), expected) + @test isequal(Array(sp), sph) + end + return plan +end + +@testset "batched_permutedims! GPU alpha/beta ($AT)" for AT in ATs + Ext = Base.get_extension(Strided, :StridedGPUArraysExt) + + @testset "every strategy, T=$T" for T in (Float32, ComplexF64) + # transpose family; on this mix BP_GROUPTILE picks the 16 edge + shapes = [(67, 51), (33, 32), (1, 64), (31, 2), (32, 32), (16, 16), (50, 3), (5, 70)] + for strategy in (Strided.BP_ELEMENTWISE, Strided.BP_THREADTILE, Strided.BP_GROUPTILE) + plan = bp_scaledcase(AT, T, shapes, (2, 1), strategy) + @test plan.strategy === strategy + strategy === Strided.BP_GROUPTILE && @test plan.tile == (16, 16) + end + end + + @testset "BP_GROUPTILE 32 edge, rank 3, and the shared elementwise kernel on the payload family" begin + # tie extents (multiples of 32 or upper half of a 32-block) keep the 32 edge; partial tiles + plan = bp_scaledcase(AT, Float32, [(64, 64), (20, 20), (50, 60), (96, 17)], (2, 1), Strided.BP_GROUPTILE) + @test plan.strategy === Strided.BP_GROUPTILE && plan.tile == (32, 32) + plan = bp_scaledcase(AT, Float32, [(33, 2, 31), (1, 32, 33), (2, 1, 1), (31, 33, 2), (16, 3, 16)], (3, 1, 2), Strided.BP_GROUPTILE) + @test plan.strategy === Strided.BP_GROUPTILE && plan.tile[1] == plan.tile[3] != 1 && plan.dstseq[1] == 3 + plan = bp_scaledcase(AT, Float32, [(5, 6, 7), (2, 33, 31), (1, 1, 32), (32, 2, 1)], (1, 3, 2), Strided.BP_THREADTILE) + @test plan.family === Strided.BP_PAYLOAD && plan.strategy === Strided.BP_THREADTILE + # one dominant tensor plus many tiny ones: the coefficient lookup follows the tensor + # lookup across many prefix boundaries + shapes = Tuple{Int, Int}[(129, 97)] + append!(shapes, [(3, 2) for _ in 1:20]) + append!(shapes, [(1, 1) for _ in 1:5]) + bp_scaledcase(AT, Float32, shapes, (2, 1), Strided.BP_ELEMENTWISE) + bp_scaledcase(AT, Float32, shapes, (2, 1), Strided.BP_GROUPTILE) + end + + @testset "offset / strided / negative-stride views with guard band, strategy=$strategy" for + strategy in (Strided.BP_ELEMENTWISE, Strided.BP_GROUPTILE) + P = (210, 210) + plan = bp_scaled_viewcase(AT, Float32, (2, 1), P, P, [ + ((3:42, 5:71), (2:68, 7:46)), # offset sub-blocks, dense + ((81:-2:3, 5:71), (2:3:200, 7:46)), # reversed stride-2 dst, stride-3 src + ((42:-1:3, 5:71), (68:-1:2, 7:46)), # both axis-1 reversed + ((100:139, 7:73), (50:116, 150:189)), # far offsets + ], strategy, Float32[2, -1, 0.5, 5], Float32[0, 1, -3, 0]) + @test plan.family === Strided.BP_TRANSPOSE && plan.strategy === strategy + # the same through the 16 edge (all extents within 16) + P = (60, 60) + plan = bp_scaled_viewcase(AT, Float32, (2, 1), P, P, [ + ((3:18, 5:16), (2:13, 7:22)), + ((33:-2:3, 5:16), (35:-3:2, 7:22)), + ((18:-1:3, 5:16), (13:-1:2, 7:22)), + ], strategy, Float32[-3, 1, 0.5], Float32[2, 0, 1]) + @test plan.strategy === strategy + strategy === Strided.BP_GROUPTILE && @test plan.tile == (16, 16) + end + + @testset "omitted alpha/beta is exactly ones/zeros; coefficients never touch the binding cache" begin + T = Float32 + shapes = [(33, 17), (5, 40), (16, 16), (2, 65)] + perm = (2, 1) + B = length(shapes) + srcs_cpu = [bp_exactdata(T, s) for s in shapes] + d0 = [bp_exactdata(T, bp_dstshape(s, perm)) for s in srcs_cpu] + srcs = [AT(s) for s in srcs_cpu] + for strategy in (Strided.BP_ELEMENTWISE, Strided.BP_GROUPTILE) + dsts = [AT(d) for d in d0] + plan = Strided.plan_batched_permutedims(dsts, srcs, perm; strategy) + runwith(a, b) = (copyto!.(dsts, d0); Strided.batched_permutedims!(dsts, srcs, plan; alpha = a, beta = b); map(Array, dsts)) + plain = runwith(nothing, nothing) + b1 = plan.devcache[] + @test b1 isa Ext._BatchedGPUBinding + for (d, s) in zip(plain, srcs_cpu) + @test isequal(d, permutedims(s, perm)) + end + @test all(map(isequal, plain, runwith(ones(T, B), zeros(T, B)))) + alpha = T[2, -1, 0.5, 5] + beta = T[1, 0, -3, 2] + @test all(map(isequal, runwith(alpha, nothing), runwith(alpha, zeros(T, B)))) + @test all(map(isequal, runwith(nothing, beta), runwith(ones(T, B), beta))) + scaled = runwith(alpha, beta) + for b in 1:B + @test isequal(scaled[b], bp_scaledref(alpha[b], srcs_cpu[b], beta[b], d0[b], perm)) + end + # same arrays, different (or no) coefficients: the cached binding is reused as-is + @test plan.devcache[] === b1 + runwith(reverse(alpha), reverse(beta)) + @test plan.devcache[] === b1 + end + end + + @testset "validation happens before any write" begin + T = Float32 + srcs = [AT(bp_exactdata(T, (3, 4))), AT(bp_exactdata(T, (5, 2)))] + dsts = [AT(fill(T(NaN), 4, 3)), AT(fill(T(NaN), 2, 5))] + plan = Strided.plan_batched_permutedims(dsts, srcs, (2, 1)) + @test_throws DimensionMismatch Strided.batched_permutedims!(dsts, srcs, plan; alpha = ones(T, 3)) + @test_throws DimensionMismatch Strided.batched_permutedims!(dsts, srcs, plan; beta = zeros(T, 1)) + @test_throws InexactError Strided.batched_permutedims!(dsts, srcs, plan; alpha = [1 + 2im, 1]) + @test all(d -> all(isnan, Array(d)), dsts) + @test plan.devcache[] === nothing # nothing was uploaded either + end +end diff --git a/test/runtests.jl b/test/runtests.jl index a091a4a..41dafb6 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -21,6 +21,7 @@ if !is_buildkite Strided.disable_threads() include("othertests.jl") include("blasmultests.jl") + include("batched_permutedims.jl") if Base.Threads.nthreads() > 1 println("Running tests multi-threaded:") @@ -28,6 +29,7 @@ if !is_buildkite Strided.set_num_threads(Base.Threads.nthreads() + 1) include("othertests.jl") include("blasmultests.jl") + include("batched_permutedims.jl") Strided.enable_threaded_mul() include("blasmultests.jl")