from mace.calculators import mace_mp from ase.io import read, write from ase.io.trajectory import Trajectory from ase.optimize import FIRE, LBFGS from matplotlib.patches import Patch import numpy as np import matplotlib.pyplot as plt import matplotlib.ticker as ticker import time # Load POSCAR structure atoms = read('POSCAR', format='vasp') atoms.set_pbc(True) atoms.set_positions(atoms.get_positions().astype(np.float64)) # Load custom-trained MACE model atoms.calc = mace_mp( model="6_MACE.model", default_dtype="float64", device="cuda:1", ) # Verify calculator works print("Initial energy:", atoms.get_potential_energy(), "eV") print("Initial max force:", np.max(np.abs(atoms.get_forces())), "eV/Å") # --- Loggers --- steps = [] energies = [] fmax_vals = [] stages = [] timestamps = [] t_run_start = time.perf_counter() class EnergyLogger: def __init__(self, atoms, stage_label): self.atoms = atoms self.stage_label = stage_label self.step = 0 def __call__(self): e = self.atoms.get_potential_energy() f = np.max(np.abs(self.atoms.get_forces())) steps.append(len(steps)) energies.append(e) fmax_vals.append(f) stages.append(self.stage_label) timestamps.append(time.perf_counter() - t_run_start) self.step += 1 class UnifiedLogger: def __init__(self, atoms, stage_label, logfile): self.atoms = atoms self.stage_label = stage_label self.logfile = logfile self.step = 0 self.t_stage_start = None def __call__(self): if self.t_stage_start is None: self.t_stage_start = time.perf_counter() e = self.atoms.get_potential_energy() f = np.max(np.abs(self.atoms.get_forces())) elapsed = time.perf_counter() - self.t_stage_start with open(self.logfile, 'a') as fh: if self.step == 0: fh.write(f"\n# --- {self.stage_label} ---\n") fh.write(f"{'Step':>6} {'Energy (eV)':>14} {'fmax (eV/Å)':>12} {'Time (s)':>10}\n") fh.write(f"{self.step:>6} {e:>14.6f} {f:>12.6f} {elapsed:>10.2f}\n") self.step += 1 LOG_FILE = 'relax_MC.log' TRAJ_FILE = 'relax_MC.traj' open(LOG_FILE, 'w').close() traj = Trajectory(TRAJ_FILE, mode='w', atoms=atoms) # ---------------------------------------------------------- # Stage 1: FIRE # ---------------------------------------------------------- fire_energy_logger = EnergyLogger(atoms, stage_label="FIRE") fire_unified_logger = UnifiedLogger(atoms, stage_label="FIRE", logfile=LOG_FILE) fire = FIRE(atoms, logfile=None, dt=0.05, maxstep=0.05, dtmax=0.5, Nmin=10, finc=1.05, fdec=0.5) fire.attach(fire_energy_logger, interval=1) fire.attach(fire_unified_logger, interval=1) fire.attach(traj.write, interval=1) t_fire_start = time.perf_counter() fire.run(fmax=0.1) t_fire = time.perf_counter() - t_fire_start nsteps_fire = fire.get_number_of_steps() print(f"✓ FIRE done : {nsteps_fire} steps | " f"fmax={np.max(np.abs(atoms.get_forces())):.4f} eV/Å | " f"time={t_fire:.1f}s ({t_fire/60:.2f} min)") # ---------------------------------------------------------- # Stage 2: LBFGS # ---------------------------------------------------------- lbfgs_energy_logger = EnergyLogger(atoms, stage_label="LBFGS") lbfgs_unified_logger = UnifiedLogger(atoms, stage_label="LBFGS", logfile=LOG_FILE) lbfgs = LBFGS(atoms, logfile=None, maxstep=0.05, memory=100) lbfgs.attach(lbfgs_energy_logger, interval=1) lbfgs.attach(lbfgs_unified_logger, interval=1) lbfgs.attach(traj.write, interval=1) t_lbfgs_start = time.perf_counter() lbfgs.run(fmax=0.01) t_lbfgs = time.perf_counter() - t_lbfgs_start traj.close() nsteps_lbfgs = lbfgs.get_number_of_steps() print(f"✓ LBFGS done: {nsteps_lbfgs} steps | " f"fmax={np.max(np.abs(atoms.get_forces())):.4f} eV/Å | " f"time={t_lbfgs:.1f}s ({t_lbfgs/60:.2f} min)") t_total = t_fire + t_lbfgs final_energy = atoms.get_potential_energy() # --- Masks and timing arrays --- stages_arr = np.array(stages) times_arr = np.array(timestamps) steps_arr = np.array(steps) fire_mask = stages_arr == "FIRE" lbfgs_mask = stages_arr == "LBFGS" x_fire = times_arr[fire_mask] x_lbfgs = times_arr[lbfgs_mask] s_fire = steps_arr[fire_mask] s_lbfgs = steps_arr[lbfgs_mask] # Per-step dt dt_fire = np.diff(x_fire, prepend=x_fire[0]) if len(x_fire) > 0 else np.array([]) dt_lbfgs = np.diff(x_lbfgs, prepend=x_lbfgs[0]) if len(x_lbfgs) > 0 else np.array([]) avg_fire = dt_fire.mean() if len(dt_fire) > 0 else 0.0 avg_lbfgs = dt_lbfgs.mean() if len(dt_lbfgs) > 0 else 0.0 transition_step = s_lbfgs[0] if any(lbfgs_mask) else None # Combined for bar chart all_steps_arr = np.concatenate([s_fire, s_lbfgs]) all_dt = np.concatenate([dt_fire, dt_lbfgs]) all_colors = ['#E07B39'] * len(s_fire) + ['#1D9E75'] * len(s_lbfgs) avg_total = all_dt.mean() if len(all_dt) > 0 else 0.0 # --- Helper --- def fmt_time(seconds): s = int(seconds) h, rem = divmod(s, 3600) m, sec = divmod(rem, 60) if h > 0: return f"{h}h {m}m {sec}s" elif m > 0: return f"{m}m {sec}s" else: return f"{sec}s" # --- Console summary --- print(f"\n✓ Done! Final energy : {final_energy:.4f} eV") print(f"Total steps : {nsteps_fire + nsteps_lbfgs} (FIRE={nsteps_fire}, LBFGS={nsteps_lbfgs})") print(f"Final fmax : {np.max(np.abs(atoms.get_forces())):.4f} eV/Å") print(f"FIRE wall time : {fmt_time(t_fire)} avg {avg_fire:.2f}s/step " f"min {dt_fire.min():.2f}s max {dt_fire.max():.2f}s") print(f"LBFGS wall time : {fmt_time(t_lbfgs)} avg {avg_lbfgs:.2f}s/step " f"min {dt_lbfgs.min():.2f}s max {dt_lbfgs.max():.2f}s") print(f"Total wall time : {fmt_time(t_total)}") # Save structures write('POSCAR-relaxed', atoms, format='vasp', vasp5=True, sort=True) write('CONTCAR-relaxed', atoms, format='vasp', direct=True) # --- Timing summary to log --- with open(LOG_FILE, 'a') as fh: fh.write(f"\n# --- Timing Summary ---\n") fh.write(f"{'Stage':<8} {'Steps':>6} {'Total':>12} {'Avg (s/step)':>13} " f"{'Min (s)':>8} {'Max (s)':>8}\n") fh.write(f"{'FIRE':<8} {nsteps_fire:>6} {fmt_time(t_fire):>12} {avg_fire:>13.3f} " f"{dt_fire.min():>8.3f} {dt_fire.max():>8.3f}\n") fh.write(f"{'LBFGS':<8} {nsteps_lbfgs:>6} {fmt_time(t_lbfgs):>12} {avg_lbfgs:>13.3f} " f"{dt_lbfgs.min():>8.3f} {dt_lbfgs.max():>8.3f}\n") fh.write(f"{'Total':<8} {nsteps_fire+nsteps_lbfgs:>6} {fmt_time(t_total):>12} " f"{'—':>13} {'—':>8} {'—':>8}\n") # --- Plot --- fig, axes = plt.subplots(3, 1, figsize=(8, 8), sharex=True) fig.suptitle('MACE Relaxation — FIRE → LBFGS', fontsize=13, fontweight='normal') ax1, ax2, ax3 = axes # --- Energy panel --- ax1.plot(s_fire, np.array(energies)[fire_mask], color='#E07B39', linewidth=1.5, marker='o', markersize=2.5, markeredgewidth=0, label='FIRE') ax1.plot(s_lbfgs, np.array(energies)[lbfgs_mask], color='#1D9E75', linewidth=1.5, marker='o', markersize=2.5, markeredgewidth=0, label='LBFGS') ax1.axhline(y=final_energy, color='#D85A30', linestyle='--', linewidth=1, label=f'Final: {final_energy:.4f} eV') if transition_step is not None: ax1.axvline(x=transition_step, color='gray', linestyle=':', linewidth=1, label='FIRE→LBFGS') ax1.set_ylabel('Energy (eV)', fontsize=11) ax1.legend(fontsize=9, framealpha=0.7) ax1.yaxis.set_major_formatter(ticker.FormatStrFormatter('%.3f')) ax1.grid(True, alpha=0.3, linestyle='--') ax1.tick_params(labelsize=9) # --- fmax panel --- ax2.plot(s_fire, np.array(fmax_vals)[fire_mask], color='#E07B39', linewidth=1.5, marker='o', markersize=2.5, markeredgewidth=0, label='FIRE') ax2.plot(s_lbfgs, np.array(fmax_vals)[lbfgs_mask], color='#1D9E75', linewidth=1.5, marker='o', markersize=2.5, markeredgewidth=0, label='LBFGS') ax2.axhline(y=0.01, color='#D85A30', linestyle='--', linewidth=1, label='Convergence (0.01 eV/Å)') ax2.axhline(y=0.1, color='gray', linestyle=':', linewidth=1, label='FIRE target (0.1 eV/Å)') if transition_step is not None: ax2.axvline(x=transition_step, color='gray', linestyle=':', linewidth=1) ax2.set_ylabel('Max Force (eV/Å)', fontsize=11) ax2.legend(fontsize=9, framealpha=0.7) ax2.set_yscale('log') ax2.grid(True, alpha=0.3, linestyle='--') ax2.tick_params(labelsize=9) # --- Wall time per step bar chart --- ax3.bar(all_steps_arr, all_dt, color=all_colors, alpha=0.75, width=0.6) ax3.axhline(y=avg_total, color='#D85A30', linestyle='--', linewidth=1, label=f'Avg: {avg_total:.3f}s | Total: {fmt_time(t_total)}') if transition_step is not None: ax3.axvline(x=transition_step, color='gray', linestyle=':', linewidth=1) legend_handles = [ plt.Line2D([0], [0], color='#D85A30', linestyle='--', linewidth=1, label=f'Avg: {avg_total:.3f}s | Total: {fmt_time(t_total)}'), ] ax3.legend(handles=legend_handles, fontsize=9, framealpha=0.7) ax3.set_ylabel('Wall Time (s)', fontsize=11) ax3.set_xlabel('Ionic Step', fontsize=11) ax3.grid(True, alpha=0.3, linestyle='--', axis='y') ax3.tick_params(labelsize=9) plt.tight_layout() plt.savefig('MC_relaxation_plot_FIRE_LBFGS.png', dpi=150, bbox_inches='tight') plt.show() print("Plot saved as MC_relaxation_plot_FIRE_LBFGS.png")