Skip to content
REAL-Lab-NUPublic

About

[KDD 2026 Oral] POLO: Preference-Guided Multi-Turn Reinforcement Learning for Sample-Efficient Lead Optimization

Topics

Resources

Stars

3 stars

Watchers

0 watching

Forks

Repository files navigation

POLO: Preference-Guided Multi-Turn Reinforcement Learning for Sample-Efficient Lead Optimization

Official code for our KDD 2026 paper. POLO uses Preference-Guided Policy Optimization (PGPO), combining trajectory-level PPO with a turn-level preference loss derived from the same molecular oracle feedback.

Installation

Use Python 3.11 and run commands from the repository root:

conda create -n polo python=3.11 pip -y
conda activate polo
python -m pip install -r requirements.txt
python -m pip install --no-deps PyTDC==1.1.15
python -m pip install flash-attn==2.7.4.post1 --no-build-isolation
python -m pip install --force-reinstall --no-deps scikit-learn==1.2.2
python -m pip install -e . --no-deps

The dependency stack uses vLLM 0.8.2 and PyTorch 2.6.0; see the vLLM installation guide for CUDA build requirements. Compatible veRL sources are included under verl/. PyTDC is installed without its optional dependency bundle to preserve the training dependencies above.

SFT Initialization

AlexanderWang915/MolOptIns is the SFT dataset. Fine-tune Qwen2.5 using LLaMA-Factory, or use AlexanderWang915/qwen2.5-1.5b-moloptins, the default policy and frozen-reference initialization.

Train

conda activate polo
TRAINER_GPUS=2,3 bash run_molecule_opt.sh \
    molecule_opt_task=qed \
    train_size=128 \
    trainer.total_training_steps=100

Append Hydra overrides to the same command, such as trainer.save_freq=20, trainer.default_local_dir=/path/to/run, or trainer.logger=[console]. Training checkpoints are saved under checkpoints/polo_experiments/polo-${train_size}-${molecule_opt_task}/global_step_N/. Resume explicitly with trainer.resume_mode=auto.

The default 1.5B full-parameter recipe uses two GPUs, five turns, one SMILES per turn, and an ECFP4 Tanimoto threshold of 0.4 relative to the original lead. Two-stage rollout filtering retains 50% of groups and 75% of their trajectories. Rollouts use TP=1 with one model replica per GPU. The per-GPU micro-batch is 2, and the global PPO mini-batch is 32.

Evaluate

Merge an FSDP actor checkpoint into a Hugging Face model directory:

python merge_fsdp_checkpoint.py \
    --base-model AlexanderWang915/qwen2.5-1.5b-moloptins \
    --checkpoint-dir /path/to/run/global_step_100/actor \
    --output-dir /path/to/run/global_step_100/actor/merged_hf

Evaluate the exported model on 200 test leads:

CUDA_VISIBLE_DEVICES=2,3 bash run_evaluate.sh \
    model_path=/path/to/run/global_step_100/actor/merged_hf \
    molecule_opt_task=qed \
    data.val_file=data/qed/test/qed_test_200.parquet \
    test_size=200 \
    num_rounds=10 \
    es_manager.val.group_size=32 \
    number_of_gpus=2 \
    actor_rollout_ref.rollout.tensor_model_parallel_size=1 \
    inference_output=results/polo-qed.json

Use test_size=100 to evaluate the first 100 leads from the same test file. Each lead uses ten rounds of 32 independent five-turn rollouts, restarting from the original molecule each round. Temperature increases from 0.9 by 0.1 per round, up to 2.0. With two GPUs and TP=1, two independent model replicas generate in parallel while one evaluator tracks all leads and aggregates metrics. For single-GPU inference, set CUDA_VISIBLE_DEVICES=2 and pass the override number_of_gpus=1.

The output JSON contains per-round and final improvement/absolute success rates, property improvements, similarity, and per-molecule best candidates. Oracle efficiency is reported after search at call thresholds from 100 to 2,000, separately for real evaluations and calls including cached reuse; these thresholds do not stop search. Initial scoring is excluded from trajectory call counts. The best candidate is retained even when it does not reach the success threshold. Run python evaluate.py --help to inspect the inference configuration and Hydra overrides.

Tasks and Data

Supported properties are qed, logp, sa, jnk3, and drd2. logp uses the TDC penalized LogP oracle. SA is minimized; the other properties are maximized. Combine property names with +, for example drd2+qed or qed+sa.

data/{task}/
    train/{task}_train_{size}.parquet
    val/{task}_val_32.parquet
    test/{task}_test_200.parquet

Training sets contain 32, 64, 128, 256, or 512 leads. Validation sets contain 32 leads; test sets of 200, 300, and 400 leads are included. Each Parquet file has a smiles column, and a matching JSON file is provided. TDC automatically downloads JNK3, DRD2, and SA oracle assets to oracle/ on first use and reuses them afterward. The initial download requires internet access and a writable working directory.

Citation

@inproceedings{wang2026polo,
  title={POLO: Preference-Guided Multi-Turn Reinforcement Learning for Sample-Efficient Lead Optimization},
  author={Wang, Ziqing and Wen, Yibo and Pattie, William and Luo, Xiao and Wu, Weimin and Hu, Jerry Yao-Chieh and Pandey, Abhishek and Liu, Han and Ding, Kaize},
  booktitle={Proceedings of the 32nd ACM SIGKDD Conference on Knowledge Discovery and Data Mining V.2},
  year={2026},
  doi={10.1145/3770855.3819016}
}

Built on RAGEN and veRL. See LICENSE and verl/LICENSE for licensing.

About

[KDD 2026 Oral] POLO: Preference-Guided Multi-Turn Reinforcement Learning for Sample-Efficient Lead Optimization

Topics

Resources

Stars

3 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages