Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
70 changes: 64 additions & 6 deletions qiskit_ibm_runtime/executor_estimator/estimator.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
from __future__ import annotations

import logging
import warnings
from copy import deepcopy
from typing import TYPE_CHECKING, Any

Expand Down Expand Up @@ -43,6 +44,33 @@

logger = logging.getLogger(__name__)


def _force_twirling_field(
field: str,
twirling_options: Any,
user_twirling_fields: set[str],
mitigation: str,
) -> None:
"""Force a twirling field to ``True``, warning if the user had explicitly set it to ``False``.

Args:
field: The twirling field name (``"enable_gates"`` or ``"enable_measure"``).
twirling_options: The ``TwirlingOptions`` instance to update in place.
user_twirling_fields: The set of field names explicitly set by the user on the twirling
options (i.e. ``self.options.twirling.model_fields_set``).
mitigation: A short description of the mitigation technique forcing the override, used in
the warning message (e.g. ``"measurement mitigation"``).
"""
if field in user_twirling_fields and getattr(twirling_options, field) is False:
warnings.warn(
f"twirling.{field}=False was explicitly set, but {mitigation} requires "
f"twirling.{field}=True. The value is being overridden to True.",
UserWarning,
stacklevel=3,
)
setattr(twirling_options, field, True)


RESILIENCE_LEVEL_DEFAULTS = {
0: {
"enable_gates": False,
Expand Down Expand Up @@ -234,18 +262,48 @@ def finalize_options(self) -> EstimatorOptions:
if finalized_options.resilience.zne_mitigation is None:
finalized_options.resilience.zne_mitigation = defaults["zne_mitigation"]

# Force-set some values based on mitigation
# Force-set some values based on mitigation, warning when a user-specified twirling
# field is being overridden.
user_twirling_fields = self.options.twirling.model_fields_set

if finalized_options.resilience.measure_mitigation is True:
finalized_options.twirling.enable_measure = True
_force_twirling_field(
"enable_measure",
finalized_options.twirling,
user_twirling_fields,
"measurement mitigation",
)

if (
finalized_options.resilience.zne_mitigation is True
and finalized_options.resilience.zne.amplifier == "pea"
):
finalized_options.twirling.enable_gates = True
finalized_options.twirling.enable_measure = True
_force_twirling_field(
"enable_gates",
finalized_options.twirling,
user_twirling_fields,
"PEA mitigation",
)
_force_twirling_field(
"enable_measure",
finalized_options.twirling,
user_twirling_fields,
"PEA mitigation",
)

if finalized_options.resilience.pec_mitigation is True:
finalized_options.twirling.enable_gates = True
finalized_options.twirling.enable_measure = True
_force_twirling_field(
"enable_gates",
finalized_options.twirling,
user_twirling_fields,
"PEC mitigation",
)
_force_twirling_field(
"enable_measure",
finalized_options.twirling,
user_twirling_fields,
"PEC mitigation",
)

return finalized_options

Expand Down
64 changes: 64 additions & 0 deletions test/unit/executor_estimator/test_estimator_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@

"""Unit tests for EstimatorV2 run method."""

import warnings
from unittest.mock import MagicMock, patch

import numpy as np
Expand Down Expand Up @@ -572,3 +573,66 @@ def test_forced_values(self, resilience_level):
finalized_options = estimator.finalize_options()
self.assertTrue(finalized_options.twirling.enable_gates)
self.assertTrue(finalized_options.twirling.enable_measure)

def test_no_warning_when_twirling_field_not_set_by_user(self):
"""No warning when the user never set the twirling field that is being overridden."""
estimator = EstimatorV2(self.backend)
# Use resilience_level=0 so enable_measure defaults to False, ensuring the only
# thing suppressing the warning is the field being absent from model_fields_set.
estimator.options.resilience_level = 0
estimator.options.resilience.measure_mitigation = True
# enable_measure was not explicitly set by the user → no warning expected
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always")
estimator.finalize_options()
user_warns = [w for w in caught if issubclass(w.category, UserWarning)]
self.assertEqual(user_warns, [])

def test_no_warning_when_user_set_field_to_true(self):
"""No warning when the user already set the field to True (no conflict)."""
estimator = EstimatorV2(self.backend)
estimator.options.twirling.enable_measure = True
estimator.options.resilience.measure_mitigation = True
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always")
estimator.finalize_options()
user_warns = [w for w in caught if issubclass(w.category, UserWarning)]
self.assertEqual(user_warns, [])

def test_warning_measure_mitigation_overrides_enable_measure_false(self):
"""Warning when measure_mitigation=True overrides user-set enable_measure=False."""
estimator = EstimatorV2(self.backend)
estimator.options.twirling.enable_measure = False
estimator.options.resilience.measure_mitigation = True
with self.assertWarns(UserWarning) as ctx:
estimator.finalize_options()
msg = str(ctx.warning)
self.assertIn("enable_measure", msg)
self.assertIn("measurement mitigation", msg)

@data("enable_gates", "enable_measure")
def test_warning_pea_overrides_twirling_field_false(self, field):
"""Warning when PEA overrides user-set enable_gates=False or enable_measure=False."""
estimator = EstimatorV2(self.backend)
setattr(estimator.options.twirling, field, False)
estimator.options.resilience.zne_mitigation = True
estimator.options.resilience.measure_mitigation = False
estimator.options.resilience.zne.amplifier = "pea"
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always")
estimator.finalize_options()
msgs = [str(w.message) for w in caught if issubclass(w.category, UserWarning)]
self.assertTrue(any(field in m and "PEA mitigation" in m for m in msgs))

@data("enable_gates", "enable_measure")
def test_warning_pec_overrides_twirling_field_false(self, field):
"""Warning when PEC overrides user-set enable_gates=False or enable_measure=False."""
estimator = EstimatorV2(self.backend)
setattr(estimator.options.twirling, field, False)
estimator.options.resilience.pec_mitigation = True
estimator.options.resilience.measure_mitigation = False
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always")
estimator.finalize_options()
msgs = [str(w.message) for w in caught if issubclass(w.category, UserWarning)]
self.assertTrue(any(field in m and "PEC mitigation" in m for m in msgs))
Loading