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)