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.
- torchspin/__init__.py +363 -0
- torchspin/_cardamom_istos.py +327 -0
- torchspin/_cardamom_propagatedm.py +413 -0
- torchspin/_cardamom_utils.py +113 -0
- torchspin/_compile.py +100 -0
- torchspin/_linalg.py +140 -0
- torchspin/angmom.py +742 -0
- torchspin/autograd.py +557 -0
- torchspin/autoguess.py +342 -0
- torchspin/batch.py +225 -0
- torchspin/blochsteady.py +274 -0
- torchspin/cardamom.py +552 -0
- torchspin/chili.py +932 -0
- torchspin/chili_sle.py +1233 -0
- torchspin/constants.py +88 -0
- torchspin/convspec.py +148 -0
- torchspin/ctafft.py +91 -0
- torchspin/curry.py +241 -0
- torchspin/data/FourierSeriesCoefficients.txt +89 -0
- torchspin/data/GaussianCascadeCoefficients.txt +56 -0
- torchspin/data/isotopedata.txt +416 -0
- torchspin/data/spacegroups.txt +637 -0
- torchspin/dataproc.py +361 -0
- torchspin/dipbackground.py +69 -0
- torchspin/dipkernel.py +102 -0
- torchspin/diptensor.py +138 -0
- torchspin/endorfrq.py +198 -0
- torchspin/endorfrq_perturb.py +302 -0
- torchspin/eprload.py +1538 -0
- torchspin/eprsave.py +165 -0
- torchspin/esfit.py +2427 -0
- torchspin/evolve.py +274 -0
- torchspin/ewrls.py +124 -0
- torchspin/excitation.py +107 -0
- torchspin/exciteprofile.py +132 -0
- torchspin/experiment.py +323 -0
- torchspin/exponfit.py +177 -0
- torchspin/fastmotion.py +182 -0
- torchspin/fdaxis.py +57 -0
- torchspin/fitgui.py +350 -0
- torchspin/garlic.py +760 -0
- torchspin/ham.py +145 -0
- torchspin/ham_cf.py +115 -0
- torchspin/ham_ee.py +124 -0
- torchspin/ham_ez.py +119 -0
- torchspin/ham_ezho.py +403 -0
- torchspin/ham_hf.py +122 -0
- torchspin/ham_nn.py +108 -0
- torchspin/ham_nq.py +100 -0
- torchspin/ham_nz.py +136 -0
- torchspin/ham_oz.py +115 -0
- torchspin/ham_so.py +118 -0
- torchspin/ham_zf.py +254 -0
- torchspin/hamsymm.py +723 -0
- torchspin/initstate.py +158 -0
- torchspin/isotopologues.py +296 -0
- torchspin/levels.py +189 -0
- torchspin/levelsplot.py +363 -0
- torchspin/lineshape.py +715 -0
- torchspin/lpsvd.py +219 -0
- torchspin/makespec.py +87 -0
- torchspin/mdhmm.py +573 -0
- torchspin/mdload.py +788 -0
- torchspin/mdtraj2oripot.py +90 -0
- torchspin/ml.py +377 -0
- torchspin/mlpsvd.py +297 -0
- torchspin/nucdata.py +310 -0
- torchspin/nucfrq2d.py +228 -0
- torchspin/orca2torchspin.py +1198 -0
- torchspin/ordering.py +78 -0
- torchspin/oripotentialplot.py +246 -0
- torchspin/orisel.py +176 -0
- torchspin/pepper.py +2019 -0
- torchspin/pepper_autograd.py +613 -0
- torchspin/photoselect.py +174 -0
- torchspin/plegendre.py +131 -0
- torchspin/propint.py +127 -0
- torchspin/pulse.py +652 -0
- torchspin/py.typed +0 -0
- torchspin/rapidscan2spc.py +101 -0
- torchspin/resfields.py +337 -0
- torchspin/resfields_batch.py +446 -0
- torchspin/resfields_eig.py +250 -0
- torchspin/resfields_perturb.py +706 -0
- torchspin/resfreqs_matrix.py +344 -0
- torchspin/resfreqs_perturb.py +299 -0
- torchspin/resonator.py +310 -0
- torchspin/resonatorprofile.py +85 -0
- torchspin/rfmixer.py +147 -0
- torchspin/rotations.py +87 -0
- torchspin/rotutils.py +880 -0
- torchspin/saffron.py +1523 -0
- torchspin/saffron_pathways.py +92 -0
- torchspin/saffron_peaks.py +580 -0
- torchspin/saffron_thyme.py +472 -0
- torchspin/salt.py +680 -0
- torchspin/sigeq.py +71 -0
- torchspin/signalprocessing.py +226 -0
- torchspin/sitetransforms.py +254 -0
- torchspin/sphgrid.py +470 -0
- torchspin/spidyan.py +1508 -0
- torchspin/spinladder.py +161 -0
- torchspin/spinops.py +265 -0
- torchspin/spinsystem.py +1427 -0
- torchspin/stackplot.py +167 -0
- torchspin/stev.py +283 -0
- torchspin/stochtraj_diffusion.py +494 -0
- torchspin/stochtraj_jump.py +195 -0
- torchspin/strainwidth.py +727 -0
- torchspin/transmitter.py +89 -0
- torchspin/utils.py +605 -0
- torchspin-0.3.0.dist-info/METADATA +267 -0
- torchspin-0.3.0.dist-info/RECORD +116 -0
- torchspin-0.3.0.dist-info/WHEEL +5 -0
- torchspin-0.3.0.dist-info/licenses/LICENSE.md +22 -0
- 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
|