Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
26 changes: 19 additions & 7 deletions qiskit_ibm_runtime/sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
import copy

from qiskit.circuit import QuantumCircuit, Parameter
from qiskit.primitives.utils import final_measurement_mapping

# TODO import BaseSampler and SamplerResult from terra once released
from .qiskit.primitives import BaseSampler, SamplerResult
Expand Down Expand Up @@ -180,14 +181,25 @@ def run(
Raises:
ValueError: If the input values are invalid.
"""
if isinstance(circuits, Iterable) and not all(
isinstance(inst, QuantumCircuit) for inst in circuits
):
raise ValueError(
"The circuits parameter has to be instances of QuantumCircuit."
)
if isinstance(circuits, QuantumCircuit):
circuits = [circuits]

for circ in circuits:
if not isinstance(circ, QuantumCircuit):
raise ValueError(
"The circuits parameter has to be instances of QuantumCircuit."
)
Comment thread
t-imamichi marked this conversation as resolved.
Outdated
if circ.num_clbits == 0:
raise ValueError("The circuits should have at least one classical bit.")
if self.options.resilience_level >= 1 and circ.num_clbits > len(
Comment thread
t-imamichi marked this conversation as resolved.
Outdated
final_measurement_mapping(circ)
):
raise ValueError(
"All classical bits should be mapped to qubits "
"in the final measurement to apply readout error mitigation."
)

circ_count = 1 if isinstance(circuits, QuantumCircuit) else len(circuits)
circ_count = len(circuits)

inputs = {
"circuits": circuits,
Expand Down
21 changes: 21 additions & 0 deletions test/integration/test_sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,3 +166,24 @@ def test_sampler_primitive_as_session(self, service):
self.assertIsInstance(result, SamplerResult)
self.assertEqual(len(result.quasi_dists), len(circuits0))
self.assertEqual(len(result.metadata), len(circuits0))

@run_integration_test
def test_sampler_raise_error(self, service):
Comment thread
kt474 marked this conversation as resolved.
Outdated
"""Test error check properly works."""
with Session(service, self.backend) as session:
sampler = Sampler(session=session)

with self.assertRaises(ValueError):
_ = sampler.run(123)
with self.assertRaises(ValueError):
_ = sampler.run([123])

circuit = QuantumCircuit(2)
with self.assertRaises(ValueError):
_ = sampler.run(circuit)

sampler.options.resilience_level = 1
circuit = QuantumCircuit(1, 2)
circuit.measure(0, 0)
with self.assertRaises(ValueError):
_ = sampler.run(circuit)