Skip to content

Commit f7c5c29

Browse files
Preserve residual shape in analytic JVPs
Co-Authored-By: Chris Rackauckas <accounts@chrisrackauckas.com>
1 parent 579761d commit f7c5c29

3 files changed

Lines changed: 22 additions & 2 deletions

File tree

lib/SciMLJacobianOperators/Project.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
name = "SciMLJacobianOperators"
22
uuid = "19f34311-ddf3-4b8b-af20-060888a46c0e"
3-
version = "0.1.15"
3+
version = "0.1.16"
44
authors = ["Avik Pal <avikpal@mit.edu> and contributors"]
55

66
[deps]

lib/SciMLJacobianOperators/src/SciMLJacobianOperators.jl

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -387,7 +387,7 @@ function prepare_jvp(
387387
return
388388
end
389389
else
390-
return @closure (v, u, p) -> reshape(f.jac(u, p) * vec(v), size(u))
390+
return @closure (v, u, p) -> reshape(f.jac(u, p) * vec(v), size(fu))
391391
end
392392
end
393393

lib/SciMLJacobianOperators/test/core_tests__item3.jl

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -104,3 +104,23 @@ prob = NonlinearProblem(
104104
@test JᵀJv JᵀJv_analytic atol = 1.0e-5
105105
end
106106
end
107+
108+
rectangular_residual(u, p) = reshape(
109+
[u[1] + u[2], 2 * u[1] - u[2], u[1] - 3 * u[2]], 3, 1
110+
)
111+
rectangular_jacobian(u, p) = [1 1; 2 -1; 1 -3]
112+
rectangular_u = [2.0, 1.0]
113+
rectangular_fu = rectangular_residual(rectangular_u, nothing)
114+
rectangular_prob = NonlinearLeastSquaresProblem(
115+
NonlinearFunction{false}(rectangular_residual; jac = rectangular_jacobian), rectangular_u
116+
)
117+
118+
@testset "Rectangular Analytic Jacobian" begin
119+
jac_op = JacobianOperator(rectangular_prob, rectangular_fu, rectangular_u)
120+
sop = StatefulJacobianOperator(jac_op, rectangular_u, rectangular_prob.p)
121+
v = [3.0, 2.0]
122+
123+
Jv = sop * v
124+
@test size(Jv) == size(rectangular_fu)
125+
@test Jv reshape(rectangular_jacobian(rectangular_u, nothing) * v, 3, 1)
126+
end

0 commit comments

Comments
 (0)