Graphormer模型批量推理脚本编写:高效处理千万级分子库
Graphormer模型批量推理脚本编写:高效处理千万级分子库
1. 引言
在药物发现和材料科学领域,处理千万级分子库已成为常态。传统单分子推理方式在面对如此庞大的数据量时显得力不从心,常常需要数周甚至更长时间才能完成计算。本文将带你从零开始,编写一个能够高效处理超大规模分子库的Graphormer推理脚本。
这个教程将聚焦于实际工程实现,而非理论细节。我们将使用Python构建一个完整的批量推理系统,包含多进程处理、批处理优化、内存映射文件等关键技术。学完本教程后,你将能够:
- 部署一个稳定高效的Graphormer推理环境
- 处理千万级SMILES列表而不耗尽内存
- 实时监控推理进度和资源使用情况
- 优雅地处理各种异常情况
- 将结果可靠地持久化存储
2. 环境准备与快速部署
2.1 Python环境安装
首先确保你的系统已安装Python 3.8或更高版本。推荐使用conda创建独立环境:
conda create -n graphormer python=3.8 conda activate graphormer2.2 依赖安装
Graphormer需要PyTorch作为后端。根据你的硬件选择安装命令:
# 仅CPU版本 pip install torch==1.12.0+cpu torchvision==0.13.0+cpu torchaudio==0.12.0 -f https://download.pytorch.org/whl/torch_stable.html # CUDA 11.3版本 pip install torch==1.12.0+cu113 torchvision==0.13.0+cu113 torchaudio==0.12.0 -f https://download.pytorch.org/whl/torch_stable.html然后安装Graphormer和其他必要依赖:
pip install graphormer rdkit tqdm numpy pandas3. 基础概念快速入门
3.1 Graphormer简介
Graphormer是一种基于Transformer架构的图神经网络,特别适合处理分子图数据。它将分子结构转化为图表示,然后利用自注意力机制捕获原子间的复杂关系。
3.2 批量推理的核心挑战
处理千万级分子库时,我们面临三个主要挑战:
- 内存限制:一次性加载所有分子会耗尽内存
- 计算效率:单进程处理速度太慢
- 稳定性:长时间运行容易因异常中断
我们的解决方案将针对这三个问题展开。
4. 分步实践操作
4.1 准备输入数据
假设我们有一个包含SMILES字符串的文本文件smi_list.txt,每行一个分子:
CCO CCN C1=CC=CC=C1 ...4.2 基础推理脚本
首先编写一个单进程的基础推理脚本:
from graphormer import GraphormerModel from rdkit import Chem from tqdm import tqdm import torch def load_model(): model = GraphormerModel.from_pretrained("graphormer-base") model.eval() return model def process_smiles(smiles, model): mol = Chem.MolFromSmiles(smiles) if mol is None: return None inputs = model.preprocess(mol) with torch.no_grad(): outputs = model(**inputs) return outputs def main(): model = load_model() with open("smi_list.txt") as f: smiles_list = [line.strip() for line in f] results = [] for smiles in tqdm(smiles_list): result = process_smiles(smiles, model) if result is not None: results.append(result) # 保存结果 torch.save(results, "results.pt") if __name__ == "__main__": main()这个基础版本虽然简单,但无法处理大规模数据。
5. 优化实现:生产级批量推理
5.1 内存映射文件处理
使用内存映射技术避免一次性加载所有SMILES:
import mmap def process_large_file(file_path, chunk_size=1000000): with open(file_path, "r+") as f: mm = mmap.mmap(f.fileno(), 0) start = 0 while True: chunk = mm[start:start+chunk_size] if not chunk: break # 处理当前chunk start += chunk_size5.2 多进程并行处理
利用Python的multiprocessing模块实现并行计算:
from multiprocessing import Pool, cpu_count import pandas as pd def process_chunk(args): chunk, model_path = args model = load_model(model_path) results = [] for smiles in chunk: result = process_smiles(smiles, model) if result is not None: results.append((smiles, result)) return results def parallel_processing(smiles_list, num_processes=None): if num_processes is None: num_processes = cpu_count() - 1 chunk_size = len(smiles_list) // num_processes chunks = [smiles_list[i:i+chunk_size] for i in range(0, len(smiles_list), chunk_size)] with Pool(num_processes) as pool: results = pool.map(process_chunk, [(chunk, "graphormer-base") for chunk in chunks]) # 合并结果 flat_results = [item for sublist in results for item in sublist] df = pd.DataFrame(flat_results, columns=["smiles", "result"]) return df5.3 进度监控与错误处理
添加进度监控和健壮的错误处理:
import time from datetime import datetime class ProgressLogger: def __init__(self, total): self.total = total self.start_time = time.time() self.processed = 0 self.errors = 0 def update(self, success=True): self.processed += 1 if not success: self.errors += 1 if self.processed % 1000 == 0: elapsed = time.time() - self.start_time remaining = (elapsed / self.processed) * (self.total - self.processed) print(f"[{datetime.now()}] Processed: {self.processed}/{self.total} | " f"Errors: {self.errors} | " f"ETA: {remaining/60:.1f} minutes") def safe_process_smiles(smiles, model, logger): try: result = process_smiles(smiles, model) logger.update(success=True) return result except Exception as e: logger.update(success=False) return None6. 完整生产级脚本
将上述组件整合成一个完整的生产级脚本:
import argparse import mmap import time from datetime import datetime from multiprocessing import Pool, cpu_count from pathlib import Path import pandas as pd import torch from graphormer import GraphormerModel from rdkit import Chem from tqdm import tqdm class GraphormerBatchProcessor: def __init__(self, model_path="graphormer-base", num_processes=None): self.model_path = model_path self.num_processes = num_processes or (cpu_count() - 1) def load_model(self): model = GraphormerModel.from_pretrained(self.model_path) model.eval() return model def process_smiles(self, smiles, model): mol = Chem.MolFromSmiles(smiles) if mol is None: return None inputs = model.preprocess(mol) with torch.no_grad(): outputs = model(**inputs) return outputs def process_chunk(self, args): chunk, chunk_id = args model = self.load_model() results = [] for smiles in chunk: result = self.process_smiles(smiles, model) if result is not None: results.append((smiles, result)) return results def process_file(self, input_path, output_path, chunk_size=1000000): # 读取文件确定总行数 total_lines = sum(1 for _ in open(input_path)) # 分块处理 results = [] with Pool(self.num_processes) as pool: with open(input_path) as f: chunks = [] current_chunk = [] for line in tqdm(f, total=total_lines, desc="Preparing chunks"): current_chunk.append(line.strip()) if len(current_chunk) >= chunk_size: chunks.append((current_chunk, len(chunks))) current_chunk = [] if current_chunk: chunks.append((current_chunk, len(chunks))) # 并行处理 for result in tqdm(pool.imap(self.process_chunk, chunks), total=len(chunks), desc="Processing chunks"): results.extend(result) # 保存结果 df = pd.DataFrame(results, columns=["smiles", "result"]) df.to_parquet(output_path) print(f"Results saved to {output_path}") if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--input", type=str, required=True, help="Input SMILES file") parser.add_argument("--output", type=str, required=True, help="Output file path") parser.add_argument("--processes", type=int, default=None, help="Number of processes") args = parser.parse_args() processor = GraphormerBatchProcessor(num_processes=args.processes) processor.process_file(args.input, args.output)7. 实用技巧与进阶
7.1 性能优化建议
- 批处理大小调整:根据GPU内存调整每批处理的分子数量
- 混合精度推理:使用
torch.cuda.amp进行混合精度计算 - 模型量化:对模型进行8位量化减少内存占用
7.2 错误处理增强
- 重试机制:对失败的计算添加自动重试
- 结果校验:检查输出结果的合理性
- 断点续传:记录已处理的位置,支持从中断处继续
7.3 结果后处理
- 结果分析:使用pandas进行统计分析
- 可视化:绘制关键指标的分布图
- 筛选:根据预测结果筛选候选分子
8. 总结
通过本教程,我们构建了一个完整的Graphormer批量推理系统,能够高效处理千万级分子库。核心优化包括内存映射文件处理、多进程并行计算、健壮的错误处理和进度监控。实际应用中,这个系统可以将原本需要数周的计算任务缩短到几小时内完成。
使用过程中,建议根据具体硬件配置调整进程数和批处理大小。对于特别大的分子库,可以考虑分布式计算框架如Ray或Dask进一步扩展处理能力。随着Graphormer模型的不断进化,这套系统也可以轻松适配新版本的模型。
获取更多AI镜像
想探索更多AI镜像和应用场景?访问 CSDN星图镜像广场,提供丰富的预置镜像,覆盖大模型推理、图像生成、视频生成、模型微调等多个领域,支持一键部署。
