Distilling ECFP4 Distance Structure into Musical Latent Space via Flow Matching
MusicMol-FM 将分子图编码为音乐 Score,使音乐距离对齐 ECFP4 距离结构,全程无需活性标签。
分子 G → GNN → hᵢ ∈ R²⁵⁶ → Flow Matching ODE → ŷᵢ ∈ [0,1]³ → music21.Score
↑
UMAP Kernel Alignment(教师:ECFP4 距离)
pip install -e . && pip install -r requirements.txt
# 准备数据
python scripts/prepare_data.py --method api --n_mols 1000 --out data/test.csv
# 训练
python -m musicmol_fm.train --config configs/fm_config.yaml --gpu_id 2
## 进一步微调
# ECFP4 对齐(无监督)
python -m musicmol_fm.tune \
--checkpoint checkpoints/fm/best.pt \
--train_csv data/CHEMBL210_train.csv \
--val_csv data/CHEMBL210_val.csv \
--test_csv data/CHEMBL210_test.csv \ # 可不传
--align_mode ecfp4 \
--out_dir checkpoints/tune/CHEMBL210
# 标签对齐(有监督)
python -m musicmol_fm.tune \
--checkpoint checkpoints/fm/best.pt \
--train_csv data/CHEMBL210_train.csv \
--val_csv data/CHEMBL210_val.csv \
--test_csv data/CHEMBL210_test.csv \
--align_mode label --label_col pIC50 \
--out_dir checkpoints/tune/CHEMBL210_label
# 批量微调 30 个靶标(你自己做好划分)
for target in CHEMBL210 CHEMBL301 CHEMBL344; do
python -m musicmol_fm.tune \
--checkpoint checkpoints/fm/best.pt \
--train_csv data/moleculeace/${target}_train.csv \
--val_csv data/moleculeace/${target}_val.csv \
--test_csv data/moleculeace/${target}_test.csv \
--align_mode label --label_col pIC50 \
--out_dir checkpoints/tune/${target}
done
# 推理:MIDI
python -m musicmol_fm.inference \
--checkpoint checkpoints/fm/best.pt \
--smiles "CCOC(=O)c1ccc2[nH]c(CN3CCN(CC3)c3ccc(OCc4ccc(Cl)cc4)cc3)nc2c1" "CCOC(=O)c1ccc2[nH]c(CN3CCN(CC3)c3ccc(SCc4ccc(Cl)cc4)cc3)nc2c1" --out_dir outputs/
python -m musicmol_fm.inference \
--checkpoint checkpoints/fm_ddp/best.pt \
--smiles_file data/input.csv --out_dir outputs/
# 推理:嵌入向量
# GNN 嵌入(256维,默认,适合 GNN-LII)
python -m musicmol_fm.inference \
--checkpoint checkpoints/fm_ddp/best.pt \
--smiles_file data/input.csv --mode embed \
--emb_source gnn --out_dir embd_output_gnn/
# → outputs/embeddings.csv: smiles + h_0 ... h_255
# 钢琴卷帘(88×64=5632维,适合音乐LII, n_time 可以修改128更精细)
python -m musicmol_fm.inference \
--checkpoint checkpoints/fm_ddp/best.pt \
--smiles_file data/input.csv --mode embed \
--emb_source pianoroll --n_pitch 88 --n_time 64 \
--out_dir embd_output_pianoroll/
# → outputs/embeddings.csv: smiles + pr_0 ... pr_5631
# 提取 Note 嵌入是音乐的平均"音调+节奏+力度"(3维,直接就是音符参数(pitch/dur/vel)的均值,每一维有明确的化学-音乐对应含义,没有经过额外的距离优化。适合:可视化、解释性分析)
python -m musicmol_fm.inference \
--checkpoint checkpoints/fm_ddp/best.pt \
--smiles_file data/input.csv --mode embed \
--emb_source note --out_dir embd_output_note/
# Option A: REST API(无需下载,~30 分钟/50万分子)
python scripts/prepare_data.py --method api --n_mols 500000 --out data/chembl_500k.csv
# Option B: 本地 SQLite(快速,需 ~500 MB)
wget https://ftp.ebi.ac.uk/pub/databases/chembl/ChEMBLdb/latest/chembl_36_sqlite.tar.gz
tar -xzf chembl_36_sqlite.tar.gz.tar.gz
# 2. 提取 + 过滤 + 规范化(约 5-10 分钟)
python scripts/prepare_data.py \
--method sqlite \
--db ./chembl_36/chembl_36_sqlite/chembl_36.db \
--out data/chembl_2m.csv \
--n_jobs 16
# 如果中途崩了,加 --resume 从断点继续
python scripts/prepare_data.py \
--method sqlite --db ./chembl_36/chembl_36_sqlite/chembl_36.db \
--out data/chembl_2m.csv --resume
# 3. 预处理成流式缓存(约 30-60 分钟,一次性)
python -m musicmol_fm.large_dataset prepare \
--csv data/chembl_2m.csv \
--cache_dir data/chembl_cache \
--n_jobs 16
过滤规则:MW 100–600 Da,重原子 ≤ 100,Ro5 违反 ≤ 1,RDKit 可解析。
# FM 模式
python -m musicmol_fm.train --config configs/fm_config.yaml --gpu_id 2
# FM 模式大规模训练,多卡模式
CUDA_VISIBLE_DEVICES=2,3,4,5,6,7 torchrun \
--nproc_per_node=6 \
--master_port=29500 \
-m musicmol_fm.train_multi_gpu \
--config configs/fm_config_large.yaml
# No-FM baseline
python -m musicmol_fm.train --config configs/no_fm_config.yaml
# 恢复训练
python -m musicmol_fm.train --config configs/fm_config.yaml \
--resume checkpoints/fm/latest.pt训练输出:
checkpoints/fm/
├── best.pt # 验证集最优
├── latest.pt # 最新(恢复用)
└── training_history.csv # 每 epoch 损失记录
training_history.csv 列:epoch, train_loss, val_loss, lr, lambda_kernel, lambda_reg
# SMILES → MIDI
python -m musicmol_fm.inference \
--checkpoint checkpoints/fm/best.pt \
--smiles "CCO" "c1ccccc1C(=O)O" \
--out_dir outputs/midi/
# CSV 批量 → MIDI
python -m musicmol_fm.inference \
--checkpoint checkpoints/fm/best.pt \
--smiles_file data/test.csv --smiles_col smiles \
--out_dir outputs/midi/
# → GNN 嵌入向量(256维)
python -m musicmol_fm.inference \
--checkpoint checkpoints/fm/best.pt \
--smiles_file data/test.csv \
--mode embed --emb_source gnn \
--out_dir outputs/emb/
# → 音符参数嵌入(3维:pitch/dur/vel)
python -m musicmol_fm.inference \
--checkpoint checkpoints/fm/best.pt \
--smiles_file data/test.csv \
--mode embed --emb_source note \
--out_dir outputs/emb/inference.py 参数:
| 参数 | 默认 | 说明 |
|---|---|---|
--checkpoint |
必填 | best.pt 路径 |
--mode |
music |
music=MIDI / embed=嵌入向量 |
--smiles |
— | 直接输入(与 --smiles_file 二选一) |
--smiles_file |
— | CSV 路径 |
--smiles_col |
smiles |
CSV SMILES 列名 |
--out_dir |
outputs |
输出目录 |
--device |
自动 | cuda / cpu |
--n_steps |
10 |
ODE 积分步数 |
--tempo |
108 |
BPM(music 模式) |
--batch_size |
64 |
推理批大小(embed 模式) |
--emb_source |
gnn |
gnn=256维 / note=3维 |
--pool |
mean |
原子→分子 pooling: mean / sum |
from musicmol_fm.inference import load_model, smiles_to_score, extract_emb
model, cfg = load_model("checkpoints/fm/best.pt", device="cuda")
# SMILES → Score
score = smiles_to_score("CCO", model, device="cuda")
score.write("midi", "ethanol.mid")
score.write("musicxml", "ethanol.xml")
# 批量嵌入
emb, valid = extract_emb(["CCO", "c1ccccc1"], model, device="cuda")
# emb.shape = (2, 256)
# LII 评估用距离矩阵
from sklearn.metrics.pairwise import cosine_distances
D = cosine_distances(emb)| 参数 | 默认 | 说明 |
|---|---|---|
csv_path |
— | SMILES CSV 路径 |
smiles_col |
smiles |
SMILES 列名 |
val_frac |
0.1 |
验证集比例 |
batch_size |
32 |
批大小(建议 16–64) |
num_workers |
4 |
DataLoader 工作进程数 |
max_atoms |
100 |
最大重原子数 |
cache_dir |
data/cache |
图数据缓存目录 |
| 参数 | 默认 | 说明 |
|---|---|---|
use_fm |
true |
true=Flow Matching / false=直接回归 |
node_dim |
256 |
GNN 隐藏层维度 |
n_gnn_layers |
4 |
MPNN 消息传递轮数(建议 3–6) |
t_dim |
128 |
时间步正弦嵌入维度(仅 FM) |
dropout |
0.1 |
GNN Dropout(推理时自动置 0) |
| 参数 | 默认 | 说明 |
|---|---|---|
kernel_type |
umap |
umap=局部自适应(推荐)/ rbf=全局高斯(更快) |
k_umap |
null |
UMAP 目标邻居数(null=自动 ⌊√B⌋,建议 5–20) |
选择建议:
umap:骨架多样性强(ChEMBL),密度自适应,无需调参rbf:快速实验,计算约快 2×
| 参数 | 默认 | 说明 |
|---|---|---|
epochs |
100 |
最大训练轮数 |
lr |
3e-4 |
AdamW 学习率 |
weight_decay |
1e-5 |
权重衰减 |
grad_clip |
1.0 |
梯度裁剪(0=不裁剪) |
seed |
42 |
随机种子 |
| 参数 | 默认 | 说明 |
|---|---|---|
lambda_fm |
1.0 |
FM 损失权重(固定) |
lambda_kernel_max |
1.0 |
Kernel 对齐最终权重 |
warmup_kernel_steps |
5000 |
λ_kernel 热启动步数(0→max) |
lambda_reg_init |
0.1 |
正则化初始权重 |
warmup_reg_steps |
10000 |
λ_reg 衰减步数(init→0) |
| 参数 | 默认 | 说明 |
|---|---|---|
checkpoint_dir |
checkpoints |
检查点根目录 |
patience |
20 |
早停耐心轮数 |
min_delta |
1e-5 |
最小改善量 |
resume |
null |
恢复训练路径(null=不恢复) |
| 参数 | 默认 | 说明 |
|---|---|---|
n_ode_steps_train |
5 |
训练时 ODE 步数(少=快) |
n_ode_steps_infer |
10 |
推理时 ODE 步数 |
log_every |
50 |
每 N 步打印日志 |
eval_every |
1 |
每 N epoch 验证 |
musicmol-fm/
├── musicmol_fm/
│ ├── __init__.py # 包导出
│ ├── dataset.py # MolDataset, collate_fn, smiles_to_batch
│ ├── model.py # GNNEncoder, VectorFieldMLP, MusicMolFM
│ ├── loss.py # umap_kernel, rbf_kernel, MusicMolLoss
│ ├── saver.py # CheckpointSaver(早停 + 原子写入)
│ ├── train.py # 训练主脚本(历史记录 → CSV)
│ ├── inference.py # smiles_to_score, extract_emb
│ └── utils.py # 日志, 随机种子, V6 规则目标
├── configs/
│ ├── fm_config.yaml # FM 模式(完整参数说明)
│ └── no_fm_config.yaml # No-FM baseline
├── scripts/
│ └── prepare_data.py # ChEMBL 数据准备
├── requirements.txt
├── setup.py
└── README.md
Direct distance alignment L = (D_music - D_ECFP4)² treats all molecule pairs equally. Kernel Alignment maps distances to similarities via RBF kernels:
K1[i,j] = exp(-D_ECFP4[i,j]² / 2σ₁²)
K2[i,j] = exp(-D_music[i,j]² / 2σ₂²)
L = mean((K1 - K2)²)
This focuses gradients on similar molecule pairs (small D, large K) — exactly where LII is most sensitive. Bandwidth σ is set via the median heuristic: σ = median(D) / √2 per batch, requiring no hyperparameter tuning.
The rule encoder (V6) produces fixed, deterministic note parameters. FM replaces this with a learned distribution conditioned on GNN atom features. Key benefits:
- Collision-free: FM can assign different pitches to chemically different atoms that happen to share the same hash bucket under V6
- Optimisable: FM parameters are differentiable with respect to the kernel alignment loss
- Generative: FM samples from a distribution, enabling stochastic exploration of the music-chemistry latent space
For the first ~2000 steps, λ_reg keeps the FM output close to V6's rule-based targets. This prevents random initialisation from disrupting the kernel loss. As training progresses, λ_reg → 0 and the model is free to improve beyond V6's capabilities.
@article{musicmol2024,
title = {MusicMol: Music as a Generative Latent Space for Molecular Representation},
author = {Your Name},
year = {2024},
}MIT