Source code for semi_cr.core.quam_components.multi_input_channel
from dataclasses import field
from collections.abc import Sequence
from typing import Any
from qm import qua
from quam.core import quam_dataclass
from quam.components.channels import SingleChannel, MWChannel
from quam.components.ports import (
LFAnalogInputPort,
LFAnalogOutputPort,
OPXPlusAnalogInputPort,
OPXPlusAnalogOutputPort,
LFFEMAnalogInputPort,
LFFEMAnalogOutputPort,
MWFEMAnalogInputPort,
MWFEMAnalogOutputPort,
)
from quam.utils.pulse import add_amplitude_scale_to_pulse_name
AnalogInputSpec = (
LFAnalogInputPort
| OPXPlusAnalogInputPort
| LFFEMAnalogInputPort
| MWFEMAnalogInputPort
| Sequence[Any]
)
AnalogOutputSpec = (
LFAnalogOutputPort
| OPXPlusAnalogOutputPort
| LFFEMAnalogOutputPort
| MWFEMAnalogOutputPort
| Sequence[Any]
)
MWAnalogInputSpec = MWFEMAnalogInputPort | Sequence[Any]
[docs]
@quam_dataclass
class MultiInputChannel(SingleChannel):
opx_inputs: list[AnalogInputSpec] = field(default_factory=list)
opx_input_offsets: list[float | None] | None = None
time_of_flight: int = 140
smearing: int = 0
def _normalize_output_port(
self,
opx_output: AnalogOutputSpec,
config: dict,
):
if isinstance(
opx_output,
(
LFAnalogOutputPort,
OPXPlusAnalogOutputPort,
LFFEMAnalogOutputPort,
MWFEMAnalogOutputPort,
),
):
opx_port = opx_output
elif len(opx_output) == 2:
opx_port = OPXPlusAnalogOutputPort(
*opx_output,
offset=self.opx_output_offset,
)
elif len(opx_output) == 3:
opx_port = LFFEMAnalogOutputPort(
*opx_output,
offset=self.opx_output_offset,
)
else:
raise ValueError(
f"Unsupported OPX output specification: {opx_output!r}"
)
opx_port.apply_to_config(config)
return opx_port
def _normalize_input_port(
self,
opx_input: AnalogInputSpec,
offset: float | None,
config: dict,
):
if isinstance(
opx_input,
(
LFAnalogInputPort,
OPXPlusAnalogInputPort,
LFFEMAnalogInputPort,
MWFEMAnalogInputPort,
),
):
opx_port = opx_input
elif len(opx_input) == 2:
opx_port = OPXPlusAnalogInputPort(
*opx_input,
offset=offset,
)
elif len(opx_input) == 3:
opx_port = LFFEMAnalogInputPort(
*opx_input,
offset=offset,
)
else:
raise ValueError(
f"Unsupported OPX input specification: {opx_input!r}"
)
opx_port.apply_to_config(config)
return opx_port
[docs]
def apply_to_config(self, config: dict):
opx_output = self._normalize_output_port(
opx_output=self.opx_output,
config=config,
)
config["elements"][self.name] = {
"intermediate_frequency": self.intermediate_frequency,
"operations": {},
}
element_config = config["elements"][self.name]
# Output side
if isinstance(opx_output, MWFEMAnalogOutputPort):
element_config["MWInput"] = {
"port": opx_output.port_tuple,
}
else:
element_config["singleInput"] = {
"port": opx_output.port_tuple,
}
# Operations / pulses
for op_name, pulse in self.operations.items():
before = set(config.get("pulses", {}).keys())
pulse.apply_to_config(config)
after = set(config.get("pulses", {}).keys())
new_pulses = list(after - before)
if len(new_pulses) == 1:
pulse_config_name = new_pulses[0]
elif op_name in config.get("pulses", {}):
pulse_config_name = op_name
elif hasattr(pulse, "id") and pulse.id in config.get("pulses", {}):
pulse_config_name = pulse.id
else:
matching = [
name for name in config.get("pulses", {})
if op_name in name
]
if len(matching) == 1:
pulse_config_name = matching[0]
else:
raise KeyError(
f"Could not determine pulse name for operation {op_name!r}. "
f"New pulses: {new_pulses}. "
f"Available pulses: {list(config.get('pulses', {}).keys())}"
)
element_config["operations"][op_name] = pulse_config_name
# Input side
element_config["time_of_flight"] = self.time_of_flight
element_config["smearing"] = self.smearing
element_config["outputs"] = {}
offsets = self.opx_input_offsets or [None] * len(self.opx_inputs)
if len(offsets) != len(self.opx_inputs):
raise ValueError(
"`opx_input_offsets` must have the same length as `opx_inputs`."
)
for k, (opx_input, offset) in enumerate(
zip(self.opx_inputs, offsets),
start=1,
):
opx_port = self._normalize_input_port(
opx_input=opx_input,
offset=offset,
config=config,
)
element_config["outputs"][f"out{k}"] = opx_port.port_tuple
[docs]
def measure(
self,
pulse_name="readout",
outputs=("out1", "out2"),
amplitude_scale=None,
):
pulse = self.operations[pulse_name]
qua_vars = []
demods = []
weights = list(pulse.integration_weights_mapping)
pulse_name_with_amp_scale = add_amplitude_scale_to_pulse_name(
pulse_name, amplitude_scale
)
for output in outputs:
I = qua.declare(qua.fixed)
Q = qua.declare(qua.fixed)
qua_vars.extend([I, Q])
demods.extend([
qua.demod.full(weights[0], I, output),
qua.demod.full(weights[1], Q, output),
])
qua.measure(
pulse_name_with_amp_scale,
self.name,
*demods,
)
return tuple(qua_vars)
[docs]
@quam_dataclass
class MultiInputMWChannel(MWChannel):
opx_inputs: list[MWAnalogInputSpec] = field(default_factory=list)
opx_input_offsets: list[float | None] | None = None
time_of_flight: int = 140
smearing: int = 0
def _normalize_input_port(
self,
opx_input: MWAnalogInputSpec,
offset: float | None,
config: dict,
) -> MWFEMAnalogInputPort:
if isinstance(opx_input, MWFEMAnalogInputPort):
opx_port = opx_input
else:
opx_port = MWFEMAnalogInputPort(
*opx_input,
offset=offset,
)
opx_port.apply_to_config(config)
return opx_port
[docs]
def apply_to_config(self, config: dict):
# Let QUAM handle:
# - MW output
# - I/Q element input structure
# - operations
# - IQ-compatible pulse registration
super().apply_to_config(config)
element_config = config["elements"][self.name]
element_config["time_of_flight"] = self.time_of_flight
element_config["smearing"] = self.smearing
element_config["outputs"] = {}
element_config["MWOutputs"] = {}
offsets = self.opx_input_offsets or [None] * len(self.opx_inputs)
if len(offsets) != len(self.opx_inputs):
raise ValueError(
"`opx_input_offsets` must have the same length as `opx_inputs`."
)
for k, (opx_input, offset) in enumerate(
zip(self.opx_inputs, offsets),
start=1,
):
opx_port = self._normalize_input_port(
opx_input=opx_input,
offset=offset,
config=config,
)
# element_config["outputs"][f"out{k}"] = opx_port.port_tuple
element_config["MWOutputs"][f"out{k}"] = {
"port": opx_port.port_tuple
}
[docs]
def measure(
self,
pulse_name: str = "readout",
outputs: tuple[str, ...] = ("out1", "out2"),
amplitude_scale: float | None = None,
):
pulse = self.operations[pulse_name]
qua_vars = []
demods = []
weights = list(pulse.integration_weights_mapping)
pulse_name_with_amp_scale = add_amplitude_scale_to_pulse_name(
pulse_name,
amplitude_scale,
)
for output in outputs:
I = qua.declare(qua.fixed)
Q = qua.declare(qua.fixed)
qua_vars.extend([I, Q])
demods.extend(
[
qua.demod.full(weights[0], I, output),
qua.demod.full(weights[1], Q, output),
]
)
qua.measure(
pulse_name_with_amp_scale,
self.name,
None,
*demods,
)
return tuple(qua_vars)