Skip to content

Commit ae089e6

Browse files
authored
Merge pull request #130 from CTUAvastLab/unthunk
This fixes the problem that Zygote started to use unthunk.
2 parents e8f8279 + bd4b649 commit ae089e6

7 files changed

Lines changed: 21 additions & 10 deletions

File tree

src/aggregations/aggregation_stack.jl

Lines changed: 12 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -67,8 +67,18 @@ AggregationStack(fs::AbstractAggregation...) = AggregationStack(fs)
6767

6868
Flux.@layer :ignore AggregationStack
6969

70-
function (a::AggregationStack)(x::Maybe{AbstractArray}, bags::AbstractBags, args...)
71-
reduce(vcat, (f(x, bags, args...) for f in a.fs))
70+
# function (a::AggregationStack)(x::Maybe{AbstractArray}, bags::AbstractBags, args...)
71+
# reduce(vcat, (f(x, bags, args...) for f in a.fs))
72+
# end
73+
74+
@generated function (a::AggregationStack{T})(x::Maybe{AbstractArray}, bags::AbstractBags, args...) where {T<:Tuple}
75+
l = T.parameters |> length
76+
chs = map(1:l) do i
77+
:(a.fs[$i](x, bags, args...))
78+
end
79+
quote
80+
vcat($(chs...))
81+
end
7282
end
7383

7484
Flux.@forward AggregationStack.fs Base.getindex, Base.firstindex, Base.lastindex, Base.first,

src/aggregations/aggregations.jl

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@ abstract type AbstractAggregation end
2222
@inline _weightsum(ws::AbstractVector, i) = ws[i]
2323

2424
# more stable definitions for r_map and p_map
25-
ChainRulesCore.rrule(::typeof(softplus), x) = softplus.(x), Δ -> (NoTangent(), Δ .* σ.(x))
25+
ChainRulesCore.rrule(::typeof(softplus), x) = softplus.(x), Δ -> (NoTangent(), unthunk(Δ) .* σ.(x))
2626

2727
# our definition of type min for Maybe{...} types
2828
_typemin(t::Type) = typemin(t)

src/aggregations/segmented_lse.jl

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -137,13 +137,13 @@ function ChainRulesCore.rrule(::typeof(segmented_lse_forw),
137137
x::AbstractMatrix, ψ::AbstractVector, r::AbstractVector, bags::AbstractBags)
138138
M = _lse_precomp(x, r, bags)
139139
y = _segmented_lse_norm(x, ψ, r, bags, M)
140-
grad = Δ -> (NoTangent(), segmented_lse_back(Δ, y, x, ψ, r, bags, M)...)
140+
grad = Δ -> (NoTangent(), segmented_lse_back(unthunk(Δ), y, x, ψ, r, bags, M)...)
141141
y, grad
142142
end
143143

144144
function ChainRulesCore.rrule(::typeof(segmented_lse_forw),
145145
x::Missing, ψ::AbstractVector, r::AbstractVector, bags::AbstractBags)
146146
y = segmented_lse_forw(x, ψ, r, bags)
147-
grad = Δ -> (NoTangent(), segmented_lse_back(Δ, x, ψ, bags)...)
147+
grad = Δ -> (NoTangent(), segmented_lse_back(unthunk(Δ), x, ψ, bags)...)
148148
y, grad
149149
end

src/aggregations/segmented_max.jl

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -97,6 +97,6 @@ end
9797

9898
function ChainRulesCore.rrule(::typeof(segmented_max_forw), args...)
9999
y = segmented_max_forw(args...)
100-
grad = Δ -> (NoTangent(), segmented_max_back(Δ, y, args...)...)
100+
grad = Δ -> (NoTangent(), segmented_max_back(unthunk(Δ), y, args...)...)
101101
y, grad
102102
end

src/aggregations/segmented_mean.jl

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -98,6 +98,6 @@ end
9898

9999
function ChainRulesCore.rrule(::typeof(segmented_mean_forw), args...)
100100
y = segmented_mean_forw(args...)
101-
grad = Δ -> (NoTangent(), segmented_mean_back(Δ, y, args...)...)
101+
grad = Δ -> (NoTangent(), segmented_mean_back(unthunk(Δ), y, args...)...)
102102
y, grad
103103
end

src/aggregations/segmented_pnorm.jl

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -141,12 +141,12 @@ end
141141
function ChainRulesCore.rrule(::typeof(segmented_pnorm_forw), a::AbstractMatrix, ψ, p, bags, w)
142142
M = _pnorm_precomp(a, bags)
143143
y = _segmented_pnorm_norm(a, ψ, p, bags, w, M)
144-
grad = Δ -> (NoTangent(), segmented_pnorm_back(Δ, y, a, ψ, p, bags, w, M)...)
144+
grad = Δ -> (NoTangent(), segmented_pnorm_back(unthunk(Δ), y, a, ψ, p, bags, w, M)...)
145145
y, grad
146146
end
147147

148148
function ChainRulesCore.rrule(::typeof(segmented_pnorm_forw), a::Missing, ψ, p, bags, w)
149149
y = segmented_pnorm_forw(a, ψ, p, bags, w)
150-
grad = Δ -> (NoTangent(), segmented_pnorm_back(Δ, y, ψ, bags)...)
150+
grad = Δ -> (NoTangent(), segmented_pnorm_back(unthunk(Δ), y, ψ, bags)...)
151151
y, grad
152152
end

src/aggregations/segmented_sum.jl

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,7 @@ function segmented_sum_forw(x::AbstractMatrix, ψ::AbstractVector, bags::Abstrac
5656
end
5757

5858
function segmented_sum_back(Δ, y, x, ψ, bags, w)
59+
Δ = unthunk(Δ)
5960
dx = zero(x)
6061
= zero(ψ)
6162
dw = isnothing(w) ? ZeroTangent() : zero(w)
@@ -96,6 +97,6 @@ end
9697

9798
function ChainRulesCore.rrule(::typeof(segmented_sum_forw), args...)
9899
y = segmented_sum_forw(args...)
99-
grad = Δ -> (NoTangent(), segmented_sum_back(Δ, y, args...)...)
100+
grad = Δ -> (NoTangent(), segmented_sum_back(unthunk(Δ), y, args...)...)
100101
y, grad
101102
end

0 commit comments

Comments
 (0)