from dataclasses import dataclass, field
from typing import Any, ClassVar
from typing import Literal
from pydantic.v1 import BaseModel, Field, validator
from semi_cr.core.lab.utils.dataclass_base import Dataclass
from enum import StrEnum, Enum
[docs]
class RangeWithUnit(BaseModel):
min_val: ClassVar[float]
max_val: ClassVar[float]
unit: ClassVar[str]
start: float = Field(..., unit="") # placeholder; overridden in schema via __init_subclass__
stop: float = Field(..., unit="")
step: float = Field(..., unit="")
num_points: int = Field(..., ge=1)
def __init_subclass__(cls, **kwargs):
# Dynamically set fields with validators & metadata on subclass creation
super().__init_subclass__(**kwargs)
# set unit metadata in v1:
for field_name in ("start", "stop", "step"):
cls.__fields__[field_name].field_info.extra["unit"] = cls.unit
@validator("start", "stop", "step")
def _range_check(cls, v):
if not (cls.min_val <= v <= cls.max_val):
raise ValueError(f"must be in [{cls.min_val}, {cls.max_val}] {cls.unit}")
return v
[docs]
class ParameterWithValidRange(BaseModel):
min_val: ClassVar[float]
max_val: ClassVar[float]
unit: ClassVar[str]
value: float
[docs]
@validator("value")
def check_range(cls, v):
if not (cls.min_val <= v <= cls.max_val):
raise ValueError(f"value must be in [{cls.min_val}, {cls.max_val}] {cls.unit}")
return v
[docs]
class TemperatureWithAttributes(RangeWithUnit):
min_val = 0
max_val = 300
unit = "K"
[docs]
class FieldWithAttributes(RangeWithUnit):
min_val = 0
max_val = 5
unit = "T"
[docs]
class PowerWithAttributes(RangeWithUnit):
min_val = -40
max_val = 20
unit = "dBm"
[docs]
class FrequencyWithAttributes(RangeWithUnit):
min_val = 10e6
max_val = 9e9
unit = "Hz"
[docs]
class TemperatureParameter(ParameterWithValidRange):
min_val = 0
max_val = 300
unit = "K"
[docs]
class RfSweepParameters(BaseModel):
# vna: str
sweep_mode: str
Sparams: list[str]
attenuation_RT: float = Field(ge=0, json_schema_extra={"unit": "dB"})
amplification_RT: float = Field(ge=0, json_schema_extra={"unit": "dB"})
averages: int = Field(ge=1)
[docs]
class DeviceParameters(BaseModel):
device_name: str
device_group: list[str]
devices: list[str]
[docs]
@dataclass
class IQChannels:
n_chan: int
[docs]
@dataclass
class RFFrequencySweep:
start: int
stop: int
step: int
n_points: int
[docs]
@dataclass
class PulseAmplitudeSweep:
start: int
stop: int
step: int
n_points: int
[docs]
@dataclass
class ReadoutPulse:
amplitude: float
length: int
# @dataclass
# class HighVoltageParameter(ParameterWithValidRange):
# min_val = -5
# max_val = 5
# unit = "V"
# @dataclass
# class LowVoltageParameter(ParameterWithValidRange):
# min_val = -5
# max_val = 5
# unit = "V"
[docs]
@dataclass
class VoltageParameter:
high: float # HighVoltageParameter
low: float # LowVoltageParameter
[docs]
@dataclass
class QDACPortNumber:
min_value: ClassVar[int] = 1
max_value: ClassVar[int] = 8
unit: ClassVar[str] = ""
[docs]
@dataclass
class DelayParameter:
min_val: ClassVar[int] = 0
max_val: ClassVar[int] = 10
unit: ClassVar[str] = "s"
[docs]
@dataclass
class SwitchingNumber:
min_val: ClassVar[int] = 1
max_val: ClassVar[int] = 10
unit: ClassVar[str] = ""
# @dataclass
# class Trigger:
# port: QDACPortNumber
# delay: DelayParameter
[docs]
@dataclass
class QDACTrigger:
voltage: VoltageParameter
port: int
delay: float
n_switching: int
[docs]
@dataclass
class DCAmplitudeSweep:
start: int
stop: int
step: int
n_points: int
# @dataclass
# class AmplitudeWithAttributes(RangeWithUnit):
# min_val = 0
# max_val = 0.5
# unit = "V"
# @dataclass
# class AmplitudeParameter(ParameterWithValidRange):
# min_val = 0
# max_val = 0.5
# unit = "V"
# @dataclass
# class DurationParameter(ParameterWithValidRange):
# min_val = 0
# max_val = 1000000
# unit = "ns"
# @dataclass
# class ReadoutPulse(BaseModel):
# amplitude: int
# length: int
[docs]
@dataclass
class InOutOPXPorts:
outputs: list[int]
inputs: list[int]
[docs]
@dataclass
class OPXTrigger:
id: int | str
port: int
intermediate_frequency: int
delay: int
buffer: int
[docs]
@dataclass
class Averaging:
enabled: bool
n_avg: int
[docs]
@dataclass
class TemperatureSweep:
start: int
stop: int
step: int
n_points: int
[docs]
@dataclass
class MagneticFieldSweep:
start: int
stop: int
step: int
n_points: int
[docs]
@dataclass
class SwitchStateSweep:
start: int
stop: int
step: int
n_points: int
[docs]
class SweepExecution(str, Enum):
EXTERNAL = "external"
QUA = "qua"
[docs]
@dataclass
class SettlingTimeSweep:
execution: SweepExecution
start: float
stop: float
step: float
n_points: int
[docs]
@dataclass
class VoltageSweep:
start: float
stop: float
step: float
n_points: int
[docs]
@dataclass
class CurrentSweep:
start: float
stop: float
step: float
n_points: int
[docs]
@dataclass
class ParameterSweep:
enabled: bool
source: str
name: str | None = None
label: str | None = None
unit: str | None = None
pin: str | None = None
pins: list[str] = field(default_factory=list)
quantity: str | None = None
metadata: dict[str, Any] = field(default_factory=dict)
[docs]
@dataclass
class ParameterSweeps:
sweeps: dict[str, ParameterSweep] = field(default_factory=dict)
[docs]
@classmethod
def from_dict(cls, data: dict[str, dict[str, Any]]) -> "ParameterSweeps":
sweeps = {}
known_fields = {
"enabled",
"source",
"name",
"label",
"unit",
"pin",
"pins"
"quantity",
}
for axis_key, cfg in data.items():
extra = {
k: v for k, v in cfg.items()
if k not in known_fields
}
sweeps[axis_key] = ParameterSweep(
enabled=str(cfg.get("enabled", "OFF")).upper() == "ON",
source=cfg["source"],
name=cfg.get("name", axis_key),
label=cfg.get("label", axis_key),
unit=cfg.get("unit"),
pin=cfg.get("pin"),
pins=cfg.get("pins"),
quantity=cfg.get("quantity"),
metadata=extra,
)
return cls(sweeps=sweeps)
def __getitem__(self, name: str) -> ParameterSweep:
return self.sweeps[name]
[docs]
def items(self):
return self.sweeps.items()
[docs]
def keys(self):
return self.sweeps.keys()
[docs]
def values(self):
return self.sweeps.values()
[docs]
def enabled_sweeps(self) -> dict[str, ParameterSweep]:
return {
name: sweep
for name, sweep in self.sweeps.items()
if sweep.enabled
}
[docs]
def get(self, name: str) -> ParameterSweep | None:
return self.sweeps.get(name)
[docs]
@dataclass
class FixedVoltageConfig:
enabled: Literal["ON", "OFF"] = "ON"
pin: str = field(default_factory=str)
quantity: str = "dc_voltage"
value: float = field(default_factory=float)
name: str = field(default_factory=str)
label: str = field(default_factory=str)
unit: str = "V"
@property
def is_enabled(self) -> bool:
return self.enabled == "ON"
[docs]
class SweepStrategy(StrEnum):
COMBINED = "combined"
INDEPENDENT = "independent"
[docs]
@dataclass
class BaseParameterSpec(Dataclass):
enabled: str = "OFF"
name: str = ""
label: str = ""
unit: str = ""
pin: str | None = None
pins: list[str] | None = None
quantity: str | None = None
[docs]
@dataclass
class ParameterSpec(BaseParameterSpec):
source: str = ""
[docs]
@dataclass
class GateSweepParameterSpec(ParameterSpec):
strategy: SweepStrategy | None = None
[docs]
@dataclass
class FixedParameterSpec(BaseParameterSpec):
value: float | int | str | None = None
[docs]
def is_enabled(value: str | bool) -> bool:
if isinstance(value, bool):
return value
return value.strip().upper() in {
"ON",
"TRUE",
"YES",
"1",
}
[docs]
@dataclass
class ProtocolParameters(Dataclass):
swept: dict[str, ParameterSpec] = field(default_factory=dict)
fixed: dict[str, FixedParameterSpec] = field(default_factory=dict)
[docs]
def enabled_swept(self) -> dict[str, ParameterSpec]:
return {
name: spec
for name, spec in self.swept.items()
if is_enabled(spec.enabled)
}
[docs]
def enabled_fixed(self) -> dict[str, FixedParameterSpec]:
return {
name: spec
for name, spec in self.fixed.items()
if is_enabled(spec.enabled)
}
[docs]
@dataclass
class GateSweepProtocolParameters(Dataclass):
swept: dict[str, GateSweepParameterSpec] = field(default_factory=dict)
fixed: dict[str, FixedParameterSpec] = field(default_factory=dict)
[docs]
def enabled_swept(self) -> dict[str, GateSweepParameterSpec]:
return {
name: spec
for name, spec in self.swept.items()
if is_enabled(spec.enabled)
}
[docs]
def enabled_fixed(self) -> dict[str, FixedParameterSpec]:
return {
name: spec
for name, spec in self.fixed.items()
if is_enabled(spec.enabled)
}