""" engines/engine_dll.py --------------------- Python ctypes 包装器:加载 C/C++/Fortran 动态链接库并调用 run_dynamics()。 用法(由 compute.py 内部调用,不直接运行): from engines.engine_dll import load_dll, run_dynamics_dll dll = load_dll("c") # 自动查找 engines/c/build/dynamics_c.dll/.so/.dylib arrays = run_dynamics_dll(dll, config, atom_data, bond_data, driver_data) # arrays: dict with keys x, y, z, vx, vy, vz shape=(n_frames, n_atoms) DLL 编译(C 版本): Windows: gcc -O3 -shared -o engines/release/dynamics_c.dll engines/src/c/dynamics_lib.c -lm Linux: gcc -O3 -shared -fPIC -o engines/release/dynamics_c.so engines/src/c/dynamics_lib.c -lm macOS: gcc -O3 -dynamiclib -o engines/release/dynamics_c.dylib engines/src/c/dynamics_lib.c -lm 或用 make dll 一键编译: cd engines/src/c && make dll """ import ctypes import os import platform import numpy as np # ── DLL 文件名后缀 ───────────────────────────────────────────── _SUFFIX = { "windows": ".dll", "linux": ".so", "darwin": ".dylib", } # ── method 字符串 → 整数 ID ──────────────────────────────────── _METHOD_ID = { "explicit_euler": 0, "euler": 0, "implicit_euler": 1, "midpoint": 2, "leapfrog": 3, } _HERE = os.path.dirname(os.path.abspath(__file__)) _DLL_NAME = { "c": "dynamics_c", "cpp": "dynamics_cpp", "c++": "dynamics_cpp", "fortran": "dynamics_f90", "f90": "dynamics_f90", # "python" 引擎通过直接 import 调用,不使用 DLL } # 引擎名规范化:将别名统一为目录名 _ENGINE_DIR = { "c": "c", "cpp": "cpp", "c++": "cpp", "fortran": "fortran", "f90": "fortran", "python": "python", } def _dll_candidates(engine: str) -> list[str]: """返回 DLL 候选路径列表(按优先级)。""" sys = platform.system().lower() ext = _SUFFIX.get(sys, ".so") eng_dir = _ENGINE_DIR.get(engine, engine) name = _DLL_NAME.get(engine, f"dynamics_{engine}") base = os.path.join(_HERE, "release", name) return [ base + ext, base + ".dll", base + ".so", base + ".dylib", ] def load_dll(engine: str = "c"): """加载指定引擎。 - C/C++/Fortran: 返回 ctypes.CDLL 对象 - Python: 返回模块对象(直接 import,无需编译) Args: engine: "c", "cpp", "fortran", 或 "python" Raises: FileNotFoundError: DLL/模块文件不存在 """ if _ENGINE_DIR.get(engine, engine) == "python": import importlib.util, sys as _sys mod_path = os.path.join(_HERE, "python", "dynamics_lib.py") if not os.path.exists(mod_path): raise FileNotFoundError(f"Python 引擎未找到: {mod_path}") spec = importlib.util.spec_from_file_location( "engines.python.dynamics_lib", mod_path) mod = importlib.util.module_from_spec(spec) spec.loader.exec_module(mod) return mod # 返回模块,不是 CDLL for p in _dll_candidates(engine): if os.path.exists(p): lib = ctypes.CDLL(p) _setup_prototype(lib) return lib raise FileNotFoundError( f"DLL 未找到(引擎 {engine}),候选路径:\n" + "\n".join(f" {p}" for p in _dll_candidates(engine)) + f"\n请先编译:cd engines/{engine} && make dll" ) def _setup_prototype(lib: ctypes.CDLL) -> None: """配置 run_dynamics 的参数类型和返回类型。""" c_dbl_p = ctypes.POINTER(ctypes.c_double) c_int_p = ctypes.POINTER(ctypes.c_int) cb_type = ctypes.CFUNCTYPE(None, ctypes.c_int, ctypes.c_int) lib.run_dynamics.restype = ctypes.c_int lib.run_dynamics.argtypes = [ ctypes.c_int, # n_atoms c_dbl_p, # pos_init [n_atoms*3] c_dbl_p, # vel_init [n_atoms*3] c_dbl_p, # masses [n_atoms] c_int_p, # fixed [n_atoms*3] ctypes.c_int, # n_bonds c_int_p, # bond_pairs [n_bonds*2] c_dbl_p, # bond_k [n_bonds] c_dbl_p, # bond_r0 [n_bonds] ctypes.c_double, # box_a ctypes.c_double, # dt ctypes.c_int, # NT ctypes.c_int, # NSTEP ctypes.c_int, # warmup_steps ctypes.c_int, # method_id ctypes.c_double, # Gx ctypes.c_double, # Gy ctypes.c_double, # Gz ctypes.c_double, # Bx ctypes.c_double, # By ctypes.c_double, # Bz ctypes.c_int, # gravity_field ctypes.c_int, # elastic_force ctypes.c_int, # damping_force ctypes.c_double, # gravity_strength ctypes.c_int, # n_drivers c_int_p, # drv_idx [n_drivers] c_dbl_p, # drv_amp [n_drivers*3] c_dbl_p, # drv_freq [n_drivers*3] c_dbl_p, # drv_phi [n_drivers*3] c_dbl_p, # drv_eq [n_drivers*3] c_dbl_p, # drv_ncycles [n_drivers] c_int_p, # drv_has_period [n_drivers] ctypes.c_int, # n_frames c_dbl_p, # out_x c_dbl_p, # out_y c_dbl_p, # out_z c_dbl_p, # out_vx c_dbl_p, # out_vy c_dbl_p, # out_vz cb_type, # progress_cb (可为 NULL) ] def _c_dbl(arr: np.ndarray): """返回 float64 C 连续数组的 ctypes 指针。""" a = np.ascontiguousarray(arr, dtype=np.float64) return a.ctypes.data_as(ctypes.POINTER(ctypes.c_double)), a def _c_int(arr: np.ndarray): """返回 int32 C 连续数组的 ctypes 指针。""" a = np.ascontiguousarray(arr, dtype=np.int32) return a.ctypes.data_as(ctypes.POINTER(ctypes.c_int)), a def _is_python_module(lib) -> bool: """判断 lib 是否为 Python 引擎模块(而非 ctypes.CDLL)。""" return not isinstance(lib, ctypes.CDLL) def _run_dynamics_python(lib, config, atom_positions, atom_velocities, atom_masses, atom_fixed, bond_pairs, bond_stiffness, bond_rest_lengths, driver_data, atom_ids, progress_cb=None) -> dict: """调用 Python 引擎的 run_dynamics(),参数/返回值格式与 ctypes 版相同。""" n = len(atom_masses) NT = int(config["NT"]) NSTEP = int(config.get("NSTEP", 1)) warmup = int(config.get("warmup_steps", 0)) dt = float(config["DT"]) box_a = float(config.get("box_a", 300.0)) method_str = str(config.get("method", "leapfrog")).lower().replace(" ", "_") method_id = _METHOD_ID.get(method_str, 3) G = config.get("G", [0.0, 0.0, 0.0]) B = config.get("B", [0.0, 0.0, 0.0]) if hasattr(G, "tolist"): G = G.tolist() if hasattr(B, "tolist"): B = B.tolist() gravity_field = int(config.get("gravity_field", 0)) elastic_force = int(config.get("elastic_force", 1)) damping_force = int(config.get("damping_force", 0)) gravity_strength = float(config.get("gravity_strength", 1.0)) record_steps = NT - warmup n_frames = max(1, record_steps // NSTEP) nd = len(driver_data) if driver_data else 0 if nd > 0: drv_idx = np.array([d["local_idx"] for d in driver_data], dtype=np.int64) drv_amp = np.array([d["amp"] for d in driver_data], dtype=np.float64) drv_freq = np.array([d["freq"] for d in driver_data], dtype=np.float64) drv_phi = np.array([d["phi"] for d in driver_data], dtype=np.float64) drv_eq = np.array([d["eq_pos"] for d in driver_data], dtype=np.float64) drv_nc = np.array([d["n_cycles"] for d in driver_data], dtype=np.float64) drv_hp = np.array([d["has_period"]for d in driver_data], dtype=np.int32) else: drv_idx = drv_amp = drv_freq = drv_phi = drv_eq = drv_nc = drv_hp = \ np.zeros(0, dtype=np.int64) out_x, out_y, out_z, out_vx, out_vy, out_vz = lib.run_dynamics( n_atoms=n, pos_init=atom_positions, vel_init=atom_velocities, masses=atom_masses, fixed=atom_fixed, n_bonds=len(bond_pairs), bond_pairs=bond_pairs, bond_k=bond_stiffness, bond_r0=bond_rest_lengths, box_a=box_a, dt=dt, NT=NT, NSTEP=NSTEP, warmup_steps=warmup, method_id=method_id, Gx=float(G[0]), Gy=float(G[1]), Gz=float(G[2]), Bx=float(B[0]), By=float(B[1]), Bz=float(B[2]), gravity_field=gravity_field, elastic_force=elastic_force, damping_force=damping_force, gravity_strength=gravity_strength, n_drivers=nd, drv_idx=drv_idx, drv_amp=drv_amp, drv_freq=drv_freq, drv_phi=drv_phi, drv_eq=drv_eq, drv_ncycles=drv_nc, drv_has_period=drv_hp, n_frames=n_frames, progress_cb=progress_cb, ) shape = (n_frames, n) t_arr = np.arange(n_frames) * NSTEP * dt + warmup * dt return { "x": out_x.reshape(shape), "y": out_y.reshape(shape), "z": out_z.reshape(shape), "vx": out_vx.reshape(shape), "vy": out_vy.reshape(shape), "vz": out_vz.reshape(shape), "t": t_arr, } def run_dynamics_dll( lib, config: dict, atom_positions: np.ndarray, # (n_atoms, 3) atom_velocities: np.ndarray, # (n_atoms, 3) atom_masses: np.ndarray, # (n_atoms,) atom_fixed: np.ndarray, # (n_atoms, 3) int, 1=固定 bond_pairs: np.ndarray, # (n_bonds, 2) int 0-based 局部索引 bond_stiffness: np.ndarray, # (n_bonds,) bond_rest_lengths: np.ndarray,# (n_bonds,) driver_data: list, # 驱动原子列表(见下文) atom_ids: np.ndarray, # (n_atoms,) 全局 atom id(用于驱动原子查找) progress_cb=None, ) -> dict: """调用 DLL 的 run_dynamics(),返回抽帧后的轨迹数组。 driver_data 格式(每个元素对应一个驱动原子): { "atom_id": int, # 全局 atom id "local_idx": int, # 在 atom_ids 数组中的位置(0-based) "amp": [ax, ay, az], "freq": [fx, fy, fz], "phi": [px, py, pz], "eq_pos": [ex, ey, ez], "n_cycles": float, # 0=不限 "has_period": int, # 0/1 } 返回: { "x": np.ndarray (n_frames, n_atoms), "y": ..., "z": ..., "vx": ..., "vy": ..., "vz": ..., "t": np.ndarray (n_frames,), # 时间轴 } """ # Python 引擎:直接调用模块函数,不走 ctypes if _is_python_module(lib): return _run_dynamics_python( lib, config, atom_positions, atom_velocities, atom_masses, atom_fixed, bond_pairs, bond_stiffness, bond_rest_lengths, driver_data, atom_ids, progress_cb) n = len(atom_masses) NT = int(config["NT"]) NSTEP = int(config.get("NSTEP", 1)) warmup = int(config.get("warmup_steps", 0)) dt = float(config["DT"]) box_a = float(config.get("box_a", 300.0)) method_str = str(config.get("method", "leapfrog")).lower().replace(" ", "_") method_id = _METHOD_ID.get(method_str, 3) G = config.get("G", [0.0, 0.0, 0.0]) B = config.get("B", [0.0, 0.0, 0.0]) if hasattr(G, "tolist"): G = G.tolist() if hasattr(B, "tolist"): B = B.tolist() gravity_field = int(config.get("gravity_field", 0)) elastic_force = int(config.get("elastic_force", 1)) damping_force = int(config.get("damping_force", 0)) gravity_strength = float(config.get("gravity_strength", 1.0)) # ── 计算帧数 ────────────────────────────────────────────── record_steps = NT - warmup n_frames = max(1, record_steps // NSTEP) # ── 驱动原子数据 ────────────────────────────────────────── nd = len(driver_data) if driver_data else 0 if nd > 0: drv_idx_arr = np.array([d["local_idx"] for d in driver_data], dtype=np.int32) drv_amp_arr = np.array([d["amp"] for d in driver_data], dtype=np.float64).ravel() drv_freq_arr = np.array([d["freq"] for d in driver_data], dtype=np.float64).ravel() drv_phi_arr = np.array([d["phi"] for d in driver_data], dtype=np.float64).ravel() drv_eq_arr = np.array([d["eq_pos"] for d in driver_data], dtype=np.float64).ravel() drv_nc_arr = np.array([d["n_cycles"] for d in driver_data], dtype=np.float64) drv_hp_arr = np.array([d["has_period"] for d in driver_data], dtype=np.int32) else: drv_idx_arr = np.zeros(1, dtype=np.int32) drv_amp_arr = np.zeros(3, dtype=np.float64) drv_freq_arr = np.zeros(3, dtype=np.float64) drv_phi_arr = np.zeros(3, dtype=np.float64) drv_eq_arr = np.zeros(3, dtype=np.float64) drv_nc_arr = np.zeros(1, dtype=np.float64) drv_hp_arr = np.zeros(1, dtype=np.int32) # ── 输出缓冲区 ──────────────────────────────────────────── out_x = np.zeros(n_frames * n, dtype=np.float64) out_y = np.zeros(n_frames * n, dtype=np.float64) out_z = np.zeros(n_frames * n, dtype=np.float64) out_vx = np.zeros(n_frames * n, dtype=np.float64) out_vy = np.zeros(n_frames * n, dtype=np.float64) out_vz = np.zeros(n_frames * n, dtype=np.float64) # ── ctypes 指针(保留 arr 引用防止 GC) ────────────────── p_pos, _pos = _c_dbl(atom_positions.ravel()) p_vel, _vel = _c_dbl(atom_velocities.ravel()) p_mass, _mass = _c_dbl(atom_masses) p_fixed, _fixed = _c_int(atom_fixed.ravel()) p_bp, _bp = _c_int(bond_pairs.ravel() if len(bond_pairs) else np.zeros(2, dtype=np.int32)) p_bk, _bk = _c_dbl(bond_stiffness if len(bond_stiffness) else np.zeros(1)) p_br0, _br0 = _c_dbl(bond_rest_lengths if len(bond_rest_lengths) else np.zeros(1)) p_didx, _didx = _c_int(drv_idx_arr) p_damp, _damp = _c_dbl(drv_amp_arr) p_dfrq, _dfrq = _c_dbl(drv_freq_arr) p_dphi, _dphi = _c_dbl(drv_phi_arr) p_deq, _deq = _c_dbl(drv_eq_arr) p_dnc, _dnc = _c_dbl(drv_nc_arr) p_dhp, _dhp = _c_int(drv_hp_arr) p_ox = out_x.ctypes.data_as(ctypes.POINTER(ctypes.c_double)) p_oy = out_y.ctypes.data_as(ctypes.POINTER(ctypes.c_double)) p_oz = out_z.ctypes.data_as(ctypes.POINTER(ctypes.c_double)) p_ovx = out_vx.ctypes.data_as(ctypes.POINTER(ctypes.c_double)) p_ovy = out_vy.ctypes.data_as(ctypes.POINTER(ctypes.c_double)) p_ovz = out_vz.ctypes.data_as(ctypes.POINTER(ctypes.c_double)) # 进度回调 cb_type = ctypes.CFUNCTYPE(None, ctypes.c_int, ctypes.c_int) if progress_cb is not None: cb = cb_type(progress_cb) else: cb = ctypes.cast(None, cb_type) ret = lib.run_dynamics( n, p_pos, p_vel, p_mass, p_fixed, len(bond_pairs), p_bp, p_bk, p_br0, box_a, dt, NT, NSTEP, warmup, method_id, float(G[0]), float(G[1]), float(G[2]), float(B[0]), float(B[1]), float(B[2]), gravity_field, elastic_force, damping_force, gravity_strength, nd, p_didx, p_damp, p_dfrq, p_dphi, p_deq, p_dnc, p_dhp, n_frames, p_ox, p_oy, p_oz, p_ovx, p_ovy, p_ovz, cb, ) if ret != 0: raise RuntimeError(f"run_dynamics() returned error code {ret}") shape = (n_frames, n) t_arr = np.arange(n_frames) * NSTEP * dt + warmup * dt return { "x": out_x.reshape(shape), "y": out_y.reshape(shape), "z": out_z.reshape(shape), "vx": out_vx.reshape(shape), "vy": out_vy.reshape(shape), "vz": out_vz.reshape(shape), "t": t_arr, } def is_dll_available(engine: str = "c") -> bool: """检查指定引擎是否可用(DLL 已编译或 Python 模块存在)。""" if _ENGINE_DIR.get(engine, engine) == "python": return os.path.exists(os.path.join(_HERE, "python", "dynamics_lib.py")) return any(os.path.exists(p) for p in _dll_candidates(engine))