cavsqueeze 1.7.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
cavsqueeze/__init__.py ADDED
@@ -0,0 +1,8 @@
1
+ """cavsqueeze: beyond-mean-field cavity-mediated spin squeezing for solid-state
2
+ clock-transition ensembles (second-order cumulant expansion with disorder,
3
+ collective decay, thermal photons and coupling inhomogeneity)."""
4
+ from .resonator import CavityParams, from_hz, thermal_occupation, loop_gap_dispersive
5
+ from .ensemble import Ensemble, equal_probability_classes, homogeneous, lineshape, product_classes, log_uniform_weights
6
+ from .cumulant import Rates, State, product_state, rotate, evolve, evolve_meanfield, wineland_xi2, collective_moments, coherence, transverse_variances
7
+
8
+ __version__ = "1.7.0"
cavsqueeze/cumulant.py ADDED
@@ -0,0 +1,435 @@
1
+ """Second-order cumulant expansion in *connected* (cumulant) variables.
2
+
3
+ This is the production solver. It integrates the same second-order cumulant
4
+ hierarchy as `cumulant_raw.py` (which is validated against exact master
5
+ equations) but in terms of connected correlations
6
+
7
+ Pc[m,n] = <sigma+_a sigma-_b> - <sigma+_a><sigma-_b>
8
+ Qc[m,n] = <sigma+_a sigma+_b> - <sigma+_a><sigma+_b>
9
+ Rc[m,n] = <sigma z_a sigma+_b> - <sigma z_a><sigma+_b>
10
+ Zc[m,n] = <sigma z_a sigma z_b> - <sigma z_a><sigma z_b>
11
+
12
+ (a in class m, b in class n, a != b). Connected correlations are O(1/N) while
13
+ raw moments are O(1); evolving them directly avoids the catastrophic
14
+ cancellation that makes raw moments useless for N > 1e8. The equations are
15
+ linear in the connected variables, with source terms of order (rate)/N, and
16
+ the first moments receive O(1/N) feedback from the correlations, so that the
17
+ two formulations are mathematically identical (tests/test_connected.py).
18
+
19
+ Model (cavity adiabatically eliminated, all rates in rad/s):
20
+
21
+ H = sum_{j,k} chi1 G_j G_k sigma+_j sigma-_k + sum_j (delta_j/2) sigma z_j
22
+ L1 = sqrt(Gd) sum_j G_j sigma-_j collective emission, Gd = Gamma_SR (n_th + 1)
23
+ L2 = sqrt(Gu) sum_j G_j sigma+_j collective absorption, Gu = Gamma_SR n_th
24
+ L_j = sqrt(gamma_phi/2) sigma z_j individual pure dephasing
25
+
26
+ The truncation rule for the connected third-order moments is
27
+ <XYB>_c = <X><YB>_c + <Y><XB>_c (Gaussian closure).
28
+ """
29
+ from __future__ import annotations
30
+
31
+ import dataclasses
32
+ import numpy as np
33
+ from scipy.integrate import solve_ivp
34
+
35
+ from .ensemble import Ensemble
36
+ from .resonator import CavityParams
37
+ from .cumulant_raw import Rates as _RatesRaw, rotation_matrix # noqa: F401
38
+
39
+
40
+ class Rates(_RatesRaw):
41
+ """Rates of the interacting classes plus the spectator (free) spins."""
42
+
43
+ spec_delta: np.ndarray = None
44
+ spec_n: np.ndarray = None
45
+ meas: float = 0.0 # continuous QND measurement of J_z: conditional variance 1/(meas t) for pure measurement
46
+ meas_eta: float = 1.0 # detection efficiency: the back-action dephasing rate is meas/(8 meas_eta)
47
+
48
+ @classmethod
49
+ def from_params(cls, params: CavityParams, ens: Ensemble) -> "Rates":
50
+ base = _RatesRaw.from_params(params, ens)
51
+ rt = cls(delta=base.delta, G=base.G, n=base.n, chi1=base.chi1, Gd=base.Gd, Gu=base.Gu, gamma_phi=base.gamma_phi)
52
+ rt.spec_delta = np.asarray(ens.spec_delta, float)
53
+ rt.spec_n = np.asarray(ens.spec_n, float)
54
+ return rt
55
+
56
+ @property
57
+ def K(self) -> int:
58
+ return len(self.spec_delta) if self.spec_delta is not None else 0
59
+
60
+ @property
61
+ def N_total(self) -> float:
62
+ return float(self.n.sum() + (self.spec_n.sum() if self.spec_n is not None else 0.0))
63
+
64
+
65
+ @dataclasses.dataclass
66
+ class State:
67
+ """Means and connected correlations of an M-class ensemble."""
68
+
69
+ s: np.ndarray
70
+ z: np.ndarray
71
+ Pc: np.ndarray
72
+ Qc: np.ndarray
73
+ Rc: np.ndarray
74
+ Zc: np.ndarray
75
+ vs: np.ndarray = None # (K,3) Bloch vectors of spectator (free, uncorrelated) spins
76
+
77
+ def __post_init__(self):
78
+ if self.vs is None:
79
+ self.vs = np.zeros((0, 3))
80
+
81
+ @property
82
+ def M(self):
83
+ return len(self.s)
84
+
85
+ def pack(self) -> np.ndarray:
86
+ return np.concatenate(
87
+ [self.s, self.z.astype(complex), self.Pc.ravel(), self.Qc.ravel(), self.Rc.ravel(), self.Zc.ravel()]
88
+ )
89
+
90
+ @classmethod
91
+ def unpack(cls, y: np.ndarray, M: int, vs=None) -> "State":
92
+ s = y[:M]
93
+ z = y[M : 2 * M]
94
+ b = y[2 * M :].reshape(4, M, M)
95
+ return cls(s=s, z=z, Pc=b[0], Qc=b[1], Rc=b[2], Zc=b[3], vs=vs)
96
+
97
+ def copy(self):
98
+ return State(self.s.copy(), self.z.copy(), self.Pc.copy(), self.Qc.copy(), self.Rc.copy(), self.Zc.copy(), self.vs.copy())
99
+
100
+ # raw moments (for comparison with cumulant_raw / exact solvers)
101
+ def raw(self):
102
+ s, z = self.s, self.z
103
+ return (
104
+ self.Pc + np.outer(s, np.conj(s)),
105
+ self.Qc + np.outer(s, s),
106
+ self.Rc + np.outer(z, s),
107
+ self.Zc + np.outer(z, z),
108
+ )
109
+
110
+
111
+ def product_state(M: int, v, K_spec: int = 0) -> State:
112
+ v = np.asarray(v, dtype=float)
113
+ s = np.full(M, 0.5 * (v[0] + 1j * v[1]), dtype=complex)
114
+ z = np.full(M, v[2], dtype=complex)
115
+ zero = np.zeros((M, M), dtype=complex)
116
+ vs = np.tile(v, (K_spec, 1))
117
+ return State(s, z, zero.copy(), zero.copy(), zero.copy(), zero.copy(), vs)
118
+
119
+
120
+ # ---------------------------------------------------------------------------
121
+ # Cartesian representation, rotations, collective moments
122
+ # ---------------------------------------------------------------------------
123
+
124
+ def to_cartesian(st: State):
125
+ """Bloch vectors v (M,3) and connected correlation tensor Cc (M,M,3,3)."""
126
+ s, z, P, Q, R, Z = st.s, st.z, st.Pc, st.Qc, st.Rc, st.Zc
127
+ v = np.stack([2 * s.real, 2 * s.imag, z.real], axis=1)
128
+ C = np.empty((st.M, st.M, 3, 3), dtype=complex)
129
+ Qc = np.conj(Q)
130
+ PT = P.T
131
+ C[:, :, 0, 0] = Q + P + PT + Qc
132
+ C[:, :, 0, 1] = -1j * (Q - P + PT - Qc)
133
+ C[:, :, 1, 0] = -1j * (Q + P - PT - Qc)
134
+ C[:, :, 1, 1] = -(Q - P - PT + Qc)
135
+ C[:, :, 0, 2] = 2 * R.T.real
136
+ C[:, :, 1, 2] = 2 * R.T.imag
137
+ C[:, :, 2, 0] = 2 * R.real
138
+ C[:, :, 2, 1] = 2 * R.imag
139
+ C[:, :, 2, 2] = Z
140
+ return v, C
141
+
142
+
143
+ def from_cartesian(v, C) -> State:
144
+ s = 0.5 * (v[:, 0] + 1j * v[:, 1])
145
+ z = v[:, 2].astype(complex)
146
+ Cxx, Cxy, Cyx, Cyy = C[:, :, 0, 0], C[:, :, 0, 1], C[:, :, 1, 0], C[:, :, 1, 1]
147
+ P = 0.25 * (Cxx - 1j * Cxy + 1j * Cyx + Cyy)
148
+ Q = 0.25 * (Cxx + 1j * Cxy + 1j * Cyx - Cyy)
149
+ R = 0.5 * (C[:, :, 2, 0] + 1j * C[:, :, 2, 1])
150
+ Z = C[:, :, 2, 2]
151
+ return State(s, z, P, Q, R, Z)
152
+
153
+
154
+ def rotate(st: State, axis, angle: float) -> State:
155
+ """Global rotation of every Bloch vector by `angle` about `axis` (a pulse)."""
156
+ Rm = rotation_matrix(axis, angle)
157
+ v, C = to_cartesian(st)
158
+ out = from_cartesian(v @ Rm.T, np.einsum("ac,bd,mncd->mnab", Rm, Rm, C))
159
+ out.vs = st.vs @ Rm.T
160
+ return out
161
+
162
+
163
+ def rotate_classes(st: State, axis, angles, spec_angles=None) -> State:
164
+ """Rotation of each class m by its own angle angles[m] about `axis` (a pulse
165
+ whose rotation angle varies between classes, e.g. with the coupling weight).
166
+ Spectator spins rotate by spec_angles (default: the mean of angles)."""
167
+ angles = np.asarray(angles, float)
168
+ Rs = np.stack([rotation_matrix(axis, a) for a in angles]) # (M,3,3)
169
+ v, C = to_cartesian(st)
170
+ vn = np.einsum("mab,mb->ma", Rs, v)
171
+ Cn = np.einsum("mac,nbd,mncd->mnab", Rs, Rs, C)
172
+ out = from_cartesian(vn, Cn)
173
+ K = st.vs.shape[0]
174
+ if K:
175
+ sa = np.full(K, float(angles.mean())) if spec_angles is None else np.asarray(spec_angles, float)
176
+ out.vs = np.stack([st.vs[k] @ rotation_matrix(axis, sa[k]).T for k in range(K)])
177
+ else:
178
+ out.vs = st.vs
179
+ return out
180
+
181
+
182
+ def collective_moments(st: State, n, weights=None, spec_n=None, spec_weights=None):
183
+ """Mean <J> and symmetrised covariance of J_alpha = 1/2 sum_a c_a sigma^alpha_a,
184
+ including spectator spins (uncorrelated, each in a pure state with Bloch vector vs)."""
185
+ n = np.asarray(n, float)
186
+ c = np.ones(st.M) if weights is None else np.asarray(weights, float)
187
+ v, C = to_cartesian(st)
188
+ nc = n * c
189
+ J = 0.5 * (nc @ v)
190
+ pair = np.einsum("m,n,mnab->ab", nc, nc, C) - np.einsum("m,mmab->ab", n * c * c, C)
191
+ same = np.sum(n * c * c) * np.eye(3) - np.einsum("m,ma,mb->ab", n * c * c, v, v)
192
+ Cov = 0.25 * (pair.real + same)
193
+ S1 = float(nc.sum())
194
+ S2 = float(np.sum(n * c * c))
195
+ K = st.vs.shape[0]
196
+ if K and spec_n is not None:
197
+ sn = np.asarray(spec_n, float)
198
+ sc = np.ones(K) if spec_weights is None else np.asarray(spec_weights, float)
199
+ J = J + 0.5 * ((sn * sc) @ st.vs)
200
+ Cov = Cov + 0.25 * (np.sum(sn * sc * sc) * np.eye(3) - np.einsum("k,ka,kb->ab", sn * sc * sc, st.vs, st.vs))
201
+ S1 += float((sn * sc).sum())
202
+ S2 += float(np.sum(sn * sc * sc))
203
+ Cov = 0.5 * (Cov + Cov.T)
204
+ return J, Cov, S1, S2
205
+
206
+
207
+ def transverse_variances(J, Cov):
208
+ """(var_min, var_max, angle) of the covariance in the plane perpendicular to J."""
209
+ Jn = np.linalg.norm(J)
210
+ e3 = J / Jn
211
+ trial = np.array([0.0, 0.0, 1.0]) if abs(e3[2]) < 0.9 else np.array([1.0, 0.0, 0.0])
212
+ e1 = np.cross(e3, trial)
213
+ e1 /= np.linalg.norm(e1)
214
+ e2 = np.cross(e3, e1)
215
+ B = np.stack([e1, e2], axis=1)
216
+ vals, vecs = np.linalg.eigh(B.T @ Cov @ B)
217
+ return vals[0], vals[1], float(np.arctan2(vecs[1, 0], vecs[0, 0])), Jn
218
+
219
+
220
+ def wineland_xi2(st: State, n, weights=None, spec_n=None):
221
+ """Wineland parameter xi_R^2 (=1 for a coherent spin state); for weighted
222
+ collective spins the coherent-state normalization S2/S1^2 is used.
223
+ Returns (xi2, angle, var_min, var_max, |J|)."""
224
+ J, Cov, S1, S2 = collective_moments(st, n, weights, spec_n=spec_n)
225
+ if np.linalg.norm(J) == 0:
226
+ return np.inf, 0.0, np.nan, np.nan, 0.0
227
+ vmin, vmax, ang, Jn = transverse_variances(J, Cov)
228
+ return float(vmin * S1**2 / (Jn**2 * S2)), ang, float(vmin), float(vmax), float(Jn)
229
+
230
+
231
+ def coherence(st: State, n, spec_n=None) -> float:
232
+ """Ramsey contrast 2|<J_perp>|/N."""
233
+ J, _, S1, _ = collective_moments(st, n, spec_n=spec_n)
234
+ return float(2.0 * np.hypot(J[0], J[1]) / S1)
235
+
236
+
237
+ # ---------------------------------------------------------------------------
238
+ # Right-hand side (connected variables)
239
+ # ---------------------------------------------------------------------------
240
+
241
+ def _rhs(t, y, rt: Rates, feedback: bool = True):
242
+ M = rt.M
243
+ st = State.unpack(y, M)
244
+ s, z, Pc, Qc, Rc, Zc = st.s, st.z, st.Pc, st.Qc, st.Rc, st.Zc
245
+ G, n, delta = rt.G, rt.n, rt.delta
246
+ chi1, Gd, Gu, gphi = rt.chi1, rt.Gd, rt.Gu, rt.gamma_phi
247
+
248
+ c = 1j * chi1 + 0.5 * (Gd - Gu)
249
+ cc = np.conj(c)
250
+ A = 1j * (delta + chi1 * G**2) + gphi + 0.5 * (Gd + Gu) * G**2
251
+ Ac = np.conj(A)
252
+ W = G * n
253
+ sc = np.conj(s)
254
+ Rcc = np.conj(Rc)
255
+
256
+ # ---------------- first moments ----------------
257
+ # sum_{k != a} G_k <sigma z_a sigma+_k> = sum_p G_p (n_p - d_pm)(Rc_mp + z_m s_p)
258
+ Ws = W @ s
259
+ fb = 1.0 if feedback else 0.0
260
+ S_zp = fb * (Rc @ W - G * np.diag(Rc)) + z * (Ws - G * s)
261
+ ds = -Ac * s + cc * G * S_zp
262
+ # sum_{k != a} G_k <sigma+_a sigma-_k> = sum_p G_p (n_p - d_pm)(Pc_mp + s_m s*_p)
263
+ S_pm = fb * (Pc @ W - G * np.diag(Pc)) + s * np.conj(Ws - G * s)
264
+ dz = -4.0 * G * np.real(c * S_pm) - G**2 * ((Gd - Gu) + (Gd + Gu) * z)
265
+
266
+ # ---------------- helpers ----------------
267
+ Gm, Gn = G[:, None], G[None, :]
268
+ zm, zn = z[:, None], z[None, :]
269
+ sm, sn = s[:, None], s[None, :]
270
+ scm, scn = sc[:, None], sc[None, :]
271
+ GG = Gm * Gn
272
+
273
+ def wsum(full, at_m, at_n):
274
+ """sum_p G_p (n_p - d_pm - d_pn) T(p) from the full sum and the p = m, p = n values."""
275
+ return full - Gm * at_m - Gn * at_n
276
+
277
+ # ---------------- Pc ----------------
278
+ # (iii) from d sigma+_a: sum_p W_p [ z_m Pc_pn + s_p Rc*_mn ]
279
+ S1 = wsum(zm * (W @ Pc)[None, :] + Rcc * Ws,
280
+ zm * Pc + Rcc * sm, # p = m: z_m Pc_mn + s_m Rc*_mn
281
+ zm * np.diag(Pc)[None, :] + Rcc * sn) # p = n: z_m Pc_nn + s_n Rc*_mn
282
+ # (iii) from d sigma-_b: sum_p W_p [ z_n Pc_mp + s*_p Rc_nm ]
283
+ S2 = wsum(zn * (Pc @ W)[:, None] + Rc.T * np.conj(Ws),
284
+ zn * np.diag(Pc)[:, None] + Rc.T * scm,
285
+ zn * Pc + Rc.T * scn)
286
+ src_a = 0.5 * (zm + zm * zn + Zc) - Rc * scn - zm * (sn * scn)
287
+ src_b = 0.5 * (zn + zm * zn + Zc) - sm * np.conj(Rc.T) - zn * (sm * scm)
288
+ dPc = (
289
+ -(Ac[:, None] + A[None, :]) * Pc
290
+ + cc * Gm * (Gn * src_a + S1)
291
+ + c * Gn * (Gm * src_b + S2)
292
+ + Gu * GG * (Zc + zm * zn)
293
+ )
294
+
295
+ # ---------------- Qc ----------------
296
+ S3 = wsum(zm * (W @ Qc)[None, :] + Rc * Ws,
297
+ zm * Qc + Rc * sm,
298
+ zm * np.diag(Qc)[None, :] + Rc * sn)
299
+ S4 = wsum(zn * (Qc @ W)[:, None] + Rc.T * Ws,
300
+ zn * np.diag(Qc)[:, None] + Rc.T * sm,
301
+ zn * Qc + Rc.T * sn)
302
+ dQc = (
303
+ -(Ac[:, None] + Ac[None, :]) * Qc
304
+ + cc * Gm * (-Gn * (Rc + zm * sn) * sn + S3)
305
+ + cc * Gn * (-Gm * sm * (Rc.T + zn * sm) + S4)
306
+ )
307
+
308
+ # ---------------- Rc ----------------
309
+ # (iii) from d sigma z_a: sum_p W_p [ c (s_m Pc_np + s*_p Qc_mn) + c* (s_p Pc_nm + s*_m Qc_pn) ]
310
+ S5 = wsum(c * (sm * (Pc @ W)[None, :] + Qc * np.conj(Ws)) + cc * (Pc.T * Ws + scm * (W @ Qc)[None, :]),
311
+ c * (sm * Pc.T + Qc * scm) + cc * (Pc.T * sm + scm * Qc),
312
+ c * (sm * np.diag(Pc)[None, :] + Qc * scn) + cc * (Pc.T * sn + scm * np.diag(Qc)[None, :]))
313
+ # (iii) from d sigma+_b: sum_p W_p [ z_n Rc_mp + s_p Zc_mn ]
314
+ S7 = wsum(zn * (Rc @ W)[:, None] + Zc * Ws,
315
+ zn * np.diag(Rc)[:, None] + Zc * sm,
316
+ zn * Rc + Zc * sn)
317
+ Pmn_raw = Pc + sm * scn
318
+ Pnm_raw = Pc.T + sn * scm
319
+ src_za = c * (0.5 * (sm - Rc.T - sm * zn) - Pmn_raw * sn) + cc * (-Pnm_raw * sn)
320
+ dRc = (
321
+ -Gm**2 * (Gd + Gu) * Rc
322
+ - 2.0 * GG * src_za
323
+ - 2.0 * Gm * S5
324
+ - Ac[None, :] * Rc
325
+ + cc * GG * (1.0 - zm) * (Rc.T + sm * zn)
326
+ + cc * Gn * S7
327
+ - 2.0 * Gd * GG * (Rc.T + sm * zn)
328
+ )
329
+
330
+ # ---------------- Zc ----------------
331
+ # (iii) from d sigma z_a: sum_p W_p [ c (s_m Rc*_np + s*_p Rc_nm) + c* (s_p Rc*_nm + s*_m Rc_np) ]
332
+ S8 = wsum(c * (sm * (Rcc @ W)[None, :] + Rc.T * np.conj(Ws)) + cc * (Rcc.T * Ws + scm * (Rc @ W)[None, :]),
333
+ c * (sm * Rcc.T + Rc.T * scm) + cc * (Rcc.T * sm + scm * Rc.T),
334
+ c * (sm * np.diag(Rcc)[None, :] + Rc.T * scn) + cc * (Rcc.T * sn + scm * np.diag(Rc)[None, :]))
335
+ # (iii) from d sigma z_b: sum_p W_p [ c (s_n Rc*_mp + s*_p Rc_mn) + c* (s_p Rc*_mn + s*_n Rc_mp) ]
336
+ S9 = wsum(c * (sn * (Rcc @ W)[:, None] + Rc * np.conj(Ws)) + cc * (Rcc * Ws + scn * (Rc @ W)[:, None]),
337
+ c * (sn * np.diag(Rcc)[:, None] + Rc * scm) + cc * (Rcc * sm + scn * np.diag(Rc)[:, None]),
338
+ c * (sn * Rcc + Rc * scn) + cc * (Rcc * sn + scn * Rc))
339
+ src_zz_a = c * Pmn_raw * (1.0 - zn) - cc * Pnm_raw * (1.0 + zn)
340
+ src_zz_b = -c * Pnm_raw * (1.0 + zm) + cc * Pmn_raw * (1.0 - zm)
341
+ dZc = (
342
+ -(Gm**2 + Gn**2) * (Gd + Gu) * Zc
343
+ - 2.0 * GG * (src_zz_a + src_zz_b)
344
+ - 2.0 * Gm * S8
345
+ - 2.0 * Gn * S9
346
+ + 4.0 * GG * (Gd * Pmn_raw + Gu * Pnm_raw)
347
+ )
348
+
349
+ if getattr(rt, "meas", 0.0):
350
+ dPc, dQc, dRc, dZc = _add_measurement(rt.meas, getattr(rt, "meas_eta", 1.0), n, st, dPc, dQc, dRc, dZc)
351
+ return State(ds, dz, dPc, dQc, dRc, dZc).pack()
352
+
353
+
354
+ def _add_measurement(Gm, eta, n, st, dPc, dQc, dRc, dZc):
355
+ """Conditioning on a continuous quantum non-demolition measurement of
356
+ J_z = (1/2) sum_i sigma z_i at rate Gm (Gaussian, Kalman form): the pair
357
+ covariance of every two spins changes as
358
+ d Cov(sigma^a_i, sigma^b_j) = -Gm Cov(sigma^a_i, J_z) Cov(sigma^b_j, J_z) dt,
359
+ with Cov(sigma^a_i, J_z) = (1/2)[delta_az - v^a_i v^z_i + sum_{j != i} C^{az}_{ij}].
360
+ For an ensemble without dynamics this gives dV/dt = -Gm V^2 for V = Var(J_z).
361
+ The back-action of the probe (photon-number fluctuations) rotates all spins
362
+ about z by a common random angle with d Var(angle)/dt = 2 Gamma_phi,
363
+ Gamma_phi = Gm/(8 eta), which adds 2 Gamma_phi (e_z x v_i)^a (e_z x v_j)^b to the
364
+ pair covariance; with eta = 1 the product Var(J_y) Var(J_z) stays minimal.
365
+ Spectator spins are not conditioned (use a discretization without spectators)."""
366
+ v, C = to_cartesian(st)
367
+ dphi = np.stack([-v[:, 1], v[:, 0], np.zeros(st.M)], axis=1) # d v / d(angle) for a rotation about z
368
+ ez = np.array([0.0, 0.0, 1.0])
369
+ same = ez[None, :] - v * v[:, [2]] # (M,3): delta_az - v^a v^z
370
+ Cz = C[:, :, :, 2] # (M,M,3): Cov(sigma^a_m, sigma^z_n)
371
+ pair = np.einsum("n,mna->ma", n, Cz) - np.einsum("mma->ma", Cz)
372
+ cvec = 0.5 * (same + pair) # (M,3) complex
373
+ dC = -Gm * np.einsum("ma,nb->mnab", cvec, cvec) + (Gm / (4.0 * eta)) * np.einsum("ma,nb->mnab", dphi, dphi)
374
+ extra = from_cartesian(np.zeros((st.M, 3)), dC)
375
+ return dPc + extra.Pc, dQc + extra.Qc, dRc + extra.Rc, dZc + extra.Zc
376
+
377
+
378
+ def _rhs_meanfield(t, y, rt: Rates):
379
+ M = rt.M
380
+ s, z = y[:M], y[M:]
381
+ G, n, delta = rt.G, rt.n, rt.delta
382
+ chi1, Gd, Gu, gphi = rt.chi1, rt.Gd, rt.Gu, rt.gamma_phi
383
+ c = 1j * chi1 + 0.5 * (Gd - Gu)
384
+ A = 1j * (delta + chi1 * G**2) + gphi + 0.5 * (Gd + Gu) * G**2
385
+ W = G * n
386
+ Ws = W @ s
387
+ ds = -np.conj(A) * s + np.conj(c) * G * z * (Ws - G * s)
388
+ dz = -4.0 * G * np.real(c * s * np.conj(Ws - G * s)) - G**2 * ((Gd - Gu) + (Gd + Gu) * z)
389
+ return np.concatenate([ds, dz])
390
+
391
+
392
+ def evolve(st: State, rt: Rates, t: float, t_eval=None, rtol=1e-8, atol=None, method="DOP853", feedback=True):
393
+ """Evolve for time t. Returns the final State (or a list at t_eval)."""
394
+ if t == 0:
395
+ return st.copy() if t_eval is None else [st.copy() for _ in t_eval]
396
+ K = getattr(rt, "K", 0)
397
+ if st.vs.shape[0] != K:
398
+ # spectators start in the same product state as the classes
399
+ v0 = np.array([2 * st.s[0].real, 2 * st.s[0].imag, st.z[0].real])
400
+ st = State(st.s, st.z, st.Pc, st.Qc, st.Rc, st.Zc, np.tile(v0, (K, 1)))
401
+ y0 = st.pack()
402
+ if atol is None:
403
+ # connected correlations are O(1/N): scale the absolute tolerance accordingly
404
+ atol = 1e-10 / max(rt.N, 1.0)
405
+ sol = solve_ivp(_rhs, (0.0, t), y0, args=(rt, feedback), t_eval=t_eval, rtol=rtol, atol=atol, method=method)
406
+ if not sol.success:
407
+ raise RuntimeError(sol.message)
408
+ M = rt.M
409
+
410
+ def spec(tau):
411
+ # exact free evolution of the spectators: <sigma+> ~ exp(+i delta t) (same
412
+ # convention as the interacting classes) with single-spin dephasing
413
+ if st.vs.shape[0] == 0:
414
+ return st.vs
415
+ ph = np.asarray(rt.spec_delta) * tau
416
+ damp = np.exp(-rt.gamma_phi * tau)
417
+ x, y, z = st.vs[:, 0], st.vs[:, 1], st.vs[:, 2]
418
+ xn = damp * (x * np.cos(ph) - y * np.sin(ph))
419
+ yn = damp * (y * np.cos(ph) + x * np.sin(ph))
420
+ return np.stack([xn, yn, z], axis=1)
421
+
422
+ if t_eval is None:
423
+ return State.unpack(sol.y[:, -1], M, spec(t))
424
+ return [State.unpack(sol.y[:, k], M, spec(sol.t[k])) for k in range(sol.y.shape[1])]
425
+
426
+
427
+ def evolve_meanfield(s, z, rt: Rates, t: float, t_eval=None, rtol=1e-8, atol=1e-11, method="DOP853"):
428
+ y0 = np.concatenate([np.asarray(s, complex), np.asarray(z, complex)])
429
+ sol = solve_ivp(_rhs_meanfield, (0.0, t), y0, args=(rt,), t_eval=t_eval, rtol=rtol, atol=atol, method=method)
430
+ if not sol.success:
431
+ raise RuntimeError(sol.message)
432
+ M = rt.M
433
+ if t_eval is None:
434
+ return sol.y[:M, -1], sol.y[M:, -1]
435
+ return sol.y[:M, :], sol.y[M:, :]