diff --git a/ext/ArrayInterfaceReverseDiffExt.jl b/ext/ArrayInterfaceReverseDiffExt.jl index 3a000def..50d52ddf 100644 --- a/ext/ArrayInterfaceReverseDiffExt.jl +++ b/ext/ArrayInterfaceReverseDiffExt.jl @@ -15,5 +15,13 @@ end function ArrayInterface.restructure(x::Array, y::ReverseDiff.TrackedArray) reshape(y, Base.size(x)...) end +function ArrayInterface.restructure(x::ReverseDiff.TrackedArray, y::ReverseDiff.TrackedArray) + reshape(y, Base.size(x)...) +end +function ArrayInterface.restructure( + x::ReverseDiff.TrackedArray, y::AbstractArray{<:ReverseDiff.TrackedReal} +) + reshape(ArrayInterface.aos_to_soa(y), Base.size(x)...) +end end # module diff --git a/ext/ArrayInterfaceTrackerExt.jl b/ext/ArrayInterfaceTrackerExt.jl index 5723d9f1..6a753d5f 100644 --- a/ext/ArrayInterfaceTrackerExt.jl +++ b/ext/ArrayInterfaceTrackerExt.jl @@ -7,11 +7,21 @@ ArrayInterface.ismutable(::Type{<:Tracker.TrackedArray}) = false ArrayInterface.ismutable(T::Type{<:Tracker.TrackedReal}) = false ArrayInterface.can_setindex(::Type{<:Tracker.TrackedArray}) = false ArrayInterface.fast_scalar_indexing(::Type{<:Tracker.TrackedArray}) = false -ArrayInterface.aos_to_soa(x::AbstractArray{<:Tracker.TrackedReal,N}) where {N} = Tracker.collect(x) +function ArrayInterface.aos_to_soa(x::AbstractArray{<:Tracker.TrackedReal, N}) where {N} + Tracker.collect(x) +end function ArrayInterface.restructure(x::Array, y::Tracker.TrackedArray) reshape(y, Base.size(x)...) end +function ArrayInterface.restructure(x::Tracker.TrackedArray, y::Tracker.TrackedArray) + reshape(y, Base.size(x)...) +end +function ArrayInterface.restructure( + x::Tracker.TrackedArray, y::AbstractArray{<:Tracker.TrackedReal} +) + reshape(ArrayInterface.aos_to_soa(y), Base.size(x)...) +end function ArrayInterface.restructure(x::Array, y::Array{<:Tracker.TrackedReal}) reshape(y, Base.size(x)...) end diff --git a/test/ad.jl b/test/ad.jl index 3c61873e..aa2bb03a 100644 --- a/test/ad.jl +++ b/test/ad.jl @@ -4,38 +4,50 @@ x = ReverseDiff.track([4.0]) x = reshape([ReverseDiff.track(rand(1, 1, 1))[1]], 1, 1, 1) @test ArrayInterface.aos_to_soa(x) isa ReverseDiff.TrackedArray @test ndims(ArrayInterface.aos_to_soa(x)) == 3 -x = reduce(vcat, ReverseDiff.track([4.0,4.0])) +x = reduce(vcat, ReverseDiff.track([4.0, 4.0])) @test ArrayInterface.aos_to_soa(x) isa ReverseDiff.TrackedArray x = [ReverseDiff.track([4.0])[1]] @test ArrayInterface.aos_to_soa(x) isa ReverseDiff.TrackedArray -x = reduce(vcat, ReverseDiff.track([4.0,4.0])) -x = [x[1],x[2]] +x = reduce(vcat, ReverseDiff.track([4.0, 4.0])) +x = [x[1], x[2]] @test ArrayInterface.aos_to_soa(x) isa ReverseDiff.TrackedArray x = Tracker.TrackedArray([4.0]) @test ArrayInterface.aos_to_soa(x) isa Tracker.TrackedArray x = [Tracker.TrackedArray([4.0])[1]] @test ArrayInterface.aos_to_soa(x) isa Tracker.TrackedArray -x = Tracker.TrackedArray([4.0,4.0]) +x = Tracker.TrackedArray([4.0, 4.0]) @test ArrayInterface.aos_to_soa(x) isa Tracker.TrackedArray -x = reduce(vcat, Tracker.TrackedArray([4.0,4.0])) -x = [x[1],x[2]] +x = reduce(vcat, Tracker.TrackedArray([4.0, 4.0])) +x = [x[1], x[2]] @test ArrayInterface.aos_to_soa(x) isa Tracker.TrackedArray x = rand(4) -y = Tracker.TrackedReal.(rand(2,2)) +y = Tracker.TrackedReal.(rand(2, 2)) @test ArrayInterface.restructure(x, y) isa Array @test eltype(ArrayInterface.restructure(x, y)) <: Tracker.TrackedReal @test size(ArrayInterface.restructure(x, y)) == (4,) -y = Tracker.TrackedArray(rand(2,2)) +y = Tracker.TrackedArray(rand(2, 2)) +@test ArrayInterface.restructure(x, y) isa Tracker.TrackedArray +@test size(ArrayInterface.restructure(x, y)) == (4,) +x = Tracker.TrackedArray(rand(4)) +@test ArrayInterface.restructure(x, y) isa Tracker.TrackedArray +@test size(ArrayInterface.restructure(x, y)) == (4,) +y = Tracker.TrackedReal.(rand(2, 2)) @test ArrayInterface.restructure(x, y) isa Tracker.TrackedArray @test size(ArrayInterface.restructure(x, y)) == (4,) x = rand(4) -y = ReverseDiff.track(rand(2,2)) +y = ReverseDiff.track(rand(2, 2)) @test ArrayInterface.restructure(x, y) isa ReverseDiff.TrackedArray @test size(ArrayInterface.restructure(x, y)) == (4,) -y = ReverseDiff.track.(rand(2,2)) +x = ReverseDiff.track(rand(4)) +@test ArrayInterface.restructure(x, y) isa ReverseDiff.TrackedArray +@test size(ArrayInterface.restructure(x, y)) == (4,) +y = ReverseDiff.track.(rand(2, 2)) +@test ArrayInterface.restructure(x, y) isa ReverseDiff.TrackedArray +@test size(ArrayInterface.restructure(x, y)) == (4,) +x = rand(4) @test ArrayInterface.restructure(x, y) isa Array @test eltype(ArrayInterface.restructure(x, y)) <: ReverseDiff.TrackedReal @test size(ArrayInterface.restructure(x, y)) == (4,)