Skip to content

Commit 08be517

Browse files
singhharsh1708Your Name
authored andcommitted
Fix jacobian2W! DimensionMismatch for ScalarOperator mass matrix
ScalarOperator reports axes(mm) == (), unlike UniformScaling which jacobian2W! already special-cases, so it fell into the boundscheck meant for full matrices. Fixes #3915.
1 parent 0464b30 commit 08be517

4 files changed

Lines changed: 43 additions & 10 deletions

File tree

lib/OrdinaryDiffEqDifferentiation/src/OrdinaryDiffEqDifferentiation.jl

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@ import StaticArraysCore: StaticArray, StaticMatrix
2222
# from the owner to satisfy `all_explicit_imports_via_owners`.
2323
using SciMLBase: UJacobianWrapper, UDerivativeWrapper, _vec, _unwrap_val
2424
import SciMLBase: SciMLBase, @set, DEIntegrator, ODEFunction, SplitFunction, DAEFunction, islinear, remake, solve!, isconstant
25-
import SciMLOperators: SciMLOperators, update_coefficients, update_coefficients!, MatrixOperator, AbstractSciMLOperator
25+
import SciMLOperators: SciMLOperators, update_coefficients, update_coefficients!, MatrixOperator, AbstractSciMLOperator, ScalarOperator
2626
import SparseMatrixColorings: ConstantColoringAlgorithm, GreedyColoringAlgorithm, ColoringProblem,
2727
ncolors, column_colors, coloring, sparsity_pattern
2828
import OrdinaryDiffEqCore

lib/OrdinaryDiffEqDifferentiation/src/derivative_utils.jl

Lines changed: 16 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -557,6 +557,13 @@ function do_newJW(integrator, alg, nlsolver, repeat_step)::NTuple{2, Bool}
557557
end
558558
end
559559

560+
# `ScalarOperator` (λ·I) reports `axes(mm) == ()` like `UniformScaling`, but unlike
561+
# `UniformScaling` it isn't matched by `isa UniformScaling` -- treat both as the same
562+
# scalar-times-identity case rather than requiring `axes(mm) == axes(W)`.
563+
_is_scalar_massmatrix(mm) = mm isa UniformScaling || mm isa ScalarOperator
564+
_scalar_massmatrix_λ(mm::UniformScaling) = mm.λ
565+
_scalar_massmatrix_λ(mm::ScalarOperator) = mm.val
566+
560567
@noinline _throwWJerror(W, J) = throw(DimensionMismatch("W: $(axes(W)), J: $(axes(J))"))
561568
@noinline function _throwWMerror(W, mass_matrix)
562569
throw(DimensionMismatch("W: $(axes(W)), mass matrix: $(axes(mass_matrix))"))
@@ -587,14 +594,14 @@ function jacobian2W!(
587594
# check size and dimension
588595
iijj = axes(W)
589596
@boundscheck (iijj == axes(J) && length(iijj) == 2) || _throwWJerror(W, J)
590-
mass_matrix isa UniformScaling ||
597+
_is_scalar_massmatrix(mass_matrix) ||
591598
@boundscheck axes(mass_matrix) == axes(W) || _throwWMerror(W, mass_matrix)
592599
@inbounds begin
593600
invdtgamma = inv(dtgamma)
594-
if mass_matrix isa UniformScaling
601+
if _is_scalar_massmatrix(mass_matrix)
595602
copyto!(W, J)
596603
idxs = diagind(W)
597-
λ = -mass_matrix.λ
604+
λ = -_scalar_massmatrix_λ(mass_matrix)
598605
if ArrayInterface.fast_scalar_indexing(J) &&
599606
ArrayInterface.fast_scalar_indexing(W)
600607
@inbounds for i in 1:size(J, 1)
@@ -618,14 +625,14 @@ function jacobian2W!(W::Matrix, mass_matrix, dtgamma::Number, J::Matrix)::Nothin
618625
# check size and dimension
619626
iijj = axes(W)
620627
@boundscheck (iijj == axes(J) && length(iijj) == 2) || _throwWJerror(W, J)
621-
mass_matrix isa UniformScaling ||
628+
_is_scalar_massmatrix(mass_matrix) ||
622629
@boundscheck axes(mass_matrix) == axes(W) || _throwWMerror(W, mass_matrix)
623630
@inbounds begin
624631
invdtgamma = inv(dtgamma)
625-
if mass_matrix isa UniformScaling
632+
if _is_scalar_massmatrix(mass_matrix)
626633
copyto!(W, J)
627634
idxs = diagind(W)
628-
λ = -mass_matrix.λ
635+
λ = -_scalar_massmatrix_λ(mass_matrix)
629636
@inbounds for i in 1:size(J, 1)
630637
W[i, i] = muladd(λ, invdtgamma, J[i, i])
631638
end
@@ -640,12 +647,12 @@ end
640647

641648
function jacobian2W(mass_matrix, dtgamma::Number, J::AbstractMatrix)
642649
# check size and dimension
643-
mass_matrix isa UniformScaling ||
650+
_is_scalar_massmatrix(mass_matrix) ||
644651
@boundscheck axes(mass_matrix) == axes(J) || _throwJMerror(J, mass_matrix)
645652
@inbounds begin
646653
invdtgamma = inv(dtgamma)
647-
if mass_matrix isa UniformScaling
648-
λ = -mass_matrix.λ
654+
if _is_scalar_massmatrix(mass_matrix)
655+
λ = -_scalar_massmatrix_λ(mass_matrix)
649656
W = J +* invdtgamma) * I
650657
else
651658
W = muladd(-mass_matrix, invdtgamma, J)

lib/OrdinaryDiffEqDifferentiation/test/runtests.jl

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@ end
3030
# Run functional tests
3131
if TEST_GROUP ("QA", "Sparse", "ModelingToolkit")
3232
@time @safetestset "DAE jacobian2W sparse" include("dae_jacobian2w_sparse_tests.jl")
33+
@time @safetestset "ScalarOperator mass matrix" include("scalar_operator_massmatrix_tests.jl")
3334
@time @safetestset "nzval helpers" include("nzval_helpers_tests.jl")
3435
@time @safetestset "OOP J_t Tracking" include("oop_jt_tracking_test.jl")
3536
@time @safetestset "Differentiation Trait Tests" include("differentiation_traits_tests.jl")
Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,25 @@
1+
using OrdinaryDiffEqDifferentiation
2+
using SciMLOperators
3+
using LinearAlgebra
4+
using Test
5+
6+
# ScalarOperator (λ·I) reports axes(mm) == (), unlike UniformScaling, so it fell
7+
# through the `mass_matrix isa UniformScaling` special case and hit the
8+
# `axes(mass_matrix) == axes(W)` boundscheck meant for full mass matrices.
9+
J = [1.0 2.0; 3.0 4.0]
10+
λ = 2.0
11+
W_expected = J - λ * inv(0.5) * I
12+
13+
W = similar(J)
14+
OrdinaryDiffEqDifferentiation.jacobian2W!(W, ScalarOperator(λ), 0.5, J)
15+
@test W W_expected
16+
17+
W_uniform = similar(J)
18+
OrdinaryDiffEqDifferentiation.jacobian2W!(W_uniform, λ * I, 0.5, J)
19+
@test W W_uniform
20+
21+
@test OrdinaryDiffEqDifferentiation.jacobian2W(ScalarOperator(λ), 0.5, J) W_expected
22+
23+
W_dense = Matrix(J)
24+
OrdinaryDiffEqDifferentiation.jacobian2W!(W_dense, ScalarOperator(λ), 0.5, Matrix(J))
25+
@test W_dense W_expected

0 commit comments

Comments
 (0)