Skip to content

Repository files navigation

MusicMol-FM

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

Python API

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

参数 默认 说明
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 调度

参数 默认 说明
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=不恢复)

ODE & 日志

参数 默认 说明
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

Design Philosophy

Why Kernel Alignment?

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.

Why Flow Matching?

The rule encoder (V6) produces fixed, deterministic note parameters. FM replaces this with a learned distribution conditioned on GNN atom features. Key benefits:

  1. Collision-free: FM can assign different pitches to chemically different atoms that happen to share the same hash bucket under V6
  2. Optimisable: FM parameters are differentiable with respect to the kernel alignment loss
  3. Generative: FM samples from a distribution, enabling stochastic exploration of the music-chemistry latent space

V6 as warm-start teacher

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.

Citation

@article{musicmol2024,
  title   = {MusicMol: Music as a Generative Latent Space for Molecular Representation},
  author  = {Your Name},
  year    = {2024},
}

License

MIT

About

No description, website, or topics provided.

Resources

Stars

7 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages