Skip to content

Commit 7bc16eb

Browse files
committed
Move to @scalar_rule
1 parent 3b9be76 commit 7bc16eb

1 file changed

Lines changed: 23 additions & 103 deletions

File tree

ext/SpecialFunctionsChainRulesCoreExt.jl

Lines changed: 23 additions & 103 deletions
Original file line numberDiff line numberDiff line change
@@ -622,110 +622,30 @@ function _beta_inc_grad(a::T, b::T, x::T) where {T<:Union{Float64, Float32}}
622622
end
623623
end
624624

625-
626-
627-
628-
629625
# Incomplete beta: beta_inc(a,b,x) -> (p, q) with q=1-p
630-
function ChainRulesCore.frule((_, Δa, Δb, Δx), ::typeof(beta_inc), a::Number, b::Number, x::Number)
631-
# primal
632-
p, q = beta_inc(a, b, x)
633-
# derivatives
634-
_a, _b, _x = map(float, promote(a, b, x))
635-
dIa, dIb, dIx = _beta_inc_grad(_a, _b, _x)
636-
Δp = muladd(dIx, Δx, muladd(dIb, Δb, dIa * Δa))
637-
Δq = -Δp
638-
Tout = typeof((p, q))
639-
return (p, q), ChainRulesCore.Tangent{Tout}(Δp, Δq)
640-
end
641-
642-
function ChainRulesCore.rrule(::typeof(beta_inc), a::Number, b::Number, x::Number)
643-
p, q = beta_inc(a, b, x)
644-
Ta = ChainRulesCore.ProjectTo(a)
645-
Tb = ChainRulesCore.ProjectTo(b)
646-
Tx = ChainRulesCore.ProjectTo(x)
647-
_a, _b, _x = map(float, promote(a, b, x))
648-
dIa, dIb, dIx = _beta_inc_grad(_a, _b, _x)
649-
function beta_inc_pullback(Δ)
650-
Δp, Δq = Δ
651-
s = Δp - Δq # because q = 1 - p
652-
= Ta(s * dIa)
653-
= Tb(s * dIb)
654-
= Tx(s * dIx)
655-
return ChainRulesCore.NoTangent(), ā, b̄, x̄
656-
end
657-
return (p, q), beta_inc_pullback
658-
end
659-
function ChainRulesCore.frule((_, Δa, Δb, Δx, Δy), ::typeof(beta_inc), a::Number, b::Number, x::Number, y::Number)
660-
p, q = beta_inc(a, b, x, y)
661-
_a, _b, _x, _y = map(float, promote(a, b, x, y))
662-
dIa, dIb, dIx = _beta_inc_grad(_a, _b, _x)
663-
Δp = muladd(dIx, Δx, muladd(-dIx, Δy, muladd(dIb, Δb, dIa * Δa)))
664-
Δq = -Δp
665-
Tout = typeof((p, q))
666-
return (p, q), ChainRulesCore.Tangent{Tout}(Δp, Δq)
667-
end
668-
669-
function ChainRulesCore.rrule(::typeof(beta_inc), a::Number, b::Number, x::Number, y::Number)
670-
p, q = beta_inc(a, b, x, y)
671-
Ta = ChainRulesCore.ProjectTo(a)
672-
Tb = ChainRulesCore.ProjectTo(b)
673-
Tx = ChainRulesCore.ProjectTo(x)
674-
Ty = ChainRulesCore.ProjectTo(y)
675-
_a, _b, _x, _y = map(float, promote(a, b, x, y))
676-
dIa, dIb, dIx = _beta_inc_grad(_a, _b, _x)
677-
function beta_inc_pullback(Δ)
678-
Δp, Δq = Δ
679-
s = Δp - Δq
680-
= Ta(s * dIa)
681-
= Tb(s * dIb)
682-
= Tx(s * dIx)
683-
= Ty(-s * dIx)
684-
return ChainRulesCore.NoTangent(), ā, b̄, x̄, ȳ
685-
end
686-
return (p, q), beta_inc_pullback
687-
end
688-
626+
ChainRulesCore.@scalar_rule(
627+
beta_inc(a::Number, b::Number, x::Number),
628+
@setup((dIa, dIb, dIx) = _beta_inc_grad(map(float, promote(a, b, x))...)),
629+
(dIa, dIb, dIx),
630+
(-dIa, -dIb, -dIx),
631+
)
632+
# Incomplete beta: beta_inc(a,b,x,y) -> (p, q) with y=1-x, q=1-p
633+
ChainRulesCore.@scalar_rule(
634+
beta_inc(a::Number, b::Number, x::Number, y::Number),
635+
@setup((dIa, dIb, dIx) = _beta_inc_grad(map(float, promote(a, b, x))...)),
636+
(dIa, dIb, dIx, -dIx),
637+
(-dIa, -dIb, -dIx, dIx),
638+
)
689639
# Inverse incomplete beta: beta_inc_inv(a,b,p) -> (x, 1-x)
690-
function ChainRulesCore.frule((_, Δa, Δb, Δp), ::typeof(beta_inc_inv), a::Number, b::Number, p::Number)
691-
x, y = beta_inc_inv(a, b, p)
692-
_a, _b, _x, _p = map(float, promote(a, b, x, p))
693-
# Implicit differentiation at solved x: I_x(a,b) = p
694-
dIa, dIb, _ = _beta_inc_grad(_a, _b, _x)
695-
# ∂I/∂x at solved x via stable log-space expression
696-
dIx_acc = exp(muladd(_a - 1, log(_x), muladd(_b - 1, log1p(-_x), -logbeta(_a, _b))))
697-
inv_dIx = inv(dIx_acc)
698-
dx_da = -dIa * inv_dIx
699-
dx_db = -dIb * inv_dIx
700-
dx_dp = inv_dIx
701-
Δx = muladd(dx_dp, Δp, muladd(dx_db, Δb, dx_da * Δa))
702-
Δy = -Δx
703-
Tout = typeof((x, y))
704-
return (x, y), ChainRulesCore.Tangent{Tout}(Δx, Δy)
705-
end
706-
707-
function ChainRulesCore.rrule(::typeof(beta_inc_inv), a::Number, b::Number, p::Number)
708-
x, y = beta_inc_inv(a, b, p)
709-
Ta = ChainRulesCore.ProjectTo(a)
710-
Tb = ChainRulesCore.ProjectTo(b)
711-
Tp = ChainRulesCore.ProjectTo(p)
712-
_a, _b, _x, _p = map(float, promote(a, b, x, p))
713-
dIa, dIb, _ = _beta_inc_grad(_a, _b, _x)
714-
# ∂I/∂x at solved x via stable log-space expression
715-
dIx_acc = exp(muladd(_a - 1, log(_x), muladd(_b - 1, log1p(-_x), -logbeta(_a, _b))))
716-
inv_dIx = inv(dIx_acc)
717-
dx_da = -dIa * inv_dIx
718-
dx_db = -dIb * inv_dIx
719-
dx_dp = inv_dIx
720-
function beta_inc_inv_pullback(Δ)
721-
Δx, Δy = Δ
722-
s = Δx - Δy
723-
= Ta(s * dx_da)
724-
= Tb(s * dx_db)
725-
= Tp(s * dx_dp)
726-
return ChainRulesCore.NoTangent(), ā, b̄, p̄
727-
end
728-
return (x, y), beta_inc_inv_pullback
729-
end
640+
ChainRulesCore.@scalar_rule(
641+
beta_inc_inv(a::Number, b::Number, p::Number),
642+
@setup(
643+
x = first(Ω), # equivalent to : x, y = beta_inc_inv(map(float, promote(a, b, p))...)
644+
(dIa, dIb, dIx) = _beta_inc_grad(a, b, x),
645+
inv_dIx = inv(dIx),
646+
),
647+
(-dIa * inv_dIx, -dIb * inv_dIx, inv_dIx),
648+
(dIa * inv_dIx, dIb * inv_dIx, -inv_dIx),
649+
)
730650

731651
end # module

0 commit comments

Comments
 (0)