diff --git a/qiskit_ibm_runtime/executor_estimator/estimator.py b/qiskit_ibm_runtime/executor_estimator/estimator.py index 05e6949ac..194a5b390 100644 --- a/qiskit_ibm_runtime/executor_estimator/estimator.py +++ b/qiskit_ibm_runtime/executor_estimator/estimator.py @@ -15,6 +15,7 @@ from __future__ import annotations import logging +import warnings from copy import deepcopy from typing import TYPE_CHECKING, Any @@ -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, @@ -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 diff --git a/test/unit/executor_estimator/test_estimator_v2.py b/test/unit/executor_estimator/test_estimator_v2.py index 12f25b729..efa1334dd 100644 --- a/test/unit/executor_estimator/test_estimator_v2.py +++ b/test/unit/executor_estimator/test_estimator_v2.py @@ -12,6 +12,7 @@ """Unit tests for EstimatorV2 run method.""" +import warnings from unittest.mock import MagicMock, patch import numpy as np @@ -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))