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