torchspin 0.3.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.
Files changed (116) hide show
  1. torchspin/__init__.py +363 -0
  2. torchspin/_cardamom_istos.py +327 -0
  3. torchspin/_cardamom_propagatedm.py +413 -0
  4. torchspin/_cardamom_utils.py +113 -0
  5. torchspin/_compile.py +100 -0
  6. torchspin/_linalg.py +140 -0
  7. torchspin/angmom.py +742 -0
  8. torchspin/autograd.py +557 -0
  9. torchspin/autoguess.py +342 -0
  10. torchspin/batch.py +225 -0
  11. torchspin/blochsteady.py +274 -0
  12. torchspin/cardamom.py +552 -0
  13. torchspin/chili.py +932 -0
  14. torchspin/chili_sle.py +1233 -0
  15. torchspin/constants.py +88 -0
  16. torchspin/convspec.py +148 -0
  17. torchspin/ctafft.py +91 -0
  18. torchspin/curry.py +241 -0
  19. torchspin/data/FourierSeriesCoefficients.txt +89 -0
  20. torchspin/data/GaussianCascadeCoefficients.txt +56 -0
  21. torchspin/data/isotopedata.txt +416 -0
  22. torchspin/data/spacegroups.txt +637 -0
  23. torchspin/dataproc.py +361 -0
  24. torchspin/dipbackground.py +69 -0
  25. torchspin/dipkernel.py +102 -0
  26. torchspin/diptensor.py +138 -0
  27. torchspin/endorfrq.py +198 -0
  28. torchspin/endorfrq_perturb.py +302 -0
  29. torchspin/eprload.py +1538 -0
  30. torchspin/eprsave.py +165 -0
  31. torchspin/esfit.py +2427 -0
  32. torchspin/evolve.py +274 -0
  33. torchspin/ewrls.py +124 -0
  34. torchspin/excitation.py +107 -0
  35. torchspin/exciteprofile.py +132 -0
  36. torchspin/experiment.py +323 -0
  37. torchspin/exponfit.py +177 -0
  38. torchspin/fastmotion.py +182 -0
  39. torchspin/fdaxis.py +57 -0
  40. torchspin/fitgui.py +350 -0
  41. torchspin/garlic.py +760 -0
  42. torchspin/ham.py +145 -0
  43. torchspin/ham_cf.py +115 -0
  44. torchspin/ham_ee.py +124 -0
  45. torchspin/ham_ez.py +119 -0
  46. torchspin/ham_ezho.py +403 -0
  47. torchspin/ham_hf.py +122 -0
  48. torchspin/ham_nn.py +108 -0
  49. torchspin/ham_nq.py +100 -0
  50. torchspin/ham_nz.py +136 -0
  51. torchspin/ham_oz.py +115 -0
  52. torchspin/ham_so.py +118 -0
  53. torchspin/ham_zf.py +254 -0
  54. torchspin/hamsymm.py +723 -0
  55. torchspin/initstate.py +158 -0
  56. torchspin/isotopologues.py +296 -0
  57. torchspin/levels.py +189 -0
  58. torchspin/levelsplot.py +363 -0
  59. torchspin/lineshape.py +715 -0
  60. torchspin/lpsvd.py +219 -0
  61. torchspin/makespec.py +87 -0
  62. torchspin/mdhmm.py +573 -0
  63. torchspin/mdload.py +788 -0
  64. torchspin/mdtraj2oripot.py +90 -0
  65. torchspin/ml.py +377 -0
  66. torchspin/mlpsvd.py +297 -0
  67. torchspin/nucdata.py +310 -0
  68. torchspin/nucfrq2d.py +228 -0
  69. torchspin/orca2torchspin.py +1198 -0
  70. torchspin/ordering.py +78 -0
  71. torchspin/oripotentialplot.py +246 -0
  72. torchspin/orisel.py +176 -0
  73. torchspin/pepper.py +2019 -0
  74. torchspin/pepper_autograd.py +613 -0
  75. torchspin/photoselect.py +174 -0
  76. torchspin/plegendre.py +131 -0
  77. torchspin/propint.py +127 -0
  78. torchspin/pulse.py +652 -0
  79. torchspin/py.typed +0 -0
  80. torchspin/rapidscan2spc.py +101 -0
  81. torchspin/resfields.py +337 -0
  82. torchspin/resfields_batch.py +446 -0
  83. torchspin/resfields_eig.py +250 -0
  84. torchspin/resfields_perturb.py +706 -0
  85. torchspin/resfreqs_matrix.py +344 -0
  86. torchspin/resfreqs_perturb.py +299 -0
  87. torchspin/resonator.py +310 -0
  88. torchspin/resonatorprofile.py +85 -0
  89. torchspin/rfmixer.py +147 -0
  90. torchspin/rotations.py +87 -0
  91. torchspin/rotutils.py +880 -0
  92. torchspin/saffron.py +1523 -0
  93. torchspin/saffron_pathways.py +92 -0
  94. torchspin/saffron_peaks.py +580 -0
  95. torchspin/saffron_thyme.py +472 -0
  96. torchspin/salt.py +680 -0
  97. torchspin/sigeq.py +71 -0
  98. torchspin/signalprocessing.py +226 -0
  99. torchspin/sitetransforms.py +254 -0
  100. torchspin/sphgrid.py +470 -0
  101. torchspin/spidyan.py +1508 -0
  102. torchspin/spinladder.py +161 -0
  103. torchspin/spinops.py +265 -0
  104. torchspin/spinsystem.py +1427 -0
  105. torchspin/stackplot.py +167 -0
  106. torchspin/stev.py +283 -0
  107. torchspin/stochtraj_diffusion.py +494 -0
  108. torchspin/stochtraj_jump.py +195 -0
  109. torchspin/strainwidth.py +727 -0
  110. torchspin/transmitter.py +89 -0
  111. torchspin/utils.py +605 -0
  112. torchspin-0.3.0.dist-info/METADATA +267 -0
  113. torchspin-0.3.0.dist-info/RECORD +116 -0
  114. torchspin-0.3.0.dist-info/WHEEL +5 -0
  115. torchspin-0.3.0.dist-info/licenses/LICENSE.md +22 -0
  116. torchspin-0.3.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,413 @@
1
+ """Density matrix propagation for cardamom.
2
+
3
+ Implements two propagation methods:
4
+
5
+ 1. **fast** (Sezer et al., JCP 128, 165106, 2008):
6
+ Operates in the m_S = -1/2 subspace only (S = 1/2).
7
+ Without hyperfine: scalar propagator ``U = exp(-i dt omega Gp_zz / 2)``.
8
+ With 14N (I=1): 3x3 matrix exponential from axis-angle decomposition.
9
+
10
+ 2. **ISTOs** (Oganesyan, PCCP 13, 4724, 2011):
11
+ Full Hilbert-space propagation using irreducible spherical tensor
12
+ operators and rank-2 Wigner D-matrices from quaternion trajectories.
13
+ """
14
+ from __future__ import annotations
15
+
16
+ import math
17
+
18
+ import numpy as np
19
+ from torchspin._cardamom_utils import tensor_traj
20
+ from torchspin.constants import GFREE
21
+
22
+
23
+ # ---------------------------------------------------------------------------
24
+ # Fast method (Sezer 2008)
25
+ # ---------------------------------------------------------------------------
26
+
27
+ def propagate_fast(
28
+ g: np.ndarray,
29
+ RTraj: np.ndarray,
30
+ omega: float,
31
+ dtSpin: float,
32
+ nSteps: int,
33
+ nTraj: int,
34
+ *,
35
+ A: np.ndarray | None = None,
36
+ RLab: np.ndarray | None = None,
37
+ groups: int = 1,
38
+ ) -> np.ndarray:
39
+ """Propagate density matrix using the fast m_S=-1/2 subspace method.
40
+
41
+ With ``groups`` > 1 the trajectories are ``groups`` consecutive blocks of
42
+ ``nTraj // groups`` (one powder orientation each) and the signal is
43
+ averaged per block, returning ``(groups, nSteps)``.
44
+
45
+ Parameters
46
+ ----------
47
+ g:
48
+ g-tensor principal values, shape ``(3,)``.
49
+ RTraj:
50
+ Rotation matrix trajectory, shape ``(3, 3, nSteps_spatial, nTraj)``.
51
+ The spatial trajectory (may be longer than nSteps).
52
+ omega:
53
+ Microwave angular frequency in rad/s (= 2*pi*mwFreq_Hz).
54
+ dtSpin:
55
+ Spin propagation time step in seconds.
56
+ nSteps:
57
+ Number of spin propagation steps.
58
+ nTraj:
59
+ Number of trajectories.
60
+ A:
61
+ Optional hyperfine tensor principal values (MHz), shape ``(3,)``.
62
+ If provided, includes nuclear spin dynamics (I=1 for 14N).
63
+ RLab:
64
+ Optional lab-frame rotation matrices, shape ``(3, 3, nSteps, nTraj)``.
65
+ Used to combine local and global dynamics.
66
+
67
+ Returns
68
+ -------
69
+ Sprho:
70
+ Density matrix trace signal, shape ``(nSteps,)``.
71
+ This is the trajectory-averaged expectation value sum_k rho_kk(t).
72
+ """
73
+ import torch
74
+
75
+ # --- Compute tensor trajectories ---
76
+ g_t = torch.tensor(g, dtype=torch.float64)
77
+ RTraj_t = torch.tensor(RTraj[:, :, :nSteps, :], dtype=torch.float64)
78
+
79
+ gTensor = tensor_traj(g_t, RTraj_t).numpy() # (3,3,nSteps,nTraj)
80
+
81
+ includeHF = A is not None
82
+ if includeHF:
83
+ A_t = torch.tensor(A, dtype=torch.float64)
84
+ ATensor = tensor_traj(A_t, RTraj_t).numpy()
85
+ # MHz → rad/s
86
+ ATensor = ATensor * 1e6 * 2.0 * np.pi
87
+
88
+ # --- Combine with lab-frame rotation if provided ---
89
+ if RLab is not None:
90
+ gTensor = _rotate_tensor_lab(gTensor, RLab[:, :, :nSteps, :])
91
+ if includeHF:
92
+ ATensor = _rotate_tensor_lab(ATensor, RLab[:, :, :nSteps, :])
93
+
94
+ # --- Compute propagators ---
95
+ gIso = np.sum(g) / 3.0
96
+ GpTensor = (gTensor - gIso) / GFREE
97
+ Gp_zz = torch.from_numpy(np.ascontiguousarray(GpTensor[2, 2, :, :])) # (nSteps, nTraj)
98
+ phase = torch.exp(-1j * dtSpin * 0.5 * omega * Gp_zz) # (nSteps, nTraj)
99
+
100
+ # --- Propagate density matrix (torch: multi-threaded elementwise ops and
101
+ # batched 3×3 matmuls; the recursion is sequential in time only) ---
102
+ if includeHF:
103
+ U = _build_hf_propagator(ATensor, phase, dtSpin) # (nSteps, nTraj, 3, 3) complex
104
+ rho = torch.zeros(nTraj, 3, 3, dtype=torch.complex128)
105
+ rho[:, 0, 0] = rho[:, 1, 1] = rho[:, 2, 2] = 0.5
106
+ tr = torch.empty(nSteps, nTraj, dtype=torch.complex128)
107
+ tr[0] = 1.5
108
+ for iStep in range(1, nSteps):
109
+ U_prev = U[iStep - 1]
110
+ # rho(t+1) = U @ rho(t) @ U (not U†, per Sezer fast method)
111
+ rho = U_prev @ rho @ U_prev
112
+ tr[iStep] = rho[:, 0, 0] + rho[:, 1, 1] + rho[:, 2, 2]
113
+ signal = tr.reshape(nSteps, groups, -1).mean(dim=2).T.numpy() # (groups, nSteps)
114
+ else:
115
+ # rho is scalar: rho(t) = 0.5 * prod_{k<t} U_k^2
116
+ U2 = phase ** 2
117
+ rho = torch.ones(nSteps, nTraj, dtype=torch.complex128) * 0.5
118
+ rho[1:] = 0.5 * torch.cumprod(U2[:-1], dim=0)
119
+ signal = rho.reshape(nSteps, groups, -1).mean(dim=2).T.numpy()
120
+
121
+ return signal if groups > 1 else signal[0]
122
+
123
+
124
+ def _build_hf_propagator(ATensor: np.ndarray, phase, dtSpin: float):
125
+ """Build the I=1 hyperfine propagator (Eqs. 35, 37, A1-A2 in Sezer 2008).
126
+
127
+ Parameters
128
+ ----------
129
+ ATensor:
130
+ Hyperfine tensor trajectory in rad/s, shape ``(3, 3, nSteps, nTraj)``.
131
+ phase:
132
+ Zeeman phase factor ``exp(-i dt/2 ω Gp_zz)``, torch ``(nSteps, nTraj)``.
133
+ dtSpin:
134
+ Spin propagation time step (s).
135
+
136
+ Returns
137
+ -------
138
+ U:
139
+ Propagator, torch complex128 ``(nSteps, nTraj, 3, 3)``.
140
+ """
141
+ import torch
142
+ Az = torch.from_numpy(np.ascontiguousarray(ATensor[:, 2, :, :])) # (3, nSteps, nTraj): A-tensor z column
143
+ a = torch.sqrt((Az ** 2).sum(dim=0)) # Eq. 24
144
+ theta = dtSpin * 0.5 * a
145
+ a_safe = torch.where(a < 1e-30, torch.full_like(a, 1e-30), a)
146
+ nx, ny, nz = Az[0] / a_safe, Az[1] / a_safe, Az[2] / a_safe
147
+ ct = torch.cos(theta) - 1.0
148
+ st = -torch.sin(theta)
149
+ s2 = math.sqrt(0.5)
150
+ nxx, nyy, nzz = nx * nx, ny * ny, nz * nz
151
+ nxy, nzx, nzy = nx * ny, nz * nx, nz * ny
152
+ U = torch.empty(a.shape + (3, 3), dtype=torch.complex128)
153
+ U[..., 0, 0] = torch.complex(1 + ct * (nzz + 0.5 * (nxx + nyy)), st * nz)
154
+ U[..., 0, 1] = torch.complex(s2 * (st * ny + ct * nzx), s2 * (st * nx - ct * nzy))
155
+ U[..., 0, 2] = torch.complex(0.5 * ct * (nxx - nyy), -ct * nxy)
156
+ U[..., 1, 0] = torch.complex(s2 * (-st * ny + ct * nzx), s2 * (st * nx + ct * nzy))
157
+ U[..., 1, 1] = torch.complex(1 + ct * (nxx + nyy), torch.zeros_like(a))
158
+ U[..., 1, 2] = torch.complex(s2 * (st * ny - ct * nzx), s2 * (st * nx + ct * nzy))
159
+ U[..., 2, 0] = torch.complex(0.5 * ct * (nxx - nyy), ct * nxy)
160
+ U[..., 2, 1] = torch.complex(s2 * (-st * ny - ct * nzx), s2 * (st * nx - ct * nzy))
161
+ U[..., 2, 2] = torch.complex(1 + ct * (nzz + 0.5 * (nxx + nyy)), -st * nz)
162
+ # Full propagator: Eq. 35 — U = exp(-i dt/2 ω Gp_zz) * expadotI
163
+ return U * phase[..., None, None]
164
+
165
+
166
+ def _rotate_tensor_lab(T: np.ndarray, RLab: np.ndarray) -> np.ndarray:
167
+ """Rotate tensor trajectory into lab frame: RLab @ T @ RLab^T.
168
+
169
+ Parameters
170
+ ----------
171
+ T:
172
+ Tensor trajectory, shape ``(3, 3, nSteps, nTraj)``.
173
+ RLab:
174
+ Lab-frame rotation matrices, shape ``(3, 3, nSteps, nTraj)``.
175
+
176
+ Returns
177
+ -------
178
+ T_lab:
179
+ Rotated tensor, shape ``(3, 3, nSteps, nTraj)``.
180
+ """
181
+ # R T Rᵀ for every (step, trajectory): batched matmul in (N, 3, 3) layout on
182
+ # torch (multi-threaded; numpy's small-matrix matmul/einsum take ~45 ms per
183
+ # 200k products, torch ~1 ms) — PERF_PLAN §2.4
184
+ import torch
185
+ Rb = torch.from_numpy(np.ascontiguousarray(np.transpose(RLab, (2, 3, 0, 1)))) # (nSteps, nTraj, 3, 3)
186
+ Tb = torch.from_numpy(np.ascontiguousarray(np.transpose(T, (2, 3, 0, 1))))
187
+ out = (Rb @ Tb @ Rb.transpose(-1, -2)).numpy()
188
+ return np.transpose(out, (2, 3, 0, 1))
189
+
190
+
191
+ # ---------------------------------------------------------------------------
192
+ # ISTOs method (Oganesyan 2011)
193
+ # ---------------------------------------------------------------------------
194
+
195
+ def _block_average_d2(D2Traj, block_length):
196
+ """Block-average D2 Wigner matrix trajectory.
197
+
198
+ Parameters
199
+ ----------
200
+ D2Traj : ndarray, shape (5, 5, nSteps, nTraj)
201
+ block_length : int
202
+ Number of frames per block.
203
+
204
+ Returns
205
+ -------
206
+ D2_avg : ndarray, shape (5, 5, nBlocks, nTraj)
207
+ """
208
+ if block_length <= 1:
209
+ return D2Traj
210
+ nSteps = D2Traj.shape[2]
211
+ nBlocks = nSteps // block_length
212
+ if nBlocks == 0:
213
+ return D2Traj
214
+ D2_trim = D2Traj[:, :, :nBlocks * block_length, :]
215
+ D2_reshaped = D2_trim.reshape(5, 5, nBlocks, block_length, -1)
216
+ return D2_reshaped.mean(axis=3)
217
+
218
+
219
+ def _sliding_window_d2(D2Traj, window_length, stride=1):
220
+ """Extract overlapping windows from D2 trajectory.
221
+
222
+ Each window position becomes a virtual trajectory, improving
223
+ statistics from a single long MD trajectory.
224
+
225
+ Parameters
226
+ ----------
227
+ D2Traj : ndarray, shape (5, 5, nSteps, nTraj)
228
+ window_length : int
229
+ stride : int
230
+
231
+ Returns
232
+ -------
233
+ D2_windowed : ndarray, shape (5, 5, window_length, nWindows * nTraj)
234
+ """
235
+ nSteps = D2Traj.shape[2]
236
+ nTraj = D2Traj.shape[3]
237
+ if window_length >= nSteps:
238
+ return D2Traj[:, :, :window_length, :]
239
+ starts = range(0, nSteps - window_length + 1, stride)
240
+ windows = []
241
+ for s in starts:
242
+ windows.append(D2Traj[:, :, s:s + window_length, :])
243
+ return np.concatenate(windows, axis=3) # (5, 5, window_length, nWindows*nTraj)
244
+
245
+
246
+ def propagate_istos(
247
+ sys_spins: list[float],
248
+ g: np.ndarray,
249
+ A: np.ndarray | None,
250
+ qTraj: np.ndarray,
251
+ omega: float,
252
+ dtSpin: float,
253
+ nSteps: int,
254
+ nTraj: int,
255
+ CenterField: float,
256
+ *,
257
+ D: np.ndarray | None = None,
258
+ gn: list[float] | None = None,
259
+ nuc_spins: list[float] | None = None,
260
+ qLab: np.ndarray | None = None,
261
+ block_length: int = 1,
262
+ d2_cache: dict | None = None,
263
+ detect_Sz: bool = False,
264
+ ) -> np.ndarray:
265
+ """Propagate density matrix using the ISTOs method (Oganesyan 2011).
266
+
267
+ Full Hilbert-space propagation using irreducible spherical tensor
268
+ operators and rank-2 Wigner D-matrices from quaternion trajectories.
269
+
270
+ Parameters
271
+ ----------
272
+ sys_spins:
273
+ Spin quantum numbers, e.g. [0.5] or [0.5, 1.0].
274
+ g:
275
+ g-tensor principal values, shape ``(3,)`` or ``(nElectrons, 3)``.
276
+ A:
277
+ Hyperfine tensor principal values in MHz, shape ``(nNuclei, 3)``
278
+ or ``(3,)`` or ``None``.
279
+ qTraj:
280
+ Quaternion trajectory, shape ``(4, nSteps_spatial, nTraj)``.
281
+ omega:
282
+ Microwave angular frequency (rad/s).
283
+ dtSpin:
284
+ Spin propagation time step (s).
285
+ nSteps:
286
+ Number of spin propagation steps.
287
+ nTraj:
288
+ Number of trajectories.
289
+ CenterField:
290
+ Center magnetic field in mT.
291
+ D:
292
+ ZFS tensor principal values in MHz, shape ``(nElectrons, 3)`` or ``None``.
293
+ gn:
294
+ Nuclear g-values.
295
+ nuc_spins:
296
+ Nuclear spin quantum numbers.
297
+ qLab:
298
+ Lab-frame quaternion trajectory, shape ``(4, nSteps, nTraj)``.
299
+
300
+ Returns
301
+ -------
302
+ Sprho:
303
+ Trace of S+ * rho(t), averaged over trajectories, shape ``(nSteps,)``.
304
+ """
305
+ from scipy.linalg import expm
306
+ from torchspin._cardamom_istos import magint, wigD
307
+ from torchspin.spinops import sop
308
+
309
+ # --- Compute IST decomposition ---
310
+ g = np.atleast_2d(np.asarray(g, dtype=float))
311
+ A_2d = np.atleast_2d(np.asarray(A, dtype=float)) if A is not None else None
312
+
313
+ T, F = magint(
314
+ sys_spins, g, CenterField,
315
+ A=A_2d, D=D, gn=gn, nuc_spins=nuc_spins,
316
+ include_nuc_zeeman=False,
317
+ )
318
+
319
+ F0 = F['F0'] * 2 * np.pi # Hz → rad/s
320
+ F2 = F['F2'] * 2 * np.pi
321
+ T0 = T['T0']
322
+ T2 = T['T2']
323
+
324
+ nInt = len(T0)
325
+ nStates = int(np.round(np.prod([2 * s + 1 for s in sys_spins])))
326
+
327
+ # --- Build Q0 (zeroth rank, time-independent) ---
328
+ Q0 = np.zeros((nStates, nStates), dtype=complex)
329
+ for k in range(nInt):
330
+ Q0 += np.conj(F0[k]) * T0[k]
331
+
332
+ # --- H0: isotropic Zeeman part (for interaction frame) ---
333
+ H0 = np.conj(F0[0]) * T0[0] # first interaction is electron Zeeman
334
+
335
+ # --- Build Q2 (second rank, 5x5 rotational basis operators) ---
336
+ Q2 = np.zeros((5, 5, nStates, nStates), dtype=complex)
337
+ for mp in range(5):
338
+ for m in range(5):
339
+ for iInt in range(nInt):
340
+ Q2[mp, m] += np.conj(F2[iInt, mp]) * T2[iInt][m]
341
+
342
+ # --- Compute D2 trajectories from quaternions (with optional caching) ---
343
+ cache_key = id(qTraj)
344
+ if d2_cache is not None and cache_key in d2_cache:
345
+ D2Traj = d2_cache[cache_key]
346
+ else:
347
+ D2Traj = wigD(qTraj[:, :nSteps, :]) # (5, 5, nSteps, nTraj)
348
+ if d2_cache is not None:
349
+ d2_cache[cache_key] = D2Traj
350
+
351
+ # --- Block averaging (optional) ---
352
+ if block_length > 1:
353
+ D2Traj = _block_average_d2(D2Traj, block_length)
354
+ nSteps = D2Traj.shape[2] # update after averaging
355
+
356
+ # --- Combine local and global dynamics ---
357
+ if qLab is not None:
358
+ D2Lab = wigD(qLab[:, :nSteps, :])
359
+ # Matrix multiply: D2Traj = D2Lab @ D2Traj per step/traj
360
+ D2Combined = np.zeros_like(D2Traj)
361
+ for iStep in range(nSteps):
362
+ for iTraj in range(nTraj):
363
+ D2Combined[:, :, iStep, iTraj] = (
364
+ D2Lab[:, :, iStep, iTraj] @ D2Traj[:, :, iStep, iTraj]
365
+ )
366
+ D2Traj = D2Combined
367
+
368
+ # --- Build Hamiltonians H(t) = Q0 + sum_{mp,m} D2(m,mp,t) * Q2{mp,m} ---
369
+ H = np.tile(Q0[:, :, np.newaxis, np.newaxis], (1, 1, nSteps, nTraj))
370
+ for mp in range(5):
371
+ for m in range(5):
372
+ # D2Traj[m, mp, :, :] is (nSteps, nTraj)
373
+ # Q2[mp, m] is (nStates, nStates)
374
+ H += D2Traj[m, mp, :, :][np.newaxis, np.newaxis, :, :] * \
375
+ Q2[mp, m, :, :, np.newaxis, np.newaxis]
376
+
377
+ # --- Build propagators ---
378
+ # Interaction frame: U = expm(-i*dt*H0) * expm(+i*dt*H)
379
+ U0 = expm(-1j * dtSpin * H0)
380
+
381
+ U = np.zeros((nStates, nStates, nSteps, nTraj), dtype=complex)
382
+ for iStep in range(nSteps):
383
+ for iTraj in range(nTraj):
384
+ U_step = expm(1j * dtSpin * H[:, :, iStep, iTraj])
385
+ U[:, :, iStep, iTraj] = U0 @ U_step
386
+
387
+ # --- Initial state: rho(0) = Sx (after pi/2 pulse) ---
388
+ Sx = sop(sys_spins, [1, 1]).numpy()
389
+ rho = np.zeros((nStates, nStates, nSteps, nTraj), dtype=complex)
390
+ rho[:, :, 0, :] = Sx[:, :, np.newaxis]
391
+
392
+ # --- Propagate: rho(t+1) = U * rho(t) * U† ---
393
+ for iStep in range(1, nSteps):
394
+ for iTraj in range(nTraj):
395
+ U_prev = U[:, :, iStep - 1, iTraj]
396
+ U_adj = U_prev.conj().T
397
+ rho[:, :, iStep, iTraj] = (
398
+ U_prev @ rho[:, :, iStep - 1, iTraj] @ U_adj
399
+ )
400
+
401
+ # --- Average over trajectories ---
402
+ rho_avg = np.mean(rho, axis=3) # (nStates, nStates, nSteps)
403
+
404
+ # --- Apply detection operator ---
405
+ if detect_Sz:
406
+ Det = sop(sys_spins, [1, 3]).numpy() # Sz
407
+ else:
408
+ Det = sop(sys_spins, [1, 4]).numpy() # S+
409
+ Sprho = np.zeros(nSteps, dtype=complex)
410
+ for iStep in range(nSteps):
411
+ Sprho[iStep] = np.trace(Det @ rho_avg[:, :, iStep])
412
+
413
+ return Sprho
@@ -0,0 +1,113 @@
1
+ """Shared helpers for cardamom trajectory-based EPR simulation.
2
+
3
+ Ports of MATLAB EasySpin private helpers:
4
+ - ``cardamom_tensortraj.m`` → ``tensor_traj()``
5
+ - spiral grid generation → ``spiral_grid()``
6
+ - Gelman-Rubin convergence diagnostic → ``gelman_rubin()``
7
+ """
8
+ from __future__ import annotations
9
+
10
+ import numpy as np
11
+ import torch
12
+
13
+
14
+ def tensor_traj(
15
+ T_diag: torch.Tensor,
16
+ R: torch.Tensor,
17
+ ) -> torch.Tensor:
18
+ """Rotate an interaction tensor along a trajectory of rotation matrices.
19
+
20
+ Computes ``R @ diag(T) @ R^T`` for each time step and trajectory.
21
+
22
+ This is a direct port of ``cardamom_tensortraj.m``.
23
+
24
+ Parameters
25
+ ----------
26
+ T_diag:
27
+ Principal values of the interaction tensor, shape ``(3,)``.
28
+ R:
29
+ Rotation matrix trajectory, shape ``(3, 3, nSteps, nTraj)``.
30
+
31
+ Returns
32
+ -------
33
+ T_traj:
34
+ Rotated tensor trajectory, shape ``(3, 3, nSteps, nTraj)``.
35
+ """
36
+ # Build diagonal tensor
37
+ if T_diag.ndim == 1 and T_diag.shape[0] == 3:
38
+ T = torch.diag(T_diag)
39
+ elif T_diag.ndim == 2 and T_diag.shape == (3, 3):
40
+ T = T_diag
41
+ else:
42
+ raise ValueError("T must be a 3-vector or 3x3 matrix.")
43
+
44
+ nSteps = R.shape[2]
45
+ nTraj = R.shape[3]
46
+
47
+ # R @ T @ R^T — batched over (nSteps, nTraj)
48
+ # Reshape R to (nSteps*nTraj, 3, 3) for batch matmul
49
+ R_flat = R.permute(2, 3, 0, 1).reshape(-1, 3, 3) # (N, 3, 3)
50
+ R_inv = R_flat.transpose(-2, -1) # R^T
51
+ T_exp = T.unsqueeze(0).expand(R_flat.shape[0], -1, -1)
52
+
53
+ T_rot = torch.bmm(R_flat, torch.bmm(T_exp, R_inv))
54
+ return T_rot.reshape(nSteps, nTraj, 3, 3).permute(2, 3, 0, 1)
55
+
56
+
57
+ def spiral_grid(n: int) -> tuple[np.ndarray, np.ndarray]:
58
+ """Generate a spherical spiral grid of *n* points.
59
+
60
+ Returns (phi, theta) arrays suitable for powder averaging in cardamom.
61
+ This is the same grid used in MATLAB's ``cardamom.m`` (lines 592-595).
62
+
63
+ Parameters
64
+ ----------
65
+ n:
66
+ Number of grid points.
67
+
68
+ Returns
69
+ -------
70
+ phi:
71
+ Azimuthal angles in radians, shape ``(n,)``.
72
+ theta:
73
+ Polar angles in radians, shape ``(n,)``.
74
+ """
75
+ pts = np.linspace(-1, 1, n)
76
+ theta = np.arccos(pts)
77
+ phi = np.sqrt(np.pi * n) * np.arcsin(pts)
78
+ return phi, theta
79
+
80
+
81
+ def gelman_rubin(chains: np.ndarray) -> float:
82
+ """Compute the Gelman-Rubin R-hat convergence diagnostic.
83
+
84
+ Parameters
85
+ ----------
86
+ chains:
87
+ Array of shape ``(n_chains, n_samples)`` containing scalar
88
+ summary statistics (e.g. autocorrelation values) from
89
+ independent trajectories.
90
+
91
+ Returns
92
+ -------
93
+ R_hat:
94
+ Gelman-Rubin R statistic. Values close to 1.0 indicate
95
+ convergence; R < 1.1 is the typical threshold.
96
+ """
97
+ n_chains, n_samples = chains.shape
98
+ if n_chains < 2:
99
+ return 1.0
100
+
101
+ chain_means = chains.mean(axis=1)
102
+ chain_vars = chains.var(axis=1, ddof=1)
103
+
104
+ grand_mean = chain_means.mean()
105
+ B = n_samples * np.var(chain_means, ddof=1) # between-chain variance
106
+ W = np.mean(chain_vars) # within-chain variance
107
+
108
+ if W < 1e-30:
109
+ return 1.0
110
+
111
+ var_hat = ((n_samples - 1) / n_samples) * W + (1.0 / n_samples) * B
112
+ R_hat = np.sqrt(var_hat / W)
113
+ return float(R_hat)
torchspin/_compile.py ADDED
@@ -0,0 +1,100 @@
1
+ """torch.compile integration for torchspin.
2
+
3
+ Provides a ``maybe_compile`` decorator that conditionally applies
4
+ ``torch.compile`` to pure-PyTorch numerical kernels.
5
+
6
+ Compilation is **off by default** to preserve bit-exact numerical
7
+ reproducibility. Enable it by setting the environment variable::
8
+
9
+ export TORCHSPIN_COMPILE=1
10
+
11
+ or by calling ``set_compile_enabled(True)`` at runtime before any
12
+ compiled function is first invoked.
13
+
14
+ Requirements: PyTorch ≥ 2.0. On older versions, ``maybe_compile``
15
+ is a no-op regardless of the flag.
16
+ """
17
+ from __future__ import annotations
18
+
19
+ import functools
20
+ import os
21
+ from typing import Any, Callable, TypeVar
22
+
23
+ import torch
24
+
25
+ F = TypeVar("F", bound=Callable[..., Any])
26
+
27
+ # ── Runtime flag ──────────────────────────────────────────────────────────────
28
+
29
+ _compile_enabled: bool | None = None # None = read from env on first call
30
+
31
+
32
+ def _resolve_flag() -> bool:
33
+ global _compile_enabled
34
+ if _compile_enabled is None:
35
+ _compile_enabled = os.environ.get("TORCHSPIN_COMPILE", "0") == "1"
36
+ return _compile_enabled
37
+
38
+
39
+ def set_compile_enabled(enabled: bool) -> None:
40
+ """Enable or disable ``torch.compile`` for torchspin kernels.
41
+
42
+ Must be called *before* the first invocation of any compiled function
43
+ (i.e. at import time or early in a script).
44
+ """
45
+ global _compile_enabled
46
+ _compile_enabled = enabled
47
+
48
+
49
+ def is_compile_available() -> bool:
50
+ """Return True if torch.compile is available.
51
+
52
+ pyproject.toml requires PyTorch ≥ 2.0, so this is always True in
53
+ production; kept for defensive use in environments with unusual
54
+ torch versions.
55
+ """
56
+ return hasattr(torch, "compile")
57
+
58
+
59
+ # ── Decorator ─────────────────────────────────────────────────────────────────
60
+
61
+ def maybe_compile(fn: F | None = None, **compile_kwargs: Any) -> F | Callable[[F], F]:
62
+ """Conditionally apply ``torch.compile`` to a function.
63
+
64
+ Usage::
65
+
66
+ @maybe_compile
67
+ def my_kernel(x: torch.Tensor) -> torch.Tensor:
68
+ ...
69
+
70
+ # With options:
71
+ @maybe_compile(mode="reduce-overhead")
72
+ def my_kernel(x: torch.Tensor) -> torch.Tensor:
73
+ ...
74
+
75
+ The decorated function behaves identically to the original when
76
+ compilation is disabled or unavailable.
77
+ """
78
+ def decorator(func: F) -> F:
79
+ # Mutable container so the closure can update the reference
80
+ _compiled: list = [None]
81
+
82
+ @functools.wraps(func)
83
+ def wrapper(*args: Any, **kwargs: Any) -> Any:
84
+ # Lazy: on first call, decide whether to swap in compiled version
85
+ if _compiled[0] is None:
86
+ if _resolve_flag() and is_compile_available():
87
+ _compiled[0] = torch.compile(func, **compile_kwargs)
88
+ else:
89
+ _compiled[0] = func
90
+ return _compiled[0](*args, **kwargs)
91
+
92
+ # Expose the original for testing / introspection
93
+ wrapper._original = func # type: ignore[attr-defined]
94
+ return wrapper # type: ignore[return-value]
95
+
96
+ if fn is not None:
97
+ # Called as @maybe_compile (no parentheses)
98
+ return decorator(fn)
99
+ # Called as @maybe_compile(...) (with keyword args)
100
+ return decorator