Skip to content

Commit 6568512

Browse files
Fix TensorNetworkDecoder float32 default causing non-physical posteriors
Signed-off-by: vedika-saravanan <vsaravanan@nvidia.com>
1 parent 291aa7f commit 6568512

3 files changed

Lines changed: 123 additions & 12 deletions

File tree

docs/sphinx/api/qec/tensor_network_decoder_api.rst

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -68,7 +68,7 @@
6868
:param logical_inds: (optional) List of logical index names
6969
:param logical_tags: (optional) List of logical tags
7070
:param contract_noise_model: (bool, optional) Whether to contract the noise model at initialization (default: True)
71-
:param dtype: (str, optional) Data type for tensors (default: "float32")
71+
:param dtype: (str, optional) Data type for tensors (default: "float64")
7272
:param device: (str, optional) Device for tensor operations ("cpu", "cuda", or "cuda:X", default: "cuda")
7373

7474
**Methods**

libs/qec/python/cudaq_qec/plugins/decoders/tensor_network_decoder.py

Lines changed: 28 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,23 @@
2222
tensor_network_from_syndrome_batch, tensor_network_from_logical_observable)
2323

2424

25+
def _safe_posterior(mass0, mass1) -> tuple:
26+
"""Validate contraction masses and compute the posterior P(logical flip).
27+
28+
Returns (posterior, is_valid). On any non-physical input (negative, NaN,
29+
infinite, or zero-sum masses) returns (0.5, False) so callers can mark the
30+
result unconverged rather than silently propagating a bad value.
31+
"""
32+
m0, m1 = float(mass0), float(mass1)
33+
if (np.isnan(m0) or np.isnan(m1) or np.isinf(m0) or np.isinf(m1) or
34+
m0 < 0.0 or m1 < 0.0):
35+
return 0.5, False
36+
denom = m0 + m1
37+
if not np.isfinite(denom) or denom == 0.0:
38+
return 0.5, False
39+
return m1 / denom, True
40+
41+
2542
@qec.decoder("tensor_network_decoder")
2643
class TensorNetworkDecoder:
2744
r"""A general class for tensor network decoders.
@@ -110,7 +127,7 @@ def __init__(
110127
logical_inds: list[str] | None = None,
111128
logical_tags: list[str] | None = None,
112129
contract_noise_model: bool = True,
113-
dtype: str = "float32",
130+
dtype: str = "float64",
114131
device: str = "cuda",
115132
) -> None:
116133
"""Initialize a sparse representation of a tensor network decoder for an arbitrary code
@@ -406,12 +423,11 @@ def decode(
406423
device_id=self.contractor_config.device_id,
407424
)
408425

426+
posterior, valid = _safe_posterior(contraction_value[0],
427+
contraction_value[1])
409428
res = qec.DecoderResult()
410-
res.converged = True
411-
res.result = [
412-
float(contraction_value[1] /
413-
(contraction_value[1] + contraction_value[0]))
414-
]
429+
res.converged = valid
430+
res.result = [posterior]
415431
return res
416432

417433
def decode_batch(
@@ -463,18 +479,20 @@ def decode_batch(
463479
)
464480

465481
probabilities = []
482+
converged = []
466483
for r in range(syndrome_batch.shape[0]):
467-
probabilities.append(
468-
float(contraction_value[r, 1] /
469-
(contraction_value[r, 1] + contraction_value[r, 0])))
484+
posterior, valid = _safe_posterior(contraction_value[r, 0],
485+
contraction_value[r, 1])
486+
probabilities.append(posterior)
487+
converged.append(valid)
470488

471489
# Python `decode_batch` override: construct a BatchDecoderResult
472490
# directly, bypassing the native decoder aggregation path. This is
473491
# the only sanctioned caller of the BatchDecoderResult constructor;
474492
# see its docstring for the supported construction surface.
475493
return qec.BatchDecoderResult(
476494
np.asarray(probabilities, dtype=np.float64).reshape((-1, 1)),
477-
np.ones(syndrome_batch.shape[0], dtype=bool),
495+
np.asarray(converged, dtype=bool),
478496
)
479497

480498
def optimize_path(

libs/qec/python/tests/test_tensor_network_decoder.py

Lines changed: 94 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -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

6868
def 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+
588681
if __name__ == "__main__":
589682
pytest.main()

0 commit comments

Comments
 (0)