from qm import qua
from qualang_tools.loops import from_array
from qualang_tools.units import unit
from quam.core import QuamComponent
from typing import cast
from semi_cr.core.lab.devices.tankinductor import TankInductor
from semi_cr.core.lab.protocols.inputs.tankinductor import TankInductorProtocolInput
from semi_cr.core.lab.measurements.axis_factories.tankinductor import TANKINDUCTOR_AXIS_FACTORIES
u = unit(coerce_to_integer=True)
[docs]
def measure_tank(
device: TankInductor,
I_st,
Q_st,
):
rf_in = device.RF_IN
amplitude_scale = rf_in._amplitude_scale
tank_circuit = cast("QuamComponent", rf_in.quam.channel)
I, Q = tank_circuit.measure(
"readout",
amplitude_scale=amplitude_scale,
)
qua.save(I, I_st)
qua.save(Q, Q_st)
[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
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_axis_from_config(axis_name, sweep_config, protocol_inputs, tank_circuit):
try:
factory = AXIS_FACTORIES[axis_name]
except KeyError:
raise ValueError(f"No axis factory registered for {axis_name!r}")
return factory(
sweep_config=sweep_config,
protocol_inputs=protocol_inputs,
tank_circuit=tank_circuit,
)
[docs]
def is_enabled(config) -> bool:
return getattr(config, "enabled", "OFF") == "ON"
[docs]
def get_enabled_tankinductor_axes(
protocol_inputs: TankInductorProtocolInput,
) -> tuple:
parameter_sweeps = protocol_inputs.parameter_sweeps
enabled_axes = list(parameter_sweeps.enabled_sweeps().keys())
if "frequency" not in enabled_axes:
raise ValueError("frequency must be enabled for this QUA program.")
return tuple(enabled_axes)
[docs]
def get_sweep_config(protocol_inputs, axis_name):
axis_config = protocol_inputs.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_tankinductor_sweep_axes(
device: TankInductor,
protocol_inputs,
):
enabled_axes = get_enabled_tankinductor_axes(
protocol_inputs=protocol_inputs,
)
axes = []
for axis_name in enabled_axes:
axis_config = protocol_inputs.parameter_sweeps[axis_name]
sweep_config = get_sweep_config(protocol_inputs, axis_name)
factory = TANKINDUCTOR_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 apply_buffers(
stream,
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
]
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)
return stream.map(qua.FUNCTIONS.average(average_axis))
[docs]
def get_tankinductor_nd_sweep_program(
device,
qua_sweep_axes,
n_avg,
magnetic_field: float = 0,
):
with qua.program() as qua_prog:
n = qua.declare(int)
I_st = qua.declare_stream()
Q_st = qua.declare_stream()
n_st = qua.declare_stream()
for axis in qua_sweep_axes:
axis.qua_var = qua.declare(axis.qua_type)
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
def measure_at_inner_point():
device.RF_IN.pulse_amplitude.set(amplitude_var)
measure_tank(
device,
I_st=I_st,
Q_st=Q_st,
)
def sweep_inner_axes():
nested_sweep(
axes=inner_axes,
measurement_fn=measure_at_inner_point,
)
def measure_at_outer_point():
if magnetic_field == 0:
qua.wait(
1000 * u.ns,
device.RF_IN.quam.channel.id,
)
with qua.for_(n, 0, n < n_avg, n + 1):
sweep_inner_axes()
# qua.save(n, n_st)
nested_sweep(
axes=outer_axes,
measurement_fn=measure_at_outer_point,
)
with qua.stream_processing():
apply_buffers(
I_st,
output_axes=qua_loop_axes,
n_avg=n_avg,
).save("I")
apply_buffers(
Q_st,
output_axes=qua_loop_axes,
n_avg=n_avg,
).save("Q")
# n_st.save("iteration")
return qua_prog