Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
79 changes: 63 additions & 16 deletions src/stridedview.jl
Original file line number Diff line number Diff line change
Expand Up @@ -39,18 +39,23 @@ function StridedView(
size::NTuple{N, Int} = size(parent),
strides::NTuple{N, Int} = strides(parent),
offset::Int = 0,
op::F = identity
op::F = identity;
normalize::Bool = true
) where {N, F}
T = Base.promote_op(op, eltype(parent))
return StridedView{T}(parent, size, strides, offset, op)
return StridedView{T}(parent, size, strides, offset, op; normalize)
end
function StridedView{T}(
parent::DenseArray,
size::NTuple{N, Int} = size(parent),
strides::NTuple{N, Int} = strides(parent),
offset::Int = 0,
op::F = identity
op::F = identity;
normalize::Bool = true
) where {T, N, F}
# `normalize = false` asserts that `parent` and `strides` are already normalized,
# e.g. because they are taken from an existing `StridedView`.
normalize || return StridedView{T, N, typeof(parent), F}(parent, size, strides, offset, op)
parent′ = _normalizeparent(parent)
strides′ = _normalizestrides(size, strides)
return StridedView{T, N, typeof(parent′), F}(parent′, size, strides′, offset, op)
Expand All @@ -71,7 +76,10 @@ function StridedView(a::Base.ReinterpretArray{T, N}) where {T, N}
throw(ArgumentError("Cannot create StridedView with reinterpretation from $S to $T"))
b.op isa FN ||
throw(ArgumentError("Cannot create StridedView with reinterpretation from view with non-identity operation"))
return StridedView{T}(b.parent, size(b), strides(b), offset(b))
return StridedView{T}(
b.parent, size(b), strides(b), offset(b), identity;
normalize = false
)
end

# trait
Expand Down Expand Up @@ -140,12 +148,15 @@ end

# Indexing with slice indices to create a new view.
function Base.getindex(a::StridedView{T, N}, I::Vararg{SliceIndex, N}) where {T, N}
newsize = _computeviewsize(a.size, I)
newstrides = _normalizestrides(newsize, _computeviewstrides(a.strides, I))
return StridedView{T}(
a.parent,
_computeviewsize(a.size, I),
_computeviewstrides(a.strides, I),
newsize,
newstrides,
a.offset + _computeviewoffset(a.strides, I),
a.op
a.op;
normalize = false
)
end

Expand All @@ -170,20 +181,30 @@ end
#----------------------------------------------------------------------------
Base.conj(a::StridedView{<:Real}) = a
function Base.conj(a::StridedView{T}) where {T <: Complex}
return StridedView{T}(a.parent, a.size, a.strides, a.offset, _conj(a.op))
newop = _conj(a.op)
return StridedView{T}(
a.parent, a.size, a.strides, a.offset, newop;
normalize = false
)
end
function Base.conj(a::StridedView)
S = Base.promote_op(a.op, eltype(a))
newop = _conj(a.op)
T = Base.promote_op(newop, S)
return StridedView{T}(a.parent, a.size, a.strides, a.offset, newop)
return StridedView{T}(
a.parent, a.size, a.strides, a.offset, newop;
normalize = false
)
end

function Base.permutedims(a::StridedView{T, N}, p) where {T, N}
_isperm(N, p) || throw(ArgumentError("Invalid permutation of length $N: $p"))
newsize = ntuple(n -> size(a, p[n]), Val(N))
newstrides = ntuple(n -> stride(a, p[n]), Val(N))
return StridedView{T}(a.parent, newsize, newstrides, a.offset, a.op)
return StridedView{T}(
a.parent, newsize, _normalizestrides(newsize, newstrides), a.offset, a.op;
normalize = false
)
end

LinearAlgebra.transpose(a::StridedView{<:Number, 2}) = permutedims(a, (2, 1))
Expand All @@ -192,29 +213,51 @@ function LinearAlgebra.adjoint(a::StridedView{<:Any, 2}) # act recursively, like
S = Base.promote_op(a.op, eltype(a))
newop = _adjoint(a.op)
T = Base.promote_op(newop, S)
return permutedims(StridedView{T}(a.parent, a.size, a.strides, a.offset, newop), (2, 1))
return permutedims(
StridedView{T}(
a.parent, a.size, a.strides, a.offset, newop;
normalize = false
), (2, 1)
)
end
function LinearAlgebra.transpose(a::StridedView{<:Any, 2}) # act recursively, like Base
S = Base.promote_op(a.op, eltype(a))
newop = _transpose(a.op)
T = Base.promote_op(newop, S)
return permutedims(StridedView{T}(a.parent, a.size, a.strides, a.offset, newop), (2, 1))
return permutedims(
StridedView{T}(
a.parent, a.size, a.strides, a.offset, newop;
normalize = false
), (2, 1)
)
end

Base.map(::FC, a::StridedView{<:Real}) = a
Base.map(::FT, a::StridedView{<:Number}) = a
Base.map(::FA, a::StridedView{<:Number}) = conj(a)
function Base.map(::FC, a::StridedView)
T = Base.promote_op(conj, eltype(a))
return StridedView{T}(a.parent, a.size, a.strides, a.offset, _conj(a.op))
newop = _conj(a.op)
return StridedView{T}(
a.parent, a.size, a.strides, a.offset, newop;
normalize = false
)
end
function Base.map(::FT, a::StridedView)
T = Base.promote_op(transpose, eltype(a))
return StridedView{T}(a.parent, a.size, a.strides, a.offset, _transpose(a.op))
newop = _transpose(a.op)
return StridedView{T}(
a.parent, a.size, a.strides, a.offset, newop;
normalize = false
)
end
function Base.map(::FA, a::StridedView)
T = Base.promote_op(adjoint, eltype(a))
return StridedView{T}(a.parent, a.size, a.strides, a.offset, _adjoint(a.op))
newop = _adjoint(a.op)
return StridedView{T}(
a.parent, a.size, a.strides, a.offset, newop;
normalize = false
)
end

# Creating or transforming StridedView by slicing
Expand Down Expand Up @@ -251,14 +294,18 @@ end
# we cannot use Base.reshape, as this also accepts indices that might not preserve stridedness
sreshape(a, args::Vararg{Int}) = sreshape(a, args)
function sreshape(a::StridedView{T}, newsize::Dims) where {T}
size(a) == newsize && return a
if any(isequal(0), newsize)
any(isequal(0), size(a)) || throw(DimensionMismatch())
newstrides = one.(newsize)
else
newstrides = _computereshapestrides(newsize, _simplifydims(size(a), strides(a))...)
end
isnothing(newstrides) && throw(ReshapeException(newsize, size(a), strides(a)))
return StridedView{T}(a.parent, newsize, newstrides, a.offset, a.op)
return StridedView{T}(
a.parent, newsize, _normalizestrides(newsize, newstrides), a.offset, a.op;
normalize = false
)
end

sreshape(a::AbstractArray, newsize::Dims) = sreshape(StridedView(a), newsize)
Expand Down
9 changes: 9 additions & 0 deletions test/runtests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -222,14 +222,23 @@ if !is_buildkite
@testset "more reshape" begin
A = randn(4, 0)
B = StridedView(A)
@test sreshape(B, size(B)) === B
@test_throws DimensionMismatch sreshape(B, (4, 1))
C = sreshape(B, (2, 1, 2, 0, 1))
@test strides(C) == (1, 2, 2, 4, 0)
@test sreshape(C, (4, 0)) == A

A = randn(4, 1, 2)
B = StridedView(A)
@test_throws DimensionMismatch sreshape(B, (4, 4))
@test sreshape(B, size(B)) === B
@test strides(permutedims(B, (2, 1, 3))) == (1, 1, 4)
@test strides(B[1:2:3, :, :]) == (2, 4, 4)
@test StridedView(
B.parent, size(B), strides(B), B.offset, B.op; normalize = false
) === B
C = sreshape(B, (2, 1, 1, 4, 1, 1))
@test strides(C) == (1, 2, 2, 2, 8, 8)
@test C == reshape(A, (2, 1, 1, 4, 1, 1))
@test sreshape(C, (4, 1, 2)) == A
end
Expand Down
Loading