fastqml 0.2.1__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.
fastqml/__init__.py ADDED
@@ -0,0 +1,6 @@
1
+ from . import gate_fusion
2
+ from . import torch_sim_real
3
+ from . import register_ops
4
+ from . import circulant
5
+
6
+ __version__ = "0.2.1"
fastqml/circulant.py ADDED
@@ -0,0 +1,50 @@
1
+ import torch
2
+
3
+
4
+ def circulant_plan(L1, L2, device, dtype=torch.float32):
5
+ N = L1 * L2
6
+ r = torch.arange(N, device=device) // L2
7
+ c = torch.arange(N, device=device) % L2
8
+ didx = ((r.unsqueeze(1) - r.unsqueeze(0)) % L1) * L2 + (c.unsqueeze(1) - c.unsqueeze(0)) % L2
9
+ th1 = 2.0 * torch.pi * torch.outer(torch.arange(L1, device=device),
10
+ torch.arange(L1, device=device)).to(dtype) / L1
11
+ th2 = 2.0 * torch.pi * torch.outer(torch.arange(L2, device=device),
12
+ torch.arange(L2, device=device)).to(dtype) / L2
13
+ return {
14
+ 'L1': L1, 'L2': L2, 'N': N, 'didx': didx,
15
+ 'W1f_re': torch.cos(th1), 'W1f_im': -torch.sin(th1),
16
+ 'W2f_re': torch.cos(th2), 'W2f_im': -torch.sin(th2),
17
+ 'W1i_re': torch.cos(th1) / L1, 'W1i_im': torch.sin(th1) / L1,
18
+ 'W2i_re': torch.cos(th2) / L2, 'W2i_im': torch.sin(th2) / L2,
19
+ }
20
+
21
+
22
+ def circulant_dense(Uk_re, Uk_im, plan):
23
+ L1, L2, N = plan['L1'], plan['L2'], plan['N']
24
+ A = Uk_re.shape[-1]
25
+ Uk = torch.complex(Uk_re, Uk_im)
26
+ Adisp = torch.fft.ifft2(Uk, dim=(0, 1)).reshape(N, A, A)
27
+ M = Adisp[plan['didx']]
28
+ M = M.permute(2, 0, 3, 1).reshape(A * N, A * N)
29
+ return M.real.contiguous(), M.imag.contiguous()
30
+
31
+
32
+ def _axis_dft(vr, vi, Wre, Wim, spec):
33
+ tr = torch.einsum(spec, vr, Wre) - torch.einsum(spec, vi, Wim)
34
+ ti = torch.einsum(spec, vr, Wim) + torch.einsum(spec, vi, Wre)
35
+ return tr, ti
36
+
37
+
38
+ def circulant_apply_k(re, im, Uk_re, Uk_im, plan):
39
+ B = re.shape[0]
40
+ L1, L2, N = plan['L1'], plan['L2'], plan['N']
41
+ A = Uk_re.shape[-1]
42
+ vr = re.view(B, A, L1, L2)
43
+ vi = im.view(B, A, L1, L2)
44
+ tr, ti = _axis_dft(vr, vi, plan['W2f_re'], plan['W2f_im'], 'baxy,ky->baxk')
45
+ ur, ui = _axis_dft(tr, ti, plan['W1f_re'], plan['W1f_im'], 'bayk,jy->bajk')
46
+ wr = torch.einsum('bijk,jkoi->bojk', ur, Uk_re) - torch.einsum('bijk,jkoi->bojk', ui, Uk_im)
47
+ wi = torch.einsum('bijk,jkoi->bojk', ur, Uk_im) + torch.einsum('bijk,jkoi->bojk', ui, Uk_re)
48
+ tr, ti = _axis_dft(wr, wi, plan['W1i_re'], plan['W1i_im'], 'bojk,xj->boxk')
49
+ outr, outi = _axis_dft(tr, ti, plan['W2i_re'], plan['W2i_im'], 'boxk,yk->boxy')
50
+ return outr.reshape(B, A * N), outi.reshape(B, A * N)
fastqml/gate_fusion.py ADDED
@@ -0,0 +1,547 @@
1
+ import math
2
+ from typing import Callable, Tuple
3
+
4
+ import torch
5
+ try:
6
+ import torch._inductor.config as _ic
7
+ _ic.triton.cudagraphs = False
8
+ except Exception:
9
+ pass
10
+
11
+
12
+ def build_ry_gates(angles: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
13
+ half = angles / 2
14
+ c, s = torch.cos(half), torch.sin(half)
15
+ G_re = torch.stack([torch.stack([c, -s], -1), torch.stack([s, c], -1)], -2)
16
+ return G_re, torch.zeros_like(G_re)
17
+
18
+
19
+ def build_rz_gates(angles: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
20
+ half = angles / 2
21
+ c, s = torch.cos(half), torch.sin(half)
22
+ z = torch.zeros_like(c)
23
+ G_re = torch.stack([torch.stack([c, z], -1), torch.stack([z, c], -1)], -2)
24
+ G_im = torch.stack([torch.stack([-s, z], -1), torch.stack([z, s], -1)], -2)
25
+ return G_re, G_im
26
+
27
+
28
+ def build_rx_gates(angles: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
29
+ half = angles / 2
30
+ c, s = torch.cos(half), torch.sin(half)
31
+ z = torch.zeros_like(c)
32
+ G_re = torch.stack([torch.stack([c, z], -1), torch.stack([z, c], -1)], -2)
33
+ G_im = torch.stack([torch.stack([z, -s], -1), torch.stack([-s, z], -1)], -2)
34
+ return G_re, G_im
35
+
36
+
37
+ def build_phase_shift_gates(phis: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
38
+ c, s = torch.cos(phis), torch.sin(phis)
39
+ one = torch.ones_like(c); z = torch.zeros_like(c)
40
+ G_re = torch.stack([torch.stack([one, z], -1), torch.stack([z, c], -1)], -2)
41
+ G_im = torch.stack([torch.stack([z, z], -1), torch.stack([z, s], -1)], -2)
42
+ return G_re, G_im
43
+
44
+
45
+ def build_rot_gates(phi: torch.Tensor, theta: torch.Tensor,
46
+ omega: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
47
+ Gr_phi, Gi_phi = build_rz_gates(phi)
48
+ Gr_th, Gi_th = build_ry_gates(theta)
49
+ Gr_om, Gi_om = build_rz_gates(omega)
50
+ a_re = torch.einsum('nij,njk->nik', Gr_th, Gr_phi) - torch.einsum('nij,njk->nik', Gi_th, Gi_phi)
51
+ a_im = torch.einsum('nij,njk->nik', Gr_th, Gi_phi) + torch.einsum('nij,njk->nik', Gi_th, Gr_phi)
52
+ out_re = torch.einsum('nij,njk->nik', Gr_om, a_re) - torch.einsum('nij,njk->nik', Gi_om, a_im)
53
+ out_im = torch.einsum('nij,njk->nik', Gr_om, a_im) + torch.einsum('nij,njk->nik', Gi_om, a_re)
54
+ return out_re, out_im
55
+
56
+
57
+ def build_ry_rz_gates(ry_angles: torch.Tensor,
58
+ rz_angles: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
59
+ cy, sy = torch.cos(ry_angles / 2), torch.sin(ry_angles / 2)
60
+ cz, sz = torch.cos(rz_angles / 2), torch.sin(rz_angles / 2)
61
+ G_re = torch.stack([
62
+ torch.stack([cz * cy, -cz * sy], -1),
63
+ torch.stack([cz * sy, cz * cy], -1),
64
+ ], -2)
65
+ G_im = torch.stack([
66
+ torch.stack([-sz * cy, sz * sy], -1),
67
+ torch.stack([ sz * sy, sz * cy], -1),
68
+ ], -2)
69
+ return G_re, G_im
70
+
71
+
72
+ def _const_2x2(re_vals, im_vals, nq, device, dtype):
73
+ G_re = torch.tensor(re_vals, device=device, dtype=dtype).unsqueeze(0).expand(nq, 2, 2)
74
+ G_im = torch.tensor(im_vals, device=device, dtype=dtype).unsqueeze(0).expand(nq, 2, 2)
75
+ return G_re.contiguous(), G_im.contiguous()
76
+
77
+
78
+ def build_x_gates(nq, device, dtype):
79
+ return _const_2x2([[0, 1], [1, 0]], [[0, 0], [0, 0]], nq, device, dtype)
80
+
81
+
82
+ def build_y_gates(nq, device, dtype):
83
+ return _const_2x2([[0, 0], [0, 0]], [[0, -1], [1, 0]], nq, device, dtype)
84
+
85
+
86
+ def build_z_gates(nq, device, dtype):
87
+ return _const_2x2([[1, 0], [0, -1]], [[0, 0], [0, 0]], nq, device, dtype)
88
+
89
+
90
+ def build_h_gates(nq, device, dtype):
91
+ s = 1.0 / math.sqrt(2.0)
92
+ return _const_2x2([[s, s], [s, -s]], [[0, 0], [0, 0]], nq, device, dtype)
93
+
94
+
95
+ def build_s_gates(nq, device, dtype):
96
+ return _const_2x2([[1, 0], [0, 0]], [[0, 0], [0, 1]], nq, device, dtype)
97
+
98
+
99
+ def build_t_gates(nq, device, dtype):
100
+ c = math.cos(math.pi / 4); s = math.sin(math.pi / 4)
101
+ return _const_2x2([[1, 0], [0, c]], [[0, 0], [0, s]], nq, device, dtype)
102
+
103
+
104
+ def build_ry_per_batch(angles):
105
+ half = angles / 2
106
+ c, s = torch.cos(half), torch.sin(half)
107
+ G_re = torch.stack([torch.stack([c, -s], -1), torch.stack([s, c], -1)], -2)
108
+ return G_re, torch.zeros_like(G_re)
109
+
110
+
111
+ def build_rz_per_batch(angles):
112
+ half = angles / 2
113
+ c, s = torch.cos(half), torch.sin(half)
114
+ z = torch.zeros_like(c)
115
+ G_re = torch.stack([torch.stack([c, z], -1), torch.stack([z, c], -1)], -2)
116
+ G_im = torch.stack([torch.stack([-s, z], -1), torch.stack([z, s], -1)], -2)
117
+ return G_re, G_im
118
+
119
+
120
+ def build_rx_per_batch(angles):
121
+ half = angles / 2
122
+ c, s = torch.cos(half), torch.sin(half)
123
+ z = torch.zeros_like(c)
124
+ G_re = torch.stack([torch.stack([c, z], -1), torch.stack([z, c], -1)], -2)
125
+ G_im = torch.stack([torch.stack([z, -s], -1), torch.stack([-s, z], -1)], -2)
126
+ return G_re, G_im
127
+
128
+
129
+ def build_ry_rx_per_batch(ry_angles, rx_angles):
130
+ cy, sy = torch.cos(ry_angles / 2), torch.sin(ry_angles / 2)
131
+ cx, sx = torch.cos(rx_angles / 2), torch.sin(rx_angles / 2)
132
+ G_re = torch.stack([
133
+ torch.stack([cx * cy, -cx * sy], -1),
134
+ torch.stack([cx * sy, cx * cy], -1),
135
+ ], -2)
136
+ G_im = torch.stack([
137
+ torch.stack([-sx * sy, -sx * cy], -1),
138
+ torch.stack([-sx * cy, sx * sy], -1),
139
+ ], -2)
140
+ return G_re, G_im
141
+
142
+
143
+ def apply_su2_layer(re: torch.Tensor, im: torch.Tensor,
144
+ G_re: torch.Tensor, G_im: torch.Tensor):
145
+ B = re.shape[0]
146
+ nq = G_re.shape[0]
147
+ U_re, U_im = G_re[0], G_im[0]
148
+ for q in range(1, nq):
149
+ Cre, Cim = G_re[q], G_im[q]
150
+ new_re = torch.kron(U_re, Cre) - torch.kron(U_im, Cim)
151
+ new_im = torch.kron(U_re, Cim) + torch.kron(U_im, Cre)
152
+ U_re, U_im = new_re, new_im
153
+ re_flat = re.reshape(B, -1)
154
+ im_flat = im.reshape(B, -1)
155
+ Ut_re, Ut_im = U_re.T, U_im.T
156
+ new_re = re_flat @ Ut_re - im_flat @ Ut_im
157
+ new_im = re_flat @ Ut_im + im_flat @ Ut_re
158
+ return new_re.reshape(re.shape), new_im.reshape(im.shape)
159
+
160
+
161
+ def apply_per_batch_su2_chain(re, im, G_re, G_im):
162
+ nq = G_re.shape[1]
163
+ for q in range(nq):
164
+ re = re.movedim(q + 1, -1)
165
+ im = im.movedim(q + 1, -1)
166
+ Gr, Gi = G_re[:, q], G_im[:, q]
167
+ new_re = torch.einsum('bij,b...j->b...i', Gr, re) - torch.einsum('bij,b...j->b...i', Gi, im)
168
+ new_im = torch.einsum('bij,b...j->b...i', Gr, im) + torch.einsum('bij,b...j->b...i', Gi, re)
169
+ re = new_re.movedim(-1, q + 1)
170
+ im = new_im.movedim(-1, q + 1)
171
+ return re, im
172
+
173
+
174
+ def _zero_4x4(n_pairs, device, dtype):
175
+ return torch.zeros(n_pairs, 4, 4, device=device, dtype=dtype)
176
+
177
+
178
+ def build_cry_4x4(angles):
179
+ n = angles.shape[0]
180
+ dev, dt = angles.device, angles.dtype
181
+ c, s = torch.cos(angles / 2), torch.sin(angles / 2)
182
+ G_re = _zero_4x4(n, dev, dt)
183
+ G_re[:, 0, 0] = 1.0; G_re[:, 1, 1] = 1.0
184
+ G_re[:, 2, 2] = c; G_re[:, 2, 3] = -s
185
+ G_re[:, 3, 2] = s; G_re[:, 3, 3] = c
186
+ return G_re, _zero_4x4(n, dev, dt)
187
+
188
+
189
+ def build_crx_4x4(angles):
190
+ n = angles.shape[0]
191
+ dev, dt = angles.device, angles.dtype
192
+ c, s = torch.cos(angles / 2), torch.sin(angles / 2)
193
+ G_re = _zero_4x4(n, dev, dt)
194
+ G_im = _zero_4x4(n, dev, dt)
195
+ G_re[:, 0, 0] = 1.0; G_re[:, 1, 1] = 1.0
196
+ G_re[:, 2, 2] = c; G_re[:, 3, 3] = c
197
+ G_im[:, 2, 3] = -s; G_im[:, 3, 2] = -s
198
+ return G_re, G_im
199
+
200
+
201
+ def build_crz_4x4(angles):
202
+ n = angles.shape[0]
203
+ dev, dt = angles.device, angles.dtype
204
+ c, s = torch.cos(angles / 2), torch.sin(angles / 2)
205
+ G_re = _zero_4x4(n, dev, dt)
206
+ G_im = _zero_4x4(n, dev, dt)
207
+ G_re[:, 0, 0] = 1.0; G_re[:, 1, 1] = 1.0
208
+ G_re[:, 2, 2] = c; G_re[:, 3, 3] = c
209
+ G_im[:, 2, 2] = -s; G_im[:, 3, 3] = s
210
+ return G_re, G_im
211
+
212
+
213
+ def build_cphase_4x4(phis):
214
+ n = phis.shape[0]
215
+ dev, dt = phis.device, phis.dtype
216
+ c, s = torch.cos(phis), torch.sin(phis)
217
+ G_re = _zero_4x4(n, dev, dt)
218
+ G_im = _zero_4x4(n, dev, dt)
219
+ G_re[:, 0, 0] = 1.0; G_re[:, 1, 1] = 1.0; G_re[:, 2, 2] = 1.0
220
+ G_re[:, 3, 3] = c; G_im[:, 3, 3] = s
221
+ return G_re, G_im
222
+
223
+
224
+ def build_ising_zz_4x4(angles):
225
+ n = angles.shape[0]
226
+ dev, dt = angles.device, angles.dtype
227
+ c, s = torch.cos(angles / 2), torch.sin(angles / 2)
228
+ G_re = _zero_4x4(n, dev, dt)
229
+ G_im = _zero_4x4(n, dev, dt)
230
+ G_re[:, 0, 0] = c; G_im[:, 0, 0] = -s
231
+ G_re[:, 1, 1] = c; G_im[:, 1, 1] = s
232
+ G_re[:, 2, 2] = c; G_im[:, 2, 2] = s
233
+ G_re[:, 3, 3] = c; G_im[:, 3, 3] = -s
234
+ return G_re, G_im
235
+
236
+
237
+ def build_ising_xx_4x4(angles):
238
+ n = angles.shape[0]
239
+ dev, dt = angles.device, angles.dtype
240
+ c, s = torch.cos(angles / 2), torch.sin(angles / 2)
241
+ G_re = _zero_4x4(n, dev, dt); G_im = _zero_4x4(n, dev, dt)
242
+ G_re[:, 0, 0] = c; G_re[:, 1, 1] = c; G_re[:, 2, 2] = c; G_re[:, 3, 3] = c
243
+ G_im[:, 0, 3] = -s; G_im[:, 3, 0] = -s
244
+ G_im[:, 1, 2] = -s; G_im[:, 2, 1] = -s
245
+ return G_re, G_im
246
+
247
+
248
+ def build_ising_yy_4x4(angles):
249
+ n = angles.shape[0]
250
+ dev, dt = angles.device, angles.dtype
251
+ c, s = torch.cos(angles / 2), torch.sin(angles / 2)
252
+ G_re = _zero_4x4(n, dev, dt); G_im = _zero_4x4(n, dev, dt)
253
+ G_re[:, 0, 0] = c; G_re[:, 1, 1] = c; G_re[:, 2, 2] = c; G_re[:, 3, 3] = c
254
+ G_im[:, 0, 3] = s; G_im[:, 3, 0] = s
255
+ G_im[:, 1, 2] = -s; G_im[:, 2, 1] = -s
256
+ return G_re, G_im
257
+
258
+
259
+ def _const_4x4(re_vals, im_vals, n_pairs, device, dtype):
260
+ G_re = torch.tensor(re_vals, device=device, dtype=dtype).unsqueeze(0).expand(n_pairs, 4, 4)
261
+ G_im = torch.tensor(im_vals, device=device, dtype=dtype).unsqueeze(0).expand(n_pairs, 4, 4)
262
+ return G_re.contiguous(), G_im.contiguous()
263
+
264
+
265
+ def build_cnot_4x4(n_pairs, device, dtype):
266
+ z = [[0.] * 4 for _ in range(4)]
267
+ re = [[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 0, 1], [0, 0, 1, 0]]
268
+ return _const_4x4(re, z, n_pairs, device, dtype)
269
+
270
+
271
+ def build_cy_4x4(n_pairs, device, dtype):
272
+ re = [[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0]]
273
+ im = [[0, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, -1], [0, 0, 1, 0]]
274
+ return _const_4x4(re, im, n_pairs, device, dtype)
275
+
276
+
277
+ def build_cz_4x4(n_pairs, device, dtype):
278
+ re = [[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 1, 0], [0, 0, 0, -1]]
279
+ z = [[0.] * 4 for _ in range(4)]
280
+ return _const_4x4(re, z, n_pairs, device, dtype)
281
+
282
+
283
+ def build_swap_4x4(n_pairs, device, dtype):
284
+ re = [[1, 0, 0, 0], [0, 0, 1, 0], [0, 1, 0, 0], [0, 0, 0, 1]]
285
+ z = [[0.] * 4 for _ in range(4)]
286
+ return _const_4x4(re, z, n_pairs, device, dtype)
287
+
288
+
289
+ def build_iswap_4x4(n_pairs, device, dtype):
290
+ re = [[1, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 1]]
291
+ im = [[0, 0, 0, 0], [0, 0, 1, 0], [0, 1, 0, 0], [0, 0, 0, 0]]
292
+ return _const_4x4(re, im, n_pairs, device, dtype)
293
+
294
+
295
+ def cry_commute(e1, e2):
296
+ a, b = e1
297
+ c, d = e2
298
+ return not (a == d or b == c)
299
+
300
+
301
+ def diagonal_commute(e1, e2):
302
+ return True
303
+
304
+
305
+ def symmetric_commute(e1, e2):
306
+ a, b = e1; c, d = e2
307
+ return not ({a, b} & {c, d})
308
+
309
+
310
+ def greedy_edge_coloring(edges, commute_fn: Callable = cry_commute):
311
+ colors = []
312
+ for k, edge in enumerate(edges):
313
+ qa, qb = edge
314
+ placed = None
315
+ for c in range(len(colors)):
316
+ disjoint = all(qa not in e and qb not in e for _, e in colors[c])
317
+ if not disjoint:
318
+ continue
319
+ ok = True
320
+ for c2 in range(c + 1, len(colors)):
321
+ for _, e_other in colors[c2]:
322
+ if not commute_fn(edge, e_other):
323
+ ok = False; break
324
+ if not ok: break
325
+ if ok:
326
+ placed = c; break
327
+ if placed is not None:
328
+ colors[placed].append((k, edge))
329
+ else:
330
+ colors.append([(k, edge)])
331
+ return colors
332
+
333
+
334
+ def identity_pad_gates(G_re, G_im, qubit_indices, nq):
335
+ device, dtype = G_re.device, G_re.dtype
336
+ eye2 = torch.eye(2, device=device, dtype=dtype)
337
+ zero2 = torch.zeros(2, 2, device=device, dtype=dtype)
338
+ qubit_set = set(qubit_indices)
339
+ re_slabs, im_slabs = [], []
340
+ g = 0
341
+ for q in range(nq):
342
+ if q in qubit_set:
343
+ re_slabs.append(G_re[g]); im_slabs.append(G_im[g]); g += 1
344
+ else:
345
+ re_slabs.append(eye2); im_slabs.append(zero2)
346
+ return torch.stack(re_slabs, dim=0), torch.stack(im_slabs, dim=0)
347
+
348
+
349
+ def apply_disjoint_cry_layer(re, im, plan, angles):
350
+ G_re, G_im = build_cry_4x4(angles)
351
+ return apply_disjoint_two_qubit_layer(re, im, plan, G_re, G_im)
352
+
353
+
354
+ def _const_8x8(re_vals, im_vals, n_triples, device, dtype):
355
+ G_re = torch.tensor(re_vals, device=device, dtype=dtype).unsqueeze(0).expand(n_triples, 8, 8)
356
+ G_im = torch.tensor(im_vals, device=device, dtype=dtype).unsqueeze(0).expand(n_triples, 8, 8)
357
+ return G_re.contiguous(), G_im.contiguous()
358
+
359
+
360
+ def build_toffoli_8x8(n_triples, device, dtype):
361
+ re = [[0.0] * 8 for _ in range(8)]
362
+ for i in range(6):
363
+ re[i][i] = 1.0
364
+ re[6][7] = 1.0
365
+ re[7][6] = 1.0
366
+ im = [[0.0] * 8 for _ in range(8)]
367
+ return _const_8x8(re, im, n_triples, device, dtype)
368
+
369
+
370
+ def build_cswap_8x8(n_triples, device, dtype):
371
+ re = [[0.0] * 8 for _ in range(8)]
372
+ for i in range(4):
373
+ re[i][i] = 1.0
374
+ re[4][4] = 1.0
375
+ re[5][6] = 1.0
376
+ re[6][5] = 1.0
377
+ re[7][7] = 1.0
378
+ im = [[0.0] * 8 for _ in range(8)]
379
+ return _const_8x8(re, im, n_triples, device, dtype)
380
+
381
+
382
+ def toffoli_commute(t1, t2):
383
+ c11, c12, t1_ = t1
384
+ c21, c22, t2_ = t2
385
+ return not (t1_ in (c21, c22) or t2_ in (c11, c12))
386
+
387
+
388
+ def cswap_commute(t1, t2):
389
+ return not (set(t1) & set(t2))
390
+
391
+
392
+ def greedy_triple_coloring(triples, commute_fn: Callable = toffoli_commute):
393
+ colors = []
394
+ for k, triple in enumerate(triples):
395
+ q_set = set(triple)
396
+ placed = None
397
+ for c in range(len(colors)):
398
+ disjoint = all(not (q_set & set(t)) for _, t in colors[c])
399
+ if not disjoint:
400
+ continue
401
+ ok = True
402
+ for c2 in range(c + 1, len(colors)):
403
+ for _, t_other in colors[c2]:
404
+ if not commute_fn(triple, t_other):
405
+ ok = False; break
406
+ if not ok: break
407
+ if ok:
408
+ placed = c; break
409
+ if placed is not None:
410
+ colors[placed].append((k, triple))
411
+ else:
412
+ colors.append([(k, triple)])
413
+ return colors
414
+
415
+
416
+ _triple_plan_cache: dict = {}
417
+
418
+
419
+ def compute_three_qubit_plans(triples, nq: int, device, dtype,
420
+ commute_fn: Callable = toffoli_commute):
421
+ if len(triples) == 0:
422
+ return []
423
+ key = (tuple(triples), nq, str(device), str(dtype), commute_fn.__name__)
424
+ if key in _triple_plan_cache:
425
+ return _triple_plan_cache[key]
426
+ colors = greedy_triple_coloring(triples, commute_fn)
427
+ plans = []
428
+ for color in colors:
429
+ k_indices = [k for k, _ in color]
430
+ color_triples = [t for _, t in color]
431
+ tripled = [q for t in color_triples for q in t]
432
+ unpaired = [q for q in range(nq) if q not in set(tripled)]
433
+ new_order = tripled + unpaired
434
+ perm = [0] + [q + 1 for q in new_order]
435
+ inv_perm = [0] * (nq + 1)
436
+ for i, p in enumerate(perm):
437
+ inv_perm[p] = i
438
+ plans.append({
439
+ 'n_triples': len(color_triples),
440
+ 'n_unpaired': len(unpaired),
441
+ 'dim': 1 << nq,
442
+ 'k_idx': torch.tensor(k_indices, dtype=torch.long, device=device),
443
+ 'perm': tuple(perm),
444
+ 'inv_perm': tuple(inv_perm),
445
+ 'triples': color_triples,
446
+ })
447
+ _triple_plan_cache[key] = plans
448
+ return plans
449
+
450
+
451
+ def apply_disjoint_three_qubit_layer(re, im, plan, G_re, G_im):
452
+ n_triples = plan['n_triples']
453
+ n_unpaired = plan['n_unpaired']
454
+ DIM = plan['dim']
455
+ B = re.shape[0]
456
+ device, dtype = re.device, re.dtype
457
+
458
+ U_re = G_re[0]; U_im = G_im[0]
459
+ for k in range(1, n_triples):
460
+ Cre, Cim = G_re[k], G_im[k]
461
+ new_re = torch.kron(U_re, Cre) - torch.kron(U_im, Cim)
462
+ new_im = torch.kron(U_re, Cim) + torch.kron(U_im, Cre)
463
+ U_re, U_im = new_re, new_im
464
+ if n_unpaired > 0:
465
+ I2 = torch.eye(2, device=device, dtype=dtype)
466
+ for _ in range(n_unpaired):
467
+ U_re = torch.kron(U_re, I2)
468
+ U_im = torch.kron(U_im, I2)
469
+
470
+ re_p = re.permute(plan['perm']).contiguous()
471
+ im_p = im.permute(plan['perm']).contiguous()
472
+ re_flat = re_p.reshape(B, DIM)
473
+ im_flat = im_p.reshape(B, DIM)
474
+ Ut_re, Ut_im = U_re.T, U_im.T
475
+ new_re = re_flat @ Ut_re - im_flat @ Ut_im
476
+ new_im = re_flat @ Ut_im + im_flat @ Ut_re
477
+ out_shape = re_p.shape
478
+ new_re = new_re.reshape(out_shape).permute(plan['inv_perm']).contiguous()
479
+ new_im = new_im.reshape(out_shape).permute(plan['inv_perm']).contiguous()
480
+ return new_re, new_im
481
+
482
+
483
+ _plan_cache: dict = {}
484
+
485
+
486
+ def compute_two_qubit_plans(edges, nq: int, device, dtype,
487
+ commute_fn: Callable = cry_commute):
488
+ if len(edges) == 0:
489
+ return []
490
+ key = (tuple(edges), nq, str(device), str(dtype), commute_fn.__name__)
491
+ if key in _plan_cache:
492
+ return _plan_cache[key]
493
+ colors = greedy_edge_coloring(edges, commute_fn)
494
+ plans = []
495
+ for color in colors:
496
+ k_indices = [k for k, _ in color]
497
+ color_edges = [e for _, e in color]
498
+ paired = [q for qa, qb in color_edges for q in (qa, qb)]
499
+ unpaired = [q for q in range(nq) if q not in set(paired)]
500
+ new_order = paired + unpaired
501
+ perm = [0] + [q + 1 for q in new_order]
502
+ inv_perm = [0] * (nq + 1)
503
+ for i, p in enumerate(perm):
504
+ inv_perm[p] = i
505
+ plans.append({
506
+ 'n_pairs': len(color_edges),
507
+ 'n_unpaired': len(unpaired),
508
+ 'dim': 1 << nq,
509
+ 'k_idx': torch.tensor(k_indices, dtype=torch.long, device=device),
510
+ 'perm': tuple(perm),
511
+ 'inv_perm': tuple(inv_perm),
512
+ 'edges': color_edges,
513
+ })
514
+ _plan_cache[key] = plans
515
+ return plans
516
+
517
+
518
+ def apply_disjoint_two_qubit_layer(re, im, plan, G_re, G_im):
519
+ n_pairs = plan['n_pairs']
520
+ n_unpaired = plan['n_unpaired']
521
+ DIM = plan['dim']
522
+ B = re.shape[0]
523
+ device, dtype = re.device, re.dtype
524
+
525
+ U_re = G_re[0]; U_im = G_im[0]
526
+ for k in range(1, n_pairs):
527
+ Cre, Cim = G_re[k], G_im[k]
528
+ new_re = torch.kron(U_re, Cre) - torch.kron(U_im, Cim)
529
+ new_im = torch.kron(U_re, Cim) + torch.kron(U_im, Cre)
530
+ U_re, U_im = new_re, new_im
531
+ if n_unpaired > 0:
532
+ I2 = torch.eye(2, device=device, dtype=dtype)
533
+ for _ in range(n_unpaired):
534
+ U_re = torch.kron(U_re, I2)
535
+ U_im = torch.kron(U_im, I2)
536
+
537
+ re_p = re.permute(plan['perm']).contiguous()
538
+ im_p = im.permute(plan['perm']).contiguous()
539
+ re_flat = re_p.reshape(B, DIM)
540
+ im_flat = im_p.reshape(B, DIM)
541
+ Ut_re, Ut_im = U_re.T, U_im.T
542
+ new_re = re_flat @ Ut_re - im_flat @ Ut_im
543
+ new_im = re_flat @ Ut_im + im_flat @ Ut_re
544
+ out_shape = re_p.shape
545
+ new_re = new_re.reshape(out_shape).permute(plan['inv_perm']).contiguous()
546
+ new_im = new_im.reshape(out_shape).permute(plan['inv_perm']).contiguous()
547
+ return new_re, new_im