# %%
from qm import qua
from qualang_tools.loops import from_array
from qualang_tools.units import unit
import numpy as np
from semi_cr.core.lab.devices.multiplexer import Multiplexer
from semi_cr.core.lab.protocols.inputs.multiplexer import MultiplexerProtocolInput
from semi_cr.core.lab.measurements.axis_factories.multiplexer import MULTIPLEXER_AXIS_FACTORIES
from semi_cr.core.lab.utils.qua import qua_cycles_to_time
from semi_cr.core.quam_components.helpers import get_trigger_component, get_mux_components
u = unit(coerce_to_integer=True)
# %%
[docs]
def measure_multiplexer(
device,
I,
Q,
I_st,
Q_st,
amplitude=1.0,
):
channel_index = 0
mux_components = get_mux_components(device)
for component_index, mux_component in enumerate(mux_components):
amplitude_scale = amplitude
# Optional: only the first component drives
if component_index > 0:
amplitude_scale = 0
if hasattr(mux_component, "opx_inputs"):
# LF MultiInputChannel / custom MultiInputMWChannel-style API
outputs = tuple(
f"out{k}"
for k in range(1, len(mux_component.opx_inputs) + 1)
)
measured = mux_component.measure(
"readout",
outputs=outputs,
amplitude_scale=amplitude_scale,
)
else:
# Standard InOutMWChannel API
measured = mux_component.measure(
"readout",
amplitude_scale=amplitude_scale,
)
for I_val, Q_val in zip(measured[0::2], measured[1::2]):
I[channel_index] = I_val
Q[channel_index] = Q_val
channel_index += 1
if channel_index != len(I):
raise ValueError(
f"Measured {channel_index} IQ channels, but expected {len(I)}."
)
for i in range(len(I)):
qua.save(I[i], I_st[i])
qua.save(Q[i], Q_st[i])
# %%
[docs]
def is_enabled(config) -> bool:
return getattr(config, "enabled", "OFF") == "ON"
[docs]
def get_enabled_multiplexer_axes(
protocol_inputs: MultiplexerProtocolInput,
) -> tuple[str, ...]:
enabled_axes = list(
protocol_inputs.parameters.enabled_swept()
)
if not ({"frequency", "mw_frequency"} & set(enabled_axes)):
raise ValueError("frequency or mw_frequency must be enabled for this QUA program.")
return tuple(enabled_axes)
[docs]
def get_sweep_config(protocol_inputs, axis_name):
parameter_sweeps = protocol_inputs.parameters.swept
axis_config = parameter_sweeps[axis_name]
source_name = getattr(axis_config, "source", None)
if source_name is None:
return axis_config
if not hasattr(protocol_inputs, source_name):
raise ValueError(
f"Axis {axis_name!r} references unknown source {source_name!r}"
)
return getattr(protocol_inputs, source_name)
[docs]
def make_multiplexer_sweep_axes(
device: Multiplexer,
protocol_inputs,
):
enabled_axes = get_enabled_multiplexer_axes(
protocol_inputs=protocol_inputs,
)
axes = []
for axis_name in enabled_axes:
parameter_sweeps = protocol_inputs.parameters.swept
axis_config = parameter_sweeps[axis_name]
sweep_config = get_sweep_config(protocol_inputs, axis_name)
factory = MULTIPLEXER_AXIS_FACTORIES[axis_name]
axes.append(
factory(
axis_config=axis_config,
sweep_config=sweep_config,
protocol_inputs=protocol_inputs,
device=device,
)
)
return axes
[docs]
def trigger_qdac_switch(trigger, mux_components):
trigger.play("const")
for mux_component in mux_components:
qua.wait_for_trigger(mux_component.id)
if len(mux_components) > 1:
qua.align(*(component.id for component in mux_components))
[docs]
def nested_sweep(
axes,
measurement_fn,
axis_index=0,
):
if axis_index == len(axes):
measurement_fn()
return
axis = axes[axis_index]
loop_values = axis.values if axis.qua_values is None else axis.qua_values
print(
"QUA axis:",
axis.name,
"values:",
axis.values,
"qua_values:",
getattr(axis, "qua_values", None),
"apply:",
axis.apply,
)
if axis.use_for_each:
with qua.for_each_(axis.qua_var, loop_values.tolist()):
if axis.apply is not None:
axis.apply(axis.qua_var)
nested_sweep(
axes=axes,
measurement_fn=measurement_fn,
axis_index=axis_index + 1,
)
else:
with qua.for_(*from_array(axis.qua_var, loop_values)):
if axis.apply is not None:
axis.apply(axis.qua_var)
nested_sweep(
axes=axes,
measurement_fn=measurement_fn,
axis_index=axis_index + 1,
)
# %%
[docs]
def make_buffer_plan(
output_axes,
n_avg,
inner_axis_names=("frequency", "amplitude"),
):
outer_axes = [
axis for axis in output_axes
if axis.name not in inner_axis_names
]
inner_axes = [
axis for axis in output_axes
if axis.name in inner_axis_names
]
return {
"inner_axes": inner_axes,
"outer_axes": outer_axes,
"n_avg": n_avg,
}
[docs]
def apply_buffers(
stream,
output_axes,
n_avg,
inner_axis_names=["frequency", "amplitude"],
):
plan = make_buffer_plan(
output_axes,
n_avg,
inner_axis_names,
)
for axis in reversed(plan["inner_axes"]):
stream = stream.buffer(len(axis.values))
stream = stream.buffer(plan["n_avg"])
stream = stream.map(qua.FUNCTIONS.average(0))
for axis in reversed(plan["outer_axes"]):
stream = stream.buffer(len(axis.values))
# outer_axes = [
# axis for axis in output_axes
# if axis.name not in inner_axis_names
# ]
# inner_axes = [
# axis for axis in output_axes
# if axis.name in inner_axis_names
# ]
# # The following commentted out part needs to be tested
# # -------------------------------------------------
# # 1. Buffer inner sweep dimensions.
# #
# # Execution order:
# #
# # n_avg
# # frequency
# # amplitude
# #
# # Therefore stream arrival order is:
# #
# # amplitude -> frequency -> n_avg
# # -------------------------------------------------
# for axis in reversed(inner_axes):
# print(
# f"buffer inner: {axis.name}, "
# f"size={len(axis.values)}"
# )
# stream = stream.buffer(len(axis.values))
# # -------------------------------------------------
# # 2. Buffer averages.
# #
# # This creates:
# #
# # [n_avg, frequency]
# #
# # or, with amplitude:
# #
# # [n_avg, frequency, amplitude]
# # -------------------------------------------------
# print(f"buffer average: n_avg, size={n_avg}")
# stream = stream.buffer(n_avg)
# # -------------------------------------------------
# # 3. Average the n_avg dimension immediately.
# #
# # n_avg is now axis 0 of this buffered object.
# # -------------------------------------------------
# print("average axis: 0")
# stream = stream.map(
# qua.FUNCTIONS.average(0)
# )
# # -------------------------------------------------
# # 4. Buffer outer sweep dimensions after averaging.
# # -------------------------------------------------
# for axis in reversed(outer_axes):
# print(
# f"buffer outer: {axis.name}, "
# f"size={len(axis.values)}"
# )
# stream = stream.buffer(len(axis.values))
# # -------------------------------------------------
# # 1. Buffer inner sweep dimensions.
# #
# # Execution order:
# #
# # n_avg
# # frequency
# # amplitude
# #
# # Therefore stream arrival order is:
# #
# # amplitude -> frequency -> n_avg
# # -------------------------------------------------
# for axis in reversed(inner_axes):
# stream = stream.buffer(len(axis.values))
# # -------------------------------------------------
# # 2. Buffer averages.
# #
# # This creates:
# #
# # [n_avg, frequency]
# #
# # or, with amplitude:
# #
# # [n_avg, frequency, amplitude]
# # -------------------------------------------------
# stream = stream.buffer(n_avg)
# # -------------------------------------------------
# # 3. Average the n_avg dimension immediately.
# #
# # n_avg is now axis 0 of this buffered object.
# # -------------------------------------------------
# stream = stream.map(
# qua.FUNCTIONS.average(0)
# )
# # -------------------------------------------------
# # 4. Buffer outer sweep dimensions after averaging.
# # -------------------------------------------------
# for axis in reversed(outer_axes):
# stream = stream.buffer(len(axis.values))
# logical_dims = [
# *outer_axes,
# "n_avg",
# *inner_axes,
# ]
# buffer_dims = list(reversed(logical_dims))
# print("logical_dims:", [
# dim if dim == "n_avg" else dim.name
# for dim in logical_dims
# ])
# print("buffer_dims:", [
# dim if dim == "n_avg" else dim.name
# for dim in buffer_dims
# ])
# print("buffer_sizes:", [
# n_avg if dim == "n_avg" else len(dim.values)
# for dim in buffer_dims
# ])
# average_axis = logical_dims.index("n_avg")
# print("average_axis:", average_axis)
# for dim in buffer_dims:
# size = n_avg if dim == "n_avg" else len(dim.values)
# stream = stream.buffer(size)
# stream = stream.map(qua.FUNCTIONS.average(average_axis))
return stream
# %%
[docs]
def debug_buffer_plan(
output_axes,
n_avg,
inner_axis_names=("frequency", "amplitude"),
):
plan = make_buffer_plan(
output_axes,
n_avg,
inner_axis_names,
)
# outer_axes = [
# axis for axis in output_axes
# if axis.name not in inner_axis_names
# ]
# inner_axes = [
# axis for axis in output_axes
# if axis.name in inner_axis_names
# ]
# logical_dims = [
# *outer_axes,
# "n_avg",
# *inner_axes,
# ]
# buffer_dims = list(reversed(logical_dims))
print("output_axes:", [axis.name for axis in output_axes])
# print("outer_axes:", [axis.name for axis in outer_axes])
# print("inner_axes:", [axis.name for axis in inner_axes])
print(
"outer_axes:",
[axis.name for axis in plan["outer_axes"]],
)
print(
"inner_axes:",
[axis.name for axis in plan["inner_axes"]],
)
# print("logical_dims:", [
# dim if dim == "n_avg" else dim.name
# for dim in logical_dims
# ])
# print("buffer_dims:", [
# dim if dim == "n_avg" else dim.name
# for dim in buffer_dims
# ])
# print("buffer_sizes:", [
# n_avg if dim == "n_avg" else len(dim.values)
# for dim in buffer_dims
# ])
# print("average_axis:", logical_dims.index("n_avg"))
# Inner axes are buffered first.
for axis in reversed(plan["inner_axes"]):
print(
f"buffer inner: {axis.name}, "
f"size={len(axis.values)}"
)
# Then collect complete inner sweeps for averaging.
print(f" buffer average: n_avg, size={plan['n_avg']}")
print(" average axis: 0")
# Outer axes are buffered only after averaging.
for axis in reversed(plan["outer_axes"]):
print(
f" buffer outer: {axis.name}, "
f"size={len(axis.values)}"
)
# %%
[docs]
def get_multiplexer_sweep_program(
device,
n_iq_channels,
qua_sweep_axes,
n_avg,
settling_cycles: int,
settling_execution,
):
with qua.program() as qua_prog:
I = [qua.declare(qua.fixed) for _ in range(n_iq_channels)]
Q = [qua.declare(qua.fixed) for _ in range(n_iq_channels)]
I_st = [qua.declare_stream() for _ in range(n_iq_channels)]
Q_st = [qua.declare_stream() for _ in range(n_iq_channels)]
n = qua.declare(int)
for axis in qua_sweep_axes:
axis.qua_var = qua.declare(axis.qua_type)
settling_time_axis = next(
(
axis
for axis in qua_sweep_axes
if axis.name == "settling_time"
),
None,
)
if settling_execution.mode == "qua":
if settling_time_axis is None:
raise RuntimeError(
"Settling time is configured for QUA execution, "
"but no settling_time QUASweepAxis is present."
)
else:
if settling_cycles is None:
raise RuntimeError(
"Fixed/external settling time requires "
"a scalar settling_cycles value."
)
# if settling_time_axis is not None:
# settling_cycles_expr = settling_time_axis.qua_var
# else:
# settling_cycles_expr = settling_cycles
frequency_axis = next(
axis for axis in qua_sweep_axes
if axis.name == "frequency"
)
amplitude_axis = next(
(axis for axis in qua_sweep_axes
if axis.name == "amplitude"),
None,
)
amplitude_var = 1.0 if amplitude_axis is None else amplitude_axis.qua_var
outer_axes = [
axis for axis in qua_sweep_axes
if axis.name not in {"frequency", "amplitude"}
]
inner_axes = [frequency_axis]
if amplitude_axis is not None:
inner_axes.append(amplitude_axis)
qua_loop_axes = outer_axes + inner_axes
trigger = get_trigger_component(device)
mux_components = get_mux_components(device)
def measure_at_inner_point():
# # worked for LF FEM with amplitude sweep, to check again!
# frequency_axis.apply(frequency_axis.qua_var)
measure_multiplexer(
device,
I=I,
Q=Q,
I_st=I_st,
Q_st=Q_st,
amplitude=amplitude_var,
)
def sweep_inner_axes():
nested_sweep(
axes=inner_axes,
measurement_fn=measure_at_inner_point,
)
def measure_at_outer_point():
trigger_qdac_switch(
trigger=trigger,
mux_components=mux_components,
)
component_ids = [
component.id
for component in mux_components
]
# settling_time = qua_cycles_to_time(
# settling_cycles,
# unit="ms",
# )
# if settling_time > 0:
# qua.wait(
# settling_cycles,
# *[component.id for component in mux_components],
# )
if settling_execution.mode == "qua":
# Internal QUA sweep:
qua.wait(
settling_time_axis.qua_var,
# *[component.id for component in mux_components],
*component_ids,
)
elif settling_cycles > 0:
# External/QCoDeS settling-time sweep:
qua.wait(
settling_cycles,
# *[component.id for component in mux_components],
*component_ids,
)
with qua.for_(n, 0, n < n_avg, n + 1):
sweep_inner_axes()
nested_sweep(
axes=outer_axes,
measurement_fn=measure_at_outer_point,
)
debug_buffer_plan(
output_axes=qua_loop_axes,
n_avg=n_avg,
)
with qua.stream_processing():
for i in range(n_iq_channels):
apply_buffers(
I_st[i],
output_axes=qua_loop_axes,
n_avg=n_avg,
).save(f"I_{i + 1}")
apply_buffers(
Q_st[i],
output_axes=qua_loop_axes,
n_avg=n_avg,
).save(f"Q_{i + 1}")
return qua_prog