Skip to content

Commit 25848db

Browse files
authored
Merge branch 'main' into melodyr/single-source-of-truth-for-hw-pin
2 parents e67aedd + 80e46f3 commit 25848db

3 files changed

Lines changed: 145 additions & 14 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: 38 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,32 @@
2222
tensor_network_from_syndrome_batch, tensor_network_from_logical_observable)
2323

2424

25+
def _safe_posterior(mass0, mass1, dtype: str) -> tuple[float, bool]:
26+
"""Validate contraction masses and compute the posterior P(logical flip).
27+
28+
Tiny negative values within floating-point round-off of zero are normalized
29+
to zero. On any materially non-physical input (negative, NaN, infinite, or
30+
zero-sum masses), return ``(0.5, False)`` so callers can mark the result
31+
unconverged rather than silently propagating a bad value.
32+
"""
33+
m0, m1 = float(mass0), float(mass1)
34+
if not np.isfinite(m0) or not np.isfinite(m1):
35+
return 0.5, False
36+
37+
scale = max(abs(m0), abs(m1))
38+
tolerance = 64.0 * np.finfo(np.dtype(dtype)).eps * scale
39+
if m0 < -tolerance or m1 < -tolerance:
40+
return 0.5, False
41+
42+
# Signed contractions can leave an exact zero a few ULPs below zero.
43+
m0 = max(m0, 0.0)
44+
m1 = max(m1, 0.0)
45+
denom = m0 + m1
46+
if not np.isfinite(denom) or denom == 0.0:
47+
return 0.5, False
48+
return m1 / denom, True
49+
50+
2551
@qec.decoder("tensor_network_decoder")
2652
class TensorNetworkDecoder:
2753
r"""A general class for tensor network decoders.
@@ -110,7 +136,7 @@ def __init__(
110136
logical_inds: list[str] | None = None,
111137
logical_tags: list[str] | None = None,
112138
contract_noise_model: bool = True,
113-
dtype: str = "float32",
139+
dtype: str = "float64",
114140
device: str = "cuda",
115141
) -> None:
116142
"""Initialize a sparse representation of a tensor network decoder for an arbitrary code
@@ -406,12 +432,11 @@ def decode(
406432
device_id=self.contractor_config.device_id,
407433
)
408434

435+
posterior, valid = _safe_posterior(contraction_value[0],
436+
contraction_value[1], self._dtype)
409437
res = qec.DecoderResult()
410-
res.converged = True
411-
res.result = [
412-
float(contraction_value[1] /
413-
(contraction_value[1] + contraction_value[0]))
414-
]
438+
res.converged = valid
439+
res.result = [posterior]
415440
return res
416441

417442
def decode_batch(
@@ -463,18 +488,21 @@ def decode_batch(
463488
)
464489

465490
probabilities = []
491+
converged = []
466492
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])))
493+
posterior, valid = _safe_posterior(contraction_value[r, 0],
494+
contraction_value[r, 1],
495+
self._dtype)
496+
probabilities.append(posterior)
497+
converged.append(valid)
470498

471499
# Python `decode_batch` override: construct a BatchDecoderResult
472500
# directly, bypassing the native decoder aggregation path. This is
473501
# the only sanctioned caller of the BatchDecoderResult constructor;
474502
# see its docstring for the supported construction surface.
475503
return qec.BatchDecoderResult(
476504
np.asarray(probabilities, dtype=np.float64).reshape((-1, 1)),
477-
np.ones(syndrome_batch.shape[0], dtype=bool),
505+
np.asarray(converged, dtype=bool),
478506
)
479507

480508
def optimize_path(

libs/qec/python/tests/test_tensor_network_decoder.py

Lines changed: 106 additions & 3 deletions
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():
@@ -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

165166
def 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+
588691
if __name__ == "__main__":
589692
pytest.main()

0 commit comments

Comments
 (0)