@@ -622,110 +622,30 @@ function _beta_inc_grad(a::T, b::T, x::T) where {T<:Union{Float64, Float32}}
622622 end
623623end
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- b̄ = Tb (s * dIb)
654- x̄ = 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- b̄ = Tb (s * dIb)
682- x̄ = 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- b̄ = Tb (s * dx_db)
725- p̄ = 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
731651end # module
0 commit comments