diff --git a/CHANGELOG.md b/CHANGELOG.md index b7673ae1..44218d99 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,7 @@ # Changelog ## Unreleased +* Faster indexing into contiguous `PyArray`s. * Bug fixes. ## 0.9.36 (2026-09-18) diff --git a/src/Wrap/PyArray.jl b/src/Wrap/PyArray.jl index dc11c085..6ab0a77b 100644 --- a/src/Wrap/PyArray.jl +++ b/src/Wrap/PyArray.jl @@ -672,11 +672,17 @@ end end pyarray_offset(x::PyArray{T,N,M,true}, i::Int) where {T,N,M} = - N == 0 ? 0 : (i - 1) * x.strides[1] -pyarray_offset(x::PyArray{T,1,M,true}, i::Int) where {T,M} = (i - 1) .* x.strides[1] -pyarray_offset(x::PyArray{T,N}, i::Vararg{Int,N}) where {T,N} = sum((i .- 1) .* x.strides) + N == 0 ? 0 : pyarray_offset1(x, i) +pyarray_offset(x::PyArray{T,1,M,true}, i::Int) where {T,M} = pyarray_offset1(x, i) +pyarray_offset(x::PyArray{T,N}, i::Vararg{Int,N}) where {T,N} = + pyarray_offset1(x, i[1]) + sum((Base.tail(i) .- 1) .* Base.tail(x.strides); init = 0) pyarray_offset(x::PyArray{T,0}) where {T} = 0 +# The type can't say the first stride is the element size (strided arrays can be linear too). +# Branching on it lets LLVM version loops on the contiguous case and vectorise them. +pyarray_offset1(x::PyArray{T,N,M,L,R}, i::Int) where {T,N,M,L,R} = + x.strides[1] == sizeof(R) ? (i - 1) * sizeof(R) : (i - 1) * x.strides[1] + function pyarray_load(::Type{T}, p::Ptr{R}) where {T,R} if R == T unsafe_load(p) diff --git a/test/Wrap.jl b/test/Wrap.jl index 09d6ece7..0089081f 100644 --- a/test/Wrap.jl +++ b/test/Wrap.jl @@ -99,6 +99,23 @@ ) @test_throws Exception PyArray(nd; array = false, buffer = true) end + @testset "linear with unit and non-unit first stride" begin + tb = pyimport("_testbuffer") + nd = tb.ndarray( + pylist(1:24), + shape = pylist([4, 6]), + format = "i", + flags = tb.ND_FORTRAN, + ) + e = reshape(1:24, 4, 6) + for (v, ev) in + ((nd, e), (nd[pyslice(nothing, nothing, 2), pyslice(nothing)], e[1:2:end, :])) + a = PyArray(v; array = false) + @test a isa PyArray{Cint,2,false,true,Cint} + @test [a[i, j] for i in axes(a, 1), j in axes(a, 2)] == ev + @test [a[i] for i in eachindex(IndexLinear(), a)] == vec(ev) + end + end end @testitem "PyDict" begin