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
|