Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 33 additions & 10 deletions qoolqit/drive.py
Original file line number Diff line number Diff line change
@@ -1,17 +1,22 @@
from __future__ import annotations

from dataclasses import dataclass
from typing import Any
from typing import Any, cast

import matplotlib.pyplot as plt
import numpy as np
from matplotlib.figure import Figure

from qoolqit.waveforms import CompositeWaveform, DelayWaveform, Waveform
from qoolqit.waveforms import CompositeWaveform, ConstantWaveform, DelayWaveform, Waveform

__all__ = ["DetuningMapModulator", "Drive"]


def _leaves(waveform: Waveform) -> list[Waveform]:
"""Returns the non-composite waveforms making up `waveform`, in order."""
return waveform.waveforms if isinstance(waveform, CompositeWaveform) else [waveform]


@dataclass(frozen=True)
class DetuningMapModulator:
"""A weighted detuning for the Detuning Map Modulator (DMM).
Expand Down Expand Up @@ -121,7 +126,11 @@ def __init__(
if dmm is not None and not isinstance(dmm, DetuningMapModulator):
raise TypeError("'dmm' must be of type DetuningMapModulator.")
self._dmm = dmm
self._phase = phase
self._phase_wf: Waveform = ConstantWaveform(duration=self._duration, value=phase)
# one (amplitude, detuning) pair per original Drive joined by `>>`, always in lockstep
# with `_leaves(self._phase_wf)` -- unlike the padded composites above, whose leaf counts
# can differ between amplitude and detuning.
self._parts: tuple[tuple[Waveform, Waveform], ...] = ((self._amplitude, self._detuning),)

@property
def amplitude(self) -> Waveform:
Expand All @@ -140,8 +149,15 @@ def dmm(self) -> DetuningMapModulator | None:

@property
def phase(self) -> float:
"""The phase value in the drive."""
return self._phase
"""The phase value in the drive.

Raises:
ValueError: If the drive was composed from segments with different phases.
"""
values = {cast(ConstantWaveform, p).value for p in _leaves(self._phase_wf)}
if len(values) > 1:
raise ValueError("Drive has a varying phase; compile it rather than reading `.phase`.")
return values.pop()

@property
def duration(self) -> float:
Expand All @@ -152,13 +168,13 @@ def __rshift__(self, other: Drive) -> Drive:

def __rrshift__(self, other: Drive) -> Drive:
if isinstance(other, Drive):
if self.phase != other.phase:
raise NotImplementedError("Composing drives with different phase not supported.")
return Drive(
new = Drive(
amplitude=CompositeWaveform(self._amplitude, other._amplitude),
detuning=CompositeWaveform(self._detuning, other._detuning),
phase=self._phase,
)
new._phase_wf = CompositeWaveform(self._phase_wf, other._phase_wf)
new._parts = self._parts + other._parts
return new
else:
raise NotImplementedError(f"Composing with object of type {type(other)} not supported.")

Expand Down Expand Up @@ -199,7 +215,7 @@ def __repr__(self) -> str:

def draw(self, return_fig: bool = False) -> Figure | None:

nrows = 3 if self.dmm is not None else 2
nrows = 4 if self.dmm is not None else 3

fig = plt.gcf()
axs = fig.subplots(nrows, 1, sharex=True)
Expand All @@ -208,6 +224,7 @@ def draw(self, return_fig: bool = False) -> Figure | None:
t_array = np.linspace(0.0, self.duration, 250)
y_amp = self.amplitude(t_array)
y_det = self.detuning(t_array)
y_phase = self._phase_wf(t_array)

# draw amplitude
axs[0].grid(True, color="lightgray", linestyle="--", linewidth=0.7)
Expand All @@ -222,6 +239,12 @@ def draw(self, return_fig: bool = False) -> Figure | None:
axs[1].plot(t_array, y_det, color="darkmagenta")
axs[1].fill_between(t_array, y_det, color="darkmagenta", alpha=0.4)

# draw phase
axs[2].grid(True, color="lightgray", linestyle="--", linewidth=0.7)
axs[2].set_ylabel("Phase")
axs[2].plot(t_array, y_phase, color="darkorange")
axs[2].fill_between(t_array, y_phase, color="darkorange", alpha=0.4)

axs[-1].set_xlabel("Time t")

# draw DMM if present
Expand Down
31 changes: 23 additions & 8 deletions qoolqit/execution/compilation_functions.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from __future__ import annotations

from enum import Enum
from typing import cast

from pulser.devices import Device as PulserDevice
from pulser.parametrized import ParamObj
Expand All @@ -10,9 +11,10 @@
from pulser.waveforms import Waveform as PulserWaveform

from qoolqit.devices import Device
from qoolqit.drive import DetuningMapModulator, Drive, Waveform
from qoolqit.drive import DetuningMapModulator, Drive, Waveform, _leaves
from qoolqit.exceptions import CompilationError
from qoolqit.register import Register
from qoolqit.waveforms import ConstantWaveform


class CompilerProfile(Enum):
Expand Down Expand Up @@ -61,6 +63,19 @@ def convert(self, waveform: Waveform) -> ParamObj | PulserWaveform:
return waveform._to_pulser(duration=pulser_duration) * self._energy


def _phase_groups(drive: Drive) -> list[tuple[Waveform, Waveform, float]]:
"""Splits a drive into contiguous (amplitude, detuning, phase) runs of constant phase."""
groups: list[tuple[Waveform, Waveform, float]] = []
for (amp, det), ph in zip(drive._parts, _leaves(drive._phase_wf)):
phase_value = cast(ConstantWaveform, ph).value
if groups and groups[-1][2] == phase_value:
prev_amp, prev_det, merged_phase = groups.pop()
groups.append((prev_amp >> amp, prev_det >> det, merged_phase))
else:
groups.append((amp, det, phase_value))
return groups


def basic_compilation(
register: Register,
drive: Drive,
Expand Down Expand Up @@ -124,19 +139,19 @@ def basic_compilation(
if device_max_duration_ratio and device._max_duration:
TIME = device_max_duration_ratio * device._max_duration / drive.duration

# Build pulser pulse and register
# Build pulser register and sequence
wf_converter = WaveformConverter(device=device, time=TIME, energy=ENERGY)
pulser_amp_wf = wf_converter.convert(drive._amplitude)
pulser_det_wf = wf_converter.convert(drive._detuning)
pulser_pulse = PulserPulse(pulser_amp_wf, pulser_det_wf, drive.phase)

pulser_register = _build_register(register, device, DISTANCE)

# Create sequence
pulser_device = device._device
pulser_sequence = PulserSequence(pulser_register, pulser_device)
pulser_sequence.declare_channel("rydberg", "rydberg_global")
pulser_sequence.add(pulser_pulse, "rydberg")

# One PulserPulse per contiguous same-phase run in the drive.
for amp_wf, det_wf, phase in _phase_groups(drive):
pulser_amp_wf = wf_converter.convert(amp_wf)
pulser_det_wf = wf_converter.convert(det_wf)
pulser_sequence.add(PulserPulse(pulser_amp_wf, pulser_det_wf, phase), "rydberg")

# Add dmm, if specified in the drive.
if drive.dmm is not None:
Expand Down
9 changes: 5 additions & 4 deletions tests/test_drive.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,10 +67,11 @@ def test_drive_init_and_composition(amp_wf: Waveform, det_wf: Waveform) -> None:
drive = drive_rand_phase >> drive_rand_phase
assert math.isclose(drive.phase, phase)

with pytest.raises(NotImplementedError):
drive1 = Drive(amplitude=amp_wf, detuning=det_wf, phase=1.0)
drive2 = Drive(amplitude=amp_wf, detuning=det_wf, phase=0.0)
drive = drive1 >> drive2
drive1 = Drive(amplitude=amp_wf, detuning=det_wf, phase=1.0)
drive2 = Drive(amplitude=amp_wf, detuning=det_wf, phase=0.0)
drive = drive1 >> drive2
with pytest.raises(ValueError, match="varying phase"):
drive.phase


def test_error_amplitude_negative() -> None:
Expand Down
Loading