Skip to content

Commit 88c4808

Browse files
committed
Support MPS path also
1 parent 6462d05 commit 88c4808

1 file changed

Lines changed: 24 additions & 18 deletions

File tree

mitsubishi/phase_3/code/euv_spectra.py

Lines changed: 24 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -47,11 +47,6 @@
4747
# %config InlineBackend.figure_format = 'retina'
4848

4949
import os
50-
# Allow PyTorch MKL and PennyLane GCC OpenMP runtimes to coexist in the same process
51-
os.environ["KMP_DUPLICATE_LIBOK"] = "TRUE"
52-
# Cap OpenMP threads to prevent CPU thread-pool mutex thrashing
53-
os.environ["OMP_NUM_THREADS"] = "8"
54-
os.environ["MKL_NUM_THREADS"] = "8"
5550
import sys
5651
import pickle
5752
import matplotlib.pyplot as plt
@@ -415,7 +410,7 @@ def generate_molecule_data(molecule_name="H2", source="qchem", local_dataset_pat
415410
if torch.cuda.is_available() and torch.cuda.device_count() <= 1:
416411
# Single H100 / H200 GPU instance (141 GB VRAM / 180 GB RAM):
417412
# 30 active-space qubits + 1 ancilla = 31 qubits (requires ~34 GB VRAM / 17 GB RAM)
418-
active_orbitals = 15
413+
active_orbitals = 14
419414
active_electrons = 16
420415
else:
421416
# 8x H100 GPU cluster (640 GB VRAM): 32 active-space qubits + 1 ancilla = 33 qubits
@@ -1216,24 +1211,23 @@ def get_subsequence_energies(op_seq):
12161211
)
12171212
return np.array(energies)
12181213

1219-
# Verify with a tiny sequence
1220-
# print(get_subsequence_energies([[op_pool[0], op_pool[1]]]))
1221-
1222-
dev_eval = qml.device("lightning.qubit", wires=num_qubits)
1223-
1224-
@qml.qnode(dev_eval)
1214+
@qml.qnode(dev)
12251215
def final_energy_circuit(gqe_ops):
1226-
"""Executes a sequence of GQE operators and measures final ground state energy directly on CPU (single pass, no GPU VRAM lock)."""
1216+
"""Executes a sequence of GQE operators and measures final ground state energy directly."""
12271217
qml.BasisState(init_state, wires=range(num_qubits))
12281218
for op in gqe_ops:
12291219
qml.apply(op)
12301220
return qml.expval(meas_hamiltonian)
12311221

12321222
def get_final_energies(gen_op_seq):
1233-
"""Evaluates final state energy for a batch of operator sequences in a single pass on CPU."""
1223+
"""Evaluates final state energy for a batch of operator sequences in a single pass."""
12341224
energies = [float(final_energy_circuit(ops)) for ops in gen_op_seq]
12351225
return np.array(energies).reshape(-1, 1)
12361226

1227+
# Test the tiny sequence using the fast single-pass function
1228+
print("Testing tiny sequence:", get_final_energies([[op_pool[0], op_pool[1]]]))
1229+
1230+
12371231
# %% [markdown]
12381232
# ## GQE Training & Dataset Hyperparameters
12391233
#
@@ -1247,7 +1241,7 @@ def get_final_energies(gen_op_seq):
12471241
# %%
12481242
GQE_TRAIN_SIZE = 128
12491243
GQE_NUM_ITERS = 10000
1250-
GQE_EVAL_FREQ = 100
1244+
GQE_EVAL_FREQ = 2500
12511245
GQE_SAMPLE_EVAL_SIZE = 32
12521246

12531247
# %% [markdown]
@@ -1676,7 +1670,11 @@ def generate(self, n_sequences, max_new_tokens, energies, temperature=1.0, devic
16761670

16771671
gen_inds = (gen_token_seq[:, 1:] - 1).cpu().numpy()
16781672
gen_op_seq = op_pool[gen_inds]
1679-
true_Es = get_final_energies(gen_op_seq)
1673+
if USE_CUDA:
1674+
true_Es = get_final_energies(gen_op_seq)
1675+
else:
1676+
true_Es = get_subsequence_energies(gen_op_seq)[:, -1].reshape(-1, 1)
1677+
16801678

16811679
mae = np.mean(np.abs(pred_Es - true_Es))
16821680
ave_E = np.mean(true_Es)
@@ -1832,7 +1830,11 @@ def generate(self, n_sequences, max_new_tokens, energies, temperature=1.0, devic
18321830

18331831
gen_inds_ = (gen_token_seq_[:, 1:] - 1).cpu().numpy()
18341832
gen_op_seq_ = op_pool[gen_inds_]
1835-
true_Es_ = get_final_energies(gen_op_seq_)
1833+
1834+
if USE_CUDA:
1835+
true_Es_ = get_final_energies(gen_op_seq_)
1836+
else:
1837+
true_Es_ = get_subsequence_energies(gen_op_seq_)[:, -1].reshape(-1, 1)
18361838

18371839
# Best model
18381840
loaded = torch.load(model_path, map_location=device, weights_only=False)
@@ -1846,7 +1848,11 @@ def generate(self, n_sequences, max_new_tokens, energies, temperature=1.0, devic
18461848

18471849
loaded_inds_ = (loaded_token_seq_[:, 1:] - 1).cpu().numpy()
18481850
loaded_op_seq_ = op_pool[loaded_inds_]
1849-
loaded_true_Es_ = get_final_energies(loaded_op_seq_)
1851+
1852+
if USE_CUDA:
1853+
loaded_true_Es_ = get_final_energies(loaded_op_seq_)
1854+
else:
1855+
loaded_true_Es_ = get_subsequence_energies(loaded_op_seq_)[:, -1].reshape(-1, 1)
18501856

18511857
# Summary table
18521858
df_compare_Es = pd.DataFrame({

0 commit comments

Comments
 (0)