@@ -62,7 +62,7 @@ def test_decoder_init_and_attributes():
6262 assert decoder .contractor_config .contractor_name == "torch"
6363 assert decoder .contractor_config .backend == "torch"
6464 assert decoder .contractor_config .device == "cpu"
65- assert decoder ._dtype == "float32 "
65+ assert decoder ._dtype == "float64 "
6666
6767
6868def test_decoder_replace_logical_observable ():
@@ -585,5 +585,98 @@ def test_invalid_combo():
585585 ContractorConfig ("torch" , "numpy" , "cpu" )
586586
587587
588+ # ---------------------------------------------------------------------------
589+ # Focused tests for _safe_posterior / mass-validation logic, using a mocked
590+ # contractor so these run without a GPU and without a real contraction.
591+ # ContractorConfig is a frozen dataclass; its contractor property reads from
592+ # the class-level _contractors dict, which is the correct mock surface.
593+ # ---------------------------------------------------------------------------
594+
595+
596+ @pytest .fixture
597+ def mock_decoder (monkeypatch ):
598+ """Fixture that returns a factory for decoders with a mocked contractor.
599+
600+ Usage:
601+ decoder = mock_decoder(np.array([m0, m1])) # single decode
602+ decoder = mock_decoder(np.array([[...],...]), rows=N) # batch decode
603+ The factory patches ContractorConfig._contractors for the active backend
604+ and pre-sets path_single/path_batch so no path optimisation is triggered.
605+ """
606+
607+ def _make (mock_output , rows = None ):
608+ H , logical , noise = make_simple_code ()
609+ decoder = qec .get_decoder ("tensor_network_decoder" ,
610+ H ,
611+ logical_obs = logical ,
612+ noise_model = noise )
613+ decoder .path_single = "auto"
614+ decoder .path_batch = "auto"
615+ if rows is not None :
616+ decoder ._batch_size = rows
617+ name = decoder .contractor_config .contractor_name
618+ monkeypatch .setitem (ContractorConfig ._contractors , name ,
619+ lambda * a , ** kw : mock_output )
620+ return decoder
621+
622+ return _make
623+
624+
625+ def test_default_dtype_is_float64 ():
626+ H , logical , noise = make_simple_code ()
627+ decoder = qec .get_decoder ("tensor_network_decoder" ,
628+ H ,
629+ logical_obs = logical ,
630+ noise_model = noise )
631+ assert decoder ._dtype == "float64"
632+
633+
634+ _BIG = np .finfo (np .float64 ).max
635+
636+
637+ @pytest .mark .parametrize (
638+ "masses" ,
639+ [
640+ [- 0.1 , 1.1 ], # negative mass
641+ [0.0 , 0.0 ], # zero denominator
642+ [float ("nan" ), 0.5 ], # NaN
643+ [float ("inf" ), 0.5 ], # infinite mass
644+ [_BIG , _BIG ], # finite masses whose sum overflows to inf
645+ ])
646+ def test_decode_invalid_masses_marks_unconverged (mock_decoder , masses ):
647+ res = mock_decoder (np .array (masses )).decode ([0.0 , 0.0 ])
648+ assert not res .converged
649+ assert res .result [0 ] == pytest .approx (0.5 )
650+
651+
652+ def test_decode_valid_masses_converged (mock_decoder ):
653+ res = mock_decoder (np .array ([0.3 , 0.7 ])).decode ([0.0 , 0.0 ])
654+ assert res .converged
655+ assert res .result [0 ] == pytest .approx (0.7 )
656+
657+
658+ @pytest .mark .parametrize (
659+ "bad_row" ,
660+ [
661+ [- 0.1 , 1.1 ], # negative mass
662+ [float ("nan" ), 0.5 ], # NaN
663+ [0.0 , 0.0 ], # zero denominator
664+ ])
665+ def test_decode_batch_invalid_row_marks_unconverged (mock_decoder , bad_row ):
666+ mock = np .array ([[0.3 , 0.7 ], bad_row ])
667+ res = mock_decoder (mock , rows = 2 ).decode_batch (np .zeros ((2 , 2 )))
668+ assert res .converged [0 ]
669+ assert not res .converged [1 ]
670+ assert res .result [0 , 0 ] == pytest .approx (0.7 )
671+ assert res .result [1 , 0 ] == pytest .approx (0.5 )
672+
673+
674+ def test_decode_batch_all_valid_all_converged (mock_decoder ):
675+ mock = np .array ([[0.2 , 0.8 ], [0.6 , 0.4 ], [0.1 , 0.9 ]])
676+ res = mock_decoder (mock , rows = 3 ).decode_batch (np .zeros ((3 , 2 )))
677+ assert np .all (res .converged )
678+ np .testing .assert_allclose (res .result [:, 0 ], [0.8 , 0.4 , 0.9 ])
679+
680+
588681if __name__ == "__main__" :
589682 pytest .main ()
0 commit comments