@@ -414,6 +414,34 @@ def run(y0__args, adjoint):
414414 assert tree_allclose (grads1 , grads3 , rtol = 1e-3 , atol = 1e-3 )
415415
416416
417+ @pytest .mark .parametrize (
418+ "solver" ,
419+ (diffrax .ShARK (), diffrax .SEA (), diffrax .SRA1 (), diffrax .SlowRK ()),
420+ )
421+ def test_backsolve_multiterm_solver_error (solver , getkey ):
422+ # https://github.com/patrick-kidger/diffrax/issues/558
423+ t0 , t1 , dt0 = 0 , 1 , 0.01
424+ bm = diffrax .VirtualBrownianTree (
425+ t0 , t1 , 1e-3 , (2 ,), key = getkey (), levy_area = diffrax .SpaceTimeLevyArea
426+ )
427+ drift = diffrax .ODETerm (lambda t , y , args : - y )
428+ diffusion = diffrax .ControlTerm (
429+ lambda t , y , args : lx .DiagonalLinearOperator (0.1 * jnp .zeros_like (y )), bm
430+ )
431+ terms = diffrax .MultiTerm (drift , diffusion )
432+
433+ @eqx .filter_jit
434+ @jax .grad
435+ def run (y0 ):
436+ sol = diffrax .diffeqsolve (
437+ terms , solver , t0 , t1 , dt0 , y0 , adjoint = diffrax .BacksolveAdjoint ()
438+ )
439+ return jnp .sum (cast (Array , sol .ys ))
440+
441+ with pytest .raises (NotImplementedError , match = "not compatible with solver" ):
442+ run (jnp .array ([1.0 , 2.0 ]))
443+
444+
417445def test_implicit_runge_kutta_direct_adjoint ():
418446 diffrax .diffeqsolve (
419447 diffrax .ODETerm (lambda t , y , args : - y ),
0 commit comments