@@ -578,6 +578,13 @@ function do_newJW(integrator, alg, nlsolver, repeat_step)::NTuple{2, Bool}
578578 end
579579end
580580
581+ # `ScalarOperator` (λ·I) reports `axes(mm) == ()` like `UniformScaling`, but unlike
582+ # `UniformScaling` it isn't matched by `isa UniformScaling` -- treat both as the same
583+ # scalar-times-identity case rather than requiring `axes(mm) == axes(W)`.
584+ _is_scalar_massmatrix (mm) = mm isa UniformScaling || mm isa ScalarOperator
585+ _scalar_massmatrix_λ (mm:: UniformScaling ) = mm. λ
586+ _scalar_massmatrix_λ (mm:: ScalarOperator ) = mm. val
587+
581588@noinline _throwWJerror (W, J) = throw (DimensionMismatch (" W: $(axes (W)) , J: $(axes (J)) " ))
582589@noinline function _throwWMerror (W, mass_matrix)
583590 throw (DimensionMismatch (" W: $(axes (W)) , mass matrix: $(axes (mass_matrix)) " ))
@@ -608,14 +615,14 @@ function jacobian2W!(
608615 # check size and dimension
609616 iijj = axes (W)
610617 @boundscheck (iijj == axes (J) && length (iijj) == 2 ) || _throwWJerror (W, J)
611- mass_matrix isa UniformScaling ||
618+ _is_scalar_massmatrix ( mass_matrix) ||
612619 @boundscheck axes (mass_matrix) == axes (W) || _throwWMerror (W, mass_matrix)
613620 @inbounds begin
614621 invdtgamma = inv (dtgamma)
615- if mass_matrix isa UniformScaling
622+ if _is_scalar_massmatrix ( mass_matrix)
616623 copyto! (W, J)
617624 idxs = diagind (W)
618- λ = - mass_matrix. λ
625+ λ = - _scalar_massmatrix_λ ( mass_matrix)
619626 if ArrayInterface. fast_scalar_indexing (J) &&
620627 ArrayInterface. fast_scalar_indexing (W)
621628 @inbounds for i in 1 : size (J, 1 )
@@ -639,14 +646,14 @@ function jacobian2W!(W::Matrix, mass_matrix, dtgamma::Number, J::Matrix)::Nothin
639646 # check size and dimension
640647 iijj = axes (W)
641648 @boundscheck (iijj == axes (J) && length (iijj) == 2 ) || _throwWJerror (W, J)
642- mass_matrix isa UniformScaling ||
649+ _is_scalar_massmatrix ( mass_matrix) ||
643650 @boundscheck axes (mass_matrix) == axes (W) || _throwWMerror (W, mass_matrix)
644651 @inbounds begin
645652 invdtgamma = inv (dtgamma)
646- if mass_matrix isa UniformScaling
653+ if _is_scalar_massmatrix ( mass_matrix)
647654 copyto! (W, J)
648655 idxs = diagind (W)
649- λ = - mass_matrix. λ
656+ λ = - _scalar_massmatrix_λ ( mass_matrix)
650657 @inbounds for i in 1 : size (J, 1 )
651658 W[i, i] = muladd (λ, invdtgamma, J[i, i])
652659 end
@@ -661,12 +668,12 @@ end
661668
662669function jacobian2W (mass_matrix, dtgamma:: Number , J:: AbstractMatrix )
663670 # check size and dimension
664- mass_matrix isa UniformScaling ||
671+ _is_scalar_massmatrix ( mass_matrix) ||
665672 @boundscheck axes (mass_matrix) == axes (J) || _throwJMerror (J, mass_matrix)
666673 @inbounds begin
667674 invdtgamma = inv (dtgamma)
668- if mass_matrix isa UniformScaling
669- λ = - mass_matrix. λ
675+ if _is_scalar_massmatrix ( mass_matrix)
676+ λ = - _scalar_massmatrix_λ ( mass_matrix)
670677 W = J + (λ * invdtgamma) * I
671678 else
672679 W = muladd (- mass_matrix, invdtgamma, J)
0 commit comments