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
6 changes: 3 additions & 3 deletions Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -5,22 +5,22 @@ version = "2.11.4"

[deps]
ArrayInterface = "4fba245c-0d91-5ea0-9b3e-6abc04ee57a9"
ChainRulesCore = "d360d2e6-b24c-11e9-a2a3-2a2ae2dbcce4"
DocStringExtensions = "ffbed154-4ef7-542d-bbb7-c09d3a79fcae"
LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e"
RecipesBase = "3cdcf5f2-1ef4-517c-9805-6587b60abb01"
Requires = "ae029012-a4dd-5104-9daa-d747884805df"
StaticArrays = "90137ffa-7385-5640-81b9-e52037218182"
Statistics = "10745b16-79ce-11e8-11f9-7d13ad32a3b2"
ZygoteRules = "700de1a5-db45-46bc-99cf-38207098b444"

[compat]
ArrayInterface = "2.7, 3.0"
ChainRulesCore = "0.10.7"
DocStringExtensions = "0.8"
RecipesBase = "0.7, 0.8, 1.0"
Requires = "0.5, 1.0"
StaticArrays = "0.12, 1.0"
ZygoteRules = "0.2"
julia = "1.3"
julia = "1.6"

[extras]
ForwardDiff = "f6369f11-7733-5829-9624-2563aa707210"
Expand Down
4 changes: 3 additions & 1 deletion src/RecursiveArrayTools.jl
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,10 @@ module RecursiveArrayTools

using DocStringExtensions
using Requires, RecipesBase, StaticArrays, Statistics,
ArrayInterface, ZygoteRules, LinearAlgebra
ArrayInterface, LinearAlgebra

import ChainRulesCore
import ChainRulesCore: NoTangent
abstract type AbstractVectorOfArray{T, N, A} <: AbstractArray{T, N} end
abstract type AbstractDiffEqArray{T, N, A} <: AbstractVectorOfArray{T, N, A} end

Expand Down
6 changes: 3 additions & 3 deletions src/init.jl
Original file line number Diff line number Diff line change
Expand Up @@ -17,16 +17,16 @@ function __init__()
return CuArrays.CuArray(reshape(reduce(hcat,vecs),size(VA.u[1])...,length(VA.u)))
end
Base.convert(::Type{<:CuArrays.CuArray},VA::AbstractVectorOfArray) = CuArrays.CuArray(VA)
@adjoint CuArrays.CuArray(xs::AbstractVectorOfArray) = CuArrays.CuArray(xs), ȳ -> (ȳ,)
ChainRules.rrule(::Type{<:CuArrays.CuArray},xs::AbstractVectorOfArray) = CuArrays.CuArray(xs), ȳ -> (NoTangent(),ȳ)
end

@require CUDA="052768ef-5323-5732-b1bb-66c8b64840ba" begin
function CUDA.CuArray(VA::AbstractVectorOfArray)
vecs = vec.(VA.u)
return CUDA.CuArray(reshape(reduce(hcat,vecs),size(VA.u[1])...,length(VA.u)))
end
Base.convert(::Type{<:CUDA.CuArray},VA::AbstractVectorOfArray) = CUDA.CuArray(VA)
@adjoint CUDA.CuArray(xs::AbstractVectorOfArray) = CUDA.CuArray(xs), ȳ -> (ȳ,)
ChainRules.rrule(::Type{<:CUDA.CuArray},xs::AbstractVectorOfArray) = CUDA.CuArray(xs), ȳ -> (NoTangent(),ȳ)
end

@require Tracker="9f7883ad-71c0-57eb-9f7f-b5c9e6d3789c" begin
Expand Down
33 changes: 18 additions & 15 deletions src/zygote.jl
Original file line number Diff line number Diff line change
@@ -1,41 +1,44 @@
ZygoteRules.@adjoint function getindex(VA::AbstractVectorOfArray, i)
function ChainRulesCore.rrule(::typeof(getindex),VA::AbstractVectorOfArray, i)
function AbstractVectorOfArray_getindex_adjoint(Δ)
Δ′ = [ (i == j ? Δ : zero(x)) for (x,j) in zip(VA.u, 1:length(VA))]
(Δ′,nothing)
(NoTangent(),Δ′,NoTangent())
end
VA[i],AbstractVectorOfArray_getindex_adjoint
end

ZygoteRules.@adjoint function getindex(VA::AbstractVectorOfArray, i, j...)
function ChainRulesCore.rrule(::typeof(getindex),VA::AbstractVectorOfArray, i, j...)
function AbstractVectorOfArray_getindex_adjoint(Δ)
Δ′ = zero(VA)
Δ′[i,j...] = Δ
(Δ′, i,map(_ -> nothing, j)...)
(NoTangent(), Δ′, i,map(_ -> NoTangent(), j)...)
end
VA[i,j...],AbstractVectorOfArray_getindex_adjoint
end

ZygoteRules.@adjoint function ArrayPartition(x::S, ::Type{Val{copy_x}} = Val{false}) where {S<:Tuple,copy_x}
function ChainRulesCore.rrule(::Type{<:ArrayPartition}, x::S, ::Type{Val{copy_x}} = Val{false}) where {S<:Tuple,copy_x}
function ArrayPartition_adjoint(_y)
y = Array(_y)
starts = vcat(0,cumsum(reduce(vcat,length.(x))))
ntuple(i -> reshape(y[starts[i]+1:starts[i+1]], size(x[i])), length(x)), nothing
NoTangent(), ntuple(i -> reshape(y[starts[i]+1:starts[i+1]], size(x[i])), length(x)), NoTangent()
end

ArrayPartition(x, Val{copy_x}), ArrayPartition_adjoint
end

ZygoteRules.@adjoint function VectorOfArray(u)
VectorOfArray(u),y -> ([y[ntuple(x->Colon(),ndims(y)-1)...,i] for i in 1:size(y)[end]],)
function ChainRulesCore.rrule(::Type{<:VectorOfArray},u)
VectorOfArray(u),y -> (NoTangent(),[y[ntuple(x->Colon(),ndims(y)-1)...,i] for i in 1:size(y)[end]])
end

ZygoteRules.@adjoint function DiffEqArray(u,t)
DiffEqArray(u,t),y -> ([y[ntuple(x->Colon(),ndims(y)-1)...,i] for i in 1:size(y)[end]],nothing)
function ChainRulesCore.rrule(::Type{<:DiffEqArray},u,t)
DiffEqArray(u,t),y -> (NoTangent(),[y[ntuple(x->Colon(),ndims(y)-1)...,i] for i in 1:size(y)[end]],NoTangent())
end

ZygoteRules.@adjoint function ZygoteRules.literal_getproperty(A::ArrayPartition, ::Val{:x})
function literal_ArrayPartition_x_adjoint(d)
(ArrayPartition((isnothing(d[i]) ? zero(A.x[i]) : d[i] for i in 1:length(d))...),)
end
A.x,literal_ArrayPartition_x_adjoint
function ChainRulesCore.rrule(::typeof(getproperty),A::ArrayPartition, s::Symbol)
if s !== :x
error("$s is not a field of ArrayPartition")
end
function literal_ArrayPartition_x_adjoint(d)
(NoTangent(),ArrayPartition((isnothing(d[i]) ? zero(A.x[i]) : d[i] for i in 1:length(d))...))
end
A.x,literal_ArrayPartition_x_adjoint
end