Skip to content

flux_base

Flux base device.

FluxDevice

Bases: Device

Source code in jaxquantum/devices/superconducting/flux_base.py
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
@struct.dataclass
class FluxDevice(Device):
    @abstractmethod
    def phi_zpf(self):
        """Return Phase ZPF."""

    def _calculate_wavefunctions_fock(self, phi_vals):
        """Calculate wavefunctions at phi_exts."""
        phi_osc = self.phi_zpf() * jnp.sqrt(2)  # length of oscillator
        basis_functions = harm_osc_wavefunctions(
            self.N_pre_diag,
            phi_vals,
            jnp.real(phi_osc),
        )
        return self.get_vec_data_in_H_eigenbasis(basis_functions)

    def _calculate_wavefunctions_charge(self, phi_vals):
        phi_vals = jnp.array(phi_vals)
        n_labels = jnp.diag(self.original_ops["n"].data)
        basis_functions = jnp.exp(
            -2j * jnp.pi * n_labels[:, None] * phi_vals
        ) / jnp.sqrt(2 * jnp.pi)
        wavefunctions = self.get_vec_data_in_H_eigenbasis(
            basis_functions,
        )
        correction = jnp.power(1j, jnp.arange(wavefunctions.shape[-2]))
        return correction[:, None] * wavefunctions

    @abstractmethod
    def potential(self, phi):
        """Return potential energy as a function of phi."""

    def plot_wavefunctions(self, phi_vals, max_n=None, which=None, ax=None, mode="abs", ylim=None, y_scale_factor=1, zero_potential=False, wavefunction_color=None):
        if self.basis == BasisTypes.fock:
            _calculate_wavefunctions = self._calculate_wavefunctions_fock
        elif self.basis == BasisTypes.charge:
            _calculate_wavefunctions = self._calculate_wavefunctions_charge
        else:
            raise NotImplementedError(
                f"The {self.basis} is not yet supported for plotting wavefunctions."
            )

        """Plot wavefunctions at phi_exts."""
        wavefunctions = _calculate_wavefunctions(phi_vals)
        energy_levels = self.eig_systems["vals"][: self.N]

        potential = self.potential(phi_vals)

        min_potential = 0 if not zero_potential else jnp.min(potential)
        if ax is None:
            fig, ax = plt.subplots(1, 1, figsize=(3.5, 2.5), dpi=1000)
        else:
            fig = ax.get_figure()

        min_val = None
        max_val = None

        assert max_n is None or which is None, "Can't specify both max_n and which"

        max_n = self.N if max_n is None else max_n
        levels = range(max_n) if which is None else which

        for n in levels:
            if mode == "abs":
                wf_vals = jnp.abs(wavefunctions[n, :]) ** 2
            elif mode == "real":
                wf_vals = wavefunctions[n, :].real
            elif mode == "imag":
                wf_vals = wavefunctions[n, :].imag

            wf_vals += energy_levels[n]
            curr_min_val = min(wf_vals)
            curr_max_val = max(wf_vals)

            if min_val is None or curr_min_val < min_val:
                min_val = curr_min_val

            if max_val is None or curr_max_val > max_val:
                max_val = curr_max_val

            extra_kwargs = {}
            if wavefunction_color is not None:
                if isinstance(wavefunction_color, list):
                    extra_kwargs["color"] = wavefunction_color[n]
                else:
                    extra_kwargs["color"] = wavefunction_color

            ax.plot(
                phi_vals, (wf_vals - min_potential)*y_scale_factor, label=f"$|${n}$\\rangle$", linestyle="-", linewidth=1, **extra_kwargs
            )

            ax.fill_between(phi_vals, (energy_levels[n] - min_potential)*y_scale_factor, (wf_vals - min_potential)*y_scale_factor, alpha=0.5, **extra_kwargs)

        ax.plot(
            phi_vals,
            (potential - min_potential)*y_scale_factor,
            label="potential",
            color="black",
            linestyle="-",
            linewidth=1,
        )

        ylim = ylim if ylim is not None else [jnp.min(jnp.array([min_val - 1 - min_potential, jnp.min(potential) - min_potential]))*y_scale_factor, (max_val + 1 - min_potential)*y_scale_factor]
        ax.set_ylim(ylim)
        ax.set_xlabel(r"$\varphi/2\pi$")
        ax.set_ylabel(r"Energy [GHz]")

        if mode == "abs":
            title_str = r"$|\psi_n(\Phi)|^2$"
        elif mode == "real":
            title_str = r"Re($\psi_n(\Phi)$)"
        elif mode == "imag":
            title_str = r"Im($\psi_n(\Phi)$)"

        ax.set_title(f"{title_str}")

        ax.legend(fontsize='xx-small')
        fig.tight_layout()

        return ax

phi_zpf() abstractmethod

Return Phase ZPF.

Source code in jaxquantum/devices/superconducting/flux_base.py
18
19
20
@abstractmethod
def phi_zpf(self):
    """Return Phase ZPF."""

potential(phi) abstractmethod

Return potential energy as a function of phi.

Source code in jaxquantum/devices/superconducting/flux_base.py
44
45
46
@abstractmethod
def potential(self, phi):
    """Return potential energy as a function of phi."""