From 771c5ba5bb267d8b9ac843abd088d37af519a9d6 Mon Sep 17 00:00:00 2001 From: Hossein Pourbozorg Date: Sun, 30 Aug 2026 16:28:23 +0330 Subject: [PATCH] rewrite `icnf_jacobian` without `Zygote.Buffer` --- src/core/utils.jl | 66 ++++++++++++++++++++++++++--------------------- 1 file changed, 37 insertions(+), 29 deletions(-) diff --git a/src/core/utils.jl b/src/core/utils.jl index 2fc407e9..f25f1c40 100644 --- a/src/core/utils.jl +++ b/src/core/utils.jl @@ -20,14 +20,14 @@ function icnf_jacobian( return y, oftype( cat(y; dims = Val(3)), - cat( + stack( ( J[i:j, i:j] for (i, j) in zip( firstindex(J, 1):size(y, 1):lastindex(J, 1), (firstindex(J, 1) + size(y, 1) - 1):size(y, 1):lastindex(J, 1), ) - )...; - dims = Val(3), + ); + dims = 3, ), ) end @@ -39,18 +39,23 @@ function icnf_jacobian( xs::AbstractMatrix{<:Real}, ) where {T <: AbstractFloat} y = f(xs) - z = similar(xs) - ChainRulesCore.@ignore_derivatives fill!(z, zero(T)) - res = Zygote.Buffer(y, size(xs, 1), size(xs, 1), size(xs, 2)) - for i in axes(xs, 1) - ChainRulesCore.@ignore_derivatives z[i, :] .= one(T) - res[i, :, :] = oftype( - y, - only(DifferentiationInterface.pullback(f, icnf.compute_mode.adback, xs, (z,))), - ) - ChainRulesCore.@ignore_derivatives z[i, :] .= zero(T) - end - return y, oftype(cat(y; dims = Val(3)), copy(res)) + ons = similar(xs, 1, size(xs, 2)) + ChainRulesCore.@ignore_derivatives fill!(ons, one(T)) + return y, + oftype( + cat(y; dims = Val(3)), + stack( + DifferentiationInterface.pullback( + f, + icnf.compute_mode.adback, + xs, + ntuple(function (i::Int) + return oftype(xs, (axes(xs, 1) .== i) * ons) + end, size(xs, 1)), + ); + dims = 1, + ), + ) end function icnf_jacobian( @@ -60,20 +65,23 @@ function icnf_jacobian( xs::AbstractMatrix{<:Real}, ) where {T <: AbstractFloat} y = f(xs) - z = similar(xs) - ChainRulesCore.@ignore_derivatives fill!(z, zero(T)) - res = Zygote.Buffer(y, size(xs, 1), size(xs, 1), size(xs, 2)) - for i in axes(xs, 1) - ChainRulesCore.@ignore_derivatives z[i, :] .= one(T) - res[:, i, :] = oftype( - y, - only( - DifferentiationInterface.pushforward(f, icnf.compute_mode.adback, xs, (z,)), - ), - ) - ChainRulesCore.@ignore_derivatives z[i, :] .= zero(T) - end - return y, oftype(cat(y; dims = Val(3)), copy(res)) + ons = similar(xs, 1, size(xs, 2)) + ChainRulesCore.@ignore_derivatives fill!(ons, one(T)) + return y, + oftype( + cat(y; dims = Val(3)), + stack( + DifferentiationInterface.pushforward( + f, + icnf.compute_mode.adback, + xs, + ntuple(function (i::Int) + return oftype(xs, (axes(xs, 1) .== i) * ons) + end, size(xs, 1)), + ); + dims = 2, + ), + ) end function icnf_jacobian(