@@ -557,6 +557,13 @@ function do_newJW(integrator, alg, nlsolver, repeat_step)::NTuple{2, Bool}
557557 end
558558end
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
641648function 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)
0 commit comments