We read every piece of feedback, and take your input very seriously.
To see all available qualifiers, see our documentation.
There was an error while loading. Please reload this page.
2 parents e8f8279 + bd4b649 commit ae089e6Copy full SHA for ae089e6
7 files changed
src/aggregations/aggregation_stack.jl
@@ -67,8 +67,18 @@ AggregationStack(fs::AbstractAggregation...) = AggregationStack(fs)
67
68
Flux.@layer :ignore AggregationStack
69
70
-function (a::AggregationStack)(x::Maybe{AbstractArray}, bags::AbstractBags, args...)
71
- reduce(vcat, (f(x, bags, args...) for f in a.fs))
+# function (a::AggregationStack)(x::Maybe{AbstractArray}, bags::AbstractBags, args...)
+# 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
82
end
83
84
Flux.@forward AggregationStack.fs Base.getindex, Base.firstindex, Base.lastindex, Base.first,
src/aggregations/aggregations.jl
@@ -22,7 +22,7 @@ abstract type AbstractAggregation end
22
@inline _weightsum(ws::AbstractVector, i) = ws[i]
23
24
# more stable definitions for r_map and p_map
25
-ChainRulesCore.rrule(::typeof(softplus), x) = softplus.(x), Δ -> (NoTangent(), Δ .* σ.(x))
+ChainRulesCore.rrule(::typeof(softplus), x) = softplus.(x), Δ -> (NoTangent(), unthunk(Δ) .* σ.(x))
26
27
# our definition of type min for Maybe{...} types
28
_typemin(t::Type) = typemin(t)
src/aggregations/segmented_lse.jl
@@ -137,13 +137,13 @@ function ChainRulesCore.rrule(::typeof(segmented_lse_forw),
137
x::AbstractMatrix, ψ::AbstractVector, r::AbstractVector, bags::AbstractBags)
138
M = _lse_precomp(x, r, bags)
139
y = _segmented_lse_norm(x, ψ, r, bags, M)
140
- grad = Δ -> (NoTangent(), segmented_lse_back(Δ, y, x, ψ, r, bags, M)...)
+ grad = Δ -> (NoTangent(), segmented_lse_back(unthunk(Δ), y, x, ψ, r, bags, M)...)
141
y, grad
142
143
144
function ChainRulesCore.rrule(::typeof(segmented_lse_forw),
145
x::Missing, ψ::AbstractVector, r::AbstractVector, bags::AbstractBags)
146
y = segmented_lse_forw(x, ψ, r, bags)
147
- grad = Δ -> (NoTangent(), segmented_lse_back(Δ, x, ψ, bags)...)
+ grad = Δ -> (NoTangent(), segmented_lse_back(unthunk(Δ), x, ψ, bags)...)
148
149
src/aggregations/segmented_max.jl
@@ -97,6 +97,6 @@ end
97
98
function ChainRulesCore.rrule(::typeof(segmented_max_forw), args...)
99
y = segmented_max_forw(args...)
100
- grad = Δ -> (NoTangent(), segmented_max_back(Δ, y, args...)...)
+ grad = Δ -> (NoTangent(), segmented_max_back(unthunk(Δ), y, args...)...)
101
102
src/aggregations/segmented_mean.jl
@@ -98,6 +98,6 @@ end
function ChainRulesCore.rrule(::typeof(segmented_mean_forw), args...)
y = segmented_mean_forw(args...)
- grad = Δ -> (NoTangent(), segmented_mean_back(Δ, y, args...)...)
+ grad = Δ -> (NoTangent(), segmented_mean_back(unthunk(Δ), y, args...)...)
103
src/aggregations/segmented_pnorm.jl
@@ -141,12 +141,12 @@ end
function ChainRulesCore.rrule(::typeof(segmented_pnorm_forw), a::AbstractMatrix, ψ, p, bags, w)
M = _pnorm_precomp(a, bags)
y = _segmented_pnorm_norm(a, ψ, p, bags, w, M)
- grad = Δ -> (NoTangent(), segmented_pnorm_back(Δ, y, a, ψ, p, bags, w, M)...)
+ grad = Δ -> (NoTangent(), segmented_pnorm_back(unthunk(Δ), y, a, ψ, p, bags, w, M)...)
function ChainRulesCore.rrule(::typeof(segmented_pnorm_forw), a::Missing, ψ, p, bags, w)
y = segmented_pnorm_forw(a, ψ, p, bags, w)
150
- grad = Δ -> (NoTangent(), segmented_pnorm_back(Δ, y, ψ, bags)...)
+ grad = Δ -> (NoTangent(), segmented_pnorm_back(unthunk(Δ), y, ψ, bags)...)
151
152
src/aggregations/segmented_sum.jl
@@ -56,6 +56,7 @@ function segmented_sum_forw(x::AbstractMatrix, ψ::AbstractVector, bags::Abstrac
56
57
58
function segmented_sum_back(Δ, y, x, ψ, bags, w)
59
+ Δ = unthunk(Δ)
60
dx = zero(x)
61
dψ = zero(ψ)
62
dw = isnothing(w) ? ZeroTangent() : zero(w)
@@ -96,6 +97,6 @@ end
96
function ChainRulesCore.rrule(::typeof(segmented_sum_forw), args...)
y = segmented_sum_forw(args...)
- grad = Δ -> (NoTangent(), segmented_sum_back(Δ, y, args...)...)
+ grad = Δ -> (NoTangent(), segmented_sum_back(unthunk(Δ), y, args...)...)
0 commit comments