Skip to content

Latest commit

 

History

History
376 lines (344 loc) · 31.4 KB

File metadata and controls

376 lines (344 loc) · 31.4 KB

Experiment Configs

Here, we specify the commands that can be used to reproduce the experiments in our paper. The specified hyperparameters were obtained from a hyperparameter sweep as described in our publication.

Unbiased data

Alanine dipeptide

  • $1 \cdot 10^6$ training samples

    • forward KL:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.00001 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=0.0 training.training_mode.loss_weight_fwd_KL_p=1.0
      
    • forward KL + LDR-L2:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.00032 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000 training.training_mode.p=2
      
    • forward KL + LDR-L1:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.001 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.5 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000
      
  • $2 \cdot 10^6$ training samples

    • forward KL:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.000032 training.train_dataset_max_samples=2_000_000 training.training_mode.loss_weight_energy_supervision=0.0 training.training_mode.loss_weight_fwd_KL_p=1.0
      
    • forward KL + LDR-L2:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.00032 training.train_dataset_max_samples=2_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000 training.training_mode.p=2
      
    • forward KL + LDR-L1:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.00032 training.train_dataset_max_samples=2_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.3 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000
      
  • $5 \cdot 10^6$ training samples

    • forward KL:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.001 training.train_dataset_max_samples=5_000_000 training.training_mode.loss_weight_energy_supervision=0.0 training.training_mode.loss_weight_fwd_KL_p=1.0
      
    • forward KL + LDR-L2:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.0001 training.train_dataset_max_samples=5_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.3 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000 training.training_mode.p=2
      
    • forward KL + LDR-L1:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.00032 training.train_dataset_max_samples=5_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.5 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000
      

Alanine hexapeptide

  • $1 \cdot 10^6$ training samples

    • forward KL:
      python train.py -cn energy_supervision.yaml +system=hexa main_temp=300.0 training.lr=0.000032 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=0.0 training.training_mode.loss_weight_fwd_KL_p=1.0 training.batch_size=2048
      
    • forward KL + LDR-L2:
      python train.py -cn energy_supervision.yaml +system=hexa main_temp=300.0 training.lr=0.000032 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=1.3 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000 training.batch_size=2048 training.training_mode.p=2
      
    • forward KL + LDR-L1:
      python train.py -cn energy_supervision.yaml +system=hexa main_temp=300.0 training.lr=0.000032 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=1.3 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000 training.batch_size=2048
      
  • $2 \cdot 10^6$ training samples

    • forward KL:
      python train.py -cn energy_supervision.yaml +system=hexa main_temp=300.0 training.lr=0.000032 training.train_dataset_max_samples=2_000_000 training.training_mode.loss_weight_energy_supervision=0.0 training.training_mode.loss_weight_fwd_KL_p=1.0 training.batch_size=2048
      
    • forward KL + LDR-L2:
      python train.py -cn energy_supervision.yaml +system=hexa main_temp=300.0 training.lr=0.0001 training.train_dataset_max_samples=2_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=1.3 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000 training.batch_size=2048 training.training_mode.p=2
      
    • forward KL + LDR-L1:
      python train.py -cn energy_supervision.yaml +system=hexa main_temp=300.0 training.lr=0.0001 training.train_dataset_max_samples=2_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.9 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000 training.batch_size=2048
      
  • $5 \cdot 10^6$ training samples

    • forward KL:
      python train.py -cn energy_supervision.yaml +system=hexa main_temp=300.0 training.lr=0.0001 training.train_dataset_max_samples=5_000_000 training.training_mode.loss_weight_energy_supervision=0.0 training.training_mode.loss_weight_fwd_KL_p=1.0 training.batch_size=2048
      
    • forward KL + LDR-L2:
      python train.py -cn energy_supervision.yaml +system=hexa main_temp=300.0 training.lr=0.001 training.train_dataset_max_samples=5_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=1.1 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000 training.batch_size=2048 training.training_mode.p=2
      
    • forward KL + LDR-L1:
      python train.py -cn energy_supervision.yaml +system=hexa main_temp=300.0 training.lr=0.001 training.train_dataset_max_samples=5_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=1.1 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=20_000 training.batch_size=2048
      

2D GMM

The GMM experiments on unbiased data, including the comparison with the path-gradient forward KL, have been performed in a separate codebase (https://github.com/henrik-schopmans/unified-path-gradients), adapted from the Vaitl et al. codebase. To run the experiments, the unified-path-gradients package should be installed inside the LDR environment of our main codebase. Then, the experiments from the paper can be reproduced:

  • $500$ training samples

    • forward KL
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator ML --lr 1.e-8 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 50 --nsamples 500 --steps 10000
      
    • forward KL (path)
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator fastPathPQ --lr 0.000001 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 50 --nsamples 500 --steps 10000
      
    • forward KL + LDR-L2
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator ES --lr 0.0001 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 50 --nsamples 500 --steps 10000 --es-use-abs False --es-fwd-kl-weight 0.3 --es-detach-mean False
      
    • forward KL + LDR-L1
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator ES --lr 0.00032 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 50 --nsamples 500 --steps 10000 --es-use-abs True --es-fwd-kl-weight 0.3 --es-detach-mean False
      
  • $1000$ training samples

    • forward KL
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator ML --lr 0.000001 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 100 --nsamples 1000 --steps 10000
      
    • forward KL (path)
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator fastPathPQ --lr 0.00001 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 100 --nsamples 1000 --steps 10000 
      
    • forward KL + LDR-L2
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator ES --lr 0.0001 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 100 --nsamples 1000 --steps 10000 --es-use-abs False --es-fwd-kl-weight 0.3 --es-detach-mean False
      
    • forward KL + LDR-L1
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator ES --lr 0.00032 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 100 --nsamples 1000 --steps 10000 --es-use-abs True --es-fwd-kl-weight 0.5 --es-detach-mean False
      
  • $10000$ training samples

    • forward KL
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator ML --lr 0.00001 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 1000 --nsamples 10000 --steps 10000
      
    • forward KL (path)
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator fastPathPQ --lr 0.0001 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 1000 --nsamples 10000 --steps 10000
      
    • forward KL + LDR-L2
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator ES --lr 0.0001 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 1000 --nsamples 10000 --steps 10000 --es-use-abs False --es-fwd-kl-weight 0.1 --es-detach-mean False
      
    • forward KL + LDR-L1
      python mgm_train.py --mgm-scale 0.5 --gradient-estimator ES --lr 0.00032 --dim 2 --hidden 160 --n-coupling-layers 15 --n-blocks 1 --batch-size 1000 --nsamples 10000 --steps 10000 --es-use-abs True --es-fwd-kl-weight 0.3 --es-detach-mean False
      

Alanine dipeptide in Cartesian coordinates

  • $1 \cdot 10^5$ training samples
    • forward KL ($\chi_3$ correction, Tan et al. (2025))
      python train.py -cn forward_kl.yaml +system=aldp_pgafm_cartesian main_temp=300.0 training.train_dataset_max_samples=100_000 training.lr=0.0001 training.max_iter=400_000 training.max_grad_norm=1.0 system.system_specifics.COM_augmentation_use_chi3=True
      
    • forward KL (Improved Gaussian correction)
      python train.py -cn forward_kl.yaml +system=aldp_pgafm_cartesian main_temp=300.0 training.train_dataset_max_samples=100_000 training.lr=0.0001 training.max_iter=400_000 training.max_grad_norm=1.0
      
    • forward KL (path)
      python train.py -cn forward_kl_path.yaml +system=aldp_pgafm_cartesian main_temp=300.0 training.train_dataset_max_samples=100_000 training.lr=0.0001 training.max_iter=100_000 training.max_grad_norm=1.0
      
    • forward KL + LDR-L2:
      python train.py -cn energy_supervision.yaml +system=aldp_pgafm_cartesian main_temp=300.0 training.train_dataset_max_samples=100_000 training.lr=0.00032 training.max_iter=400_000 training.max_grad_norm=1.0 training.training_mode.loss_weight_energy_supervision=0.5 training.training_mode.loss_weight_fwd_KL_p=1.0 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=150_000 training.training_mode.p=2
      
    • forward KL + LDR-L1:
      python train.py -cn energy_supervision.yaml +system=aldp_pgafm_cartesian main_temp=300.0 training.train_dataset_max_samples=100_000 training.lr=0.00032 training.max_iter=400_000 training.max_grad_norm=1.0 training.training_mode.loss_weight_energy_supervision=1.3 training.training_mode.loss_weight_fwd_KL_p=1.0 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=150_000
      

Biased data

Pre-training on biased dataset:

python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=1.e-6 training.training_mode.loss_weight_energy_supervision=0.0 training.training_mode.loss_weight_fwd_KL_p=1.0 +overwrites=aldp_300K_biased.yaml

The final checkpoint *.pt obtained from pre-training needs to be used in place of <checkpoint_path> in the following self-refinement fine-tuning experiments.

Self-refinement fine-tuning:

  • $10^5$ IS samples
    • forward KL:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=1.e-7 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=0.0 training.training_mode.loss_weight_fwd_KL_p=1.0 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=100_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.train_fraction_from_p_hat=0.0
      
    • forward KL + LDR-L2 (LD only on IS samples):
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.000001 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=100_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=2 training.training_mode.p_hat_energy_supervision_rel_weight=0.0 training.training_mode.train_fraction_from_p_hat=0.0
      
    • forward KL + LDR-L2:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.00001 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=100_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.train_fraction_from_p_hat=0.5 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=2
      
    • forward KL + LDR-L1 (LD only on IS samples):
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.0000032 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=100_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=1 training.training_mode.p_hat_energy_supervision_rel_weight=0.0 training.training_mode.train_fraction_from_p_hat=0.0
      
    • forward KL + LDR-L1:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.00001 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=100_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.train_fraction_from_p_hat=0.5 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=1
      
  • $10^6$ IS samples
    • forward KL:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.000001 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=0.0 training.training_mode.loss_weight_fwd_KL_p=1.0 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=1_000_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.train_fraction_from_p_hat=0.0
      
    • forward KL + LDR-L2 (LD only on IS samples):
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.000032 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=1_000_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=2 training.training_mode.p_hat_energy_supervision_rel_weight=0.0 training.training_mode.train_fraction_from_p_hat=0.0
      
    • forward KL + LDR-L2:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.000032 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=1_000_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.train_fraction_from_p_hat=0.5 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=2
      
    • forward KL + LDR-L1 (LD only on IS samples):
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.000032 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=1_000_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=1 training.training_mode.p_hat_energy_supervision_rel_weight=0.0 training.training_mode.train_fraction_from_p_hat=0.0
      
    • forward KL + LDR-L1:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.000032 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=1_000_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.train_fraction_from_p_hat=0.5 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=1
      
  • $10^7$ IS samples
    • forward KL:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.00001 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=0.0 training.training_mode.loss_weight_fwd_KL_p=1.0 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=10_000_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.train_fraction_from_p_hat=0.0
      
    • forward KL + LDR-L2 (LD only on IS samples):
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.0001 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=10_000_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=2 training.training_mode.p_hat_energy_supervision_rel_weight=0.0 training.training_mode.train_fraction_from_p_hat=0.0
      
    • forward KL + LDR-L2:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.0001 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.1 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=10_000_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.train_fraction_from_p_hat=0.5 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=2
      
    • forward KL + LDR-L1 (LD only on IS samples):
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.0001 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.3 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=10_000_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=1 training.training_mode.p_hat_energy_supervision_rel_weight=0.0 training.training_mode.train_fraction_from_p_hat=0.0
      
    • forward KL + LDR-L1:
      python train.py -cn energy_supervision.yaml +system=aldp main_temp=300.0 training.lr=0.0001 training.train_dataset_max_samples=1_000_000 training.training_mode.loss_weight_energy_supervision=1.0 training.training_mode.loss_weight_fwd_KL_p=0.3 +overwrites=aldp_300K_biased.yaml training.checkpoint_path=<checkpoint_path> training.training_mode.p_data_source=IS_resampled training.training_mode.IS_NO_samples=10_000_000 training.training_mode.IS_resample_to=10_000_000 training.training_mode.train_fraction_from_p_hat=0.5 training.training_mode.loss_weight_energy_supervision_scheduler=cosine training.training_mode.loss_weight_energy_supervision_scheduler_T_max=50_000 training.training_mode.p=1
      

Variational energy-based training

FAB

  • Alanine dipeptide

    python train.py -cd configs/paper/aldp/ -cn fab_300K.yaml
    
  • Alanine hexapeptide

    python train.py -cd configs/paper/hexa/ -cn fab_300K.yaml
    

TA-BG

  • Reverse KL pre-training (Note: Requires additionally the 1200K validation data from https://doi.org/10.5281/zenodo.15526429)

    • Alanine dipeptide

      python train.py -cd configs/paper/aldp/ -cn rev_kl_1200K.yaml
      
    • Alanine hexapeptide

      python train.py -cd configs/paper/hexa/ -cn rev_kl_1200K.yaml
      

    The final checkpoint *.pt obtained from high-temperature reverse KL pre-training needs to be used in place of <checkpoint_path> in the annealing experiments below.

  • Annealing

    • Alanine dipeptide

      python train.py -cd configs/paper/aldp/ -cn annealing.yaml training.checkpoint_path=<checkpoint_path>
      
    • Alanine hexapeptide

      python train.py -cd configs/paper/hexa/ -cn annealing.yaml training.checkpoint_path=<checkpoint_path>
      

CMT (TA-BG + TR)

  • Alanine dipeptide

    • $1 \cdot 10^7$ target evals
      • CMT
        python train.py -cd configs/paper/aldp/ -cn cmt.yaml training.training_mode.n_samples_per_step=50_000 training.lr=0.0000032
        
      • CMT + LDR-L2
        python train.py -cd configs/paper/aldp/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=50_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=0.1 training.lr=0.00032 training.training_mode.energy_supervision_abs=False
        
      • CMT + LDR-L1
        python train.py -cd configs/paper/aldp/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=50_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=0.3 training.lr=0.000032
        
    • $2 \cdot 10^7$ target evals
      • CMT
        python train.py -cd configs/paper/aldp/ -cn cmt.yaml training.training_mode.n_samples_per_step=100_000 training.lr=0.00001
        
      • CMT + LDR-L2
        python train.py -cd configs/paper/aldp/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=100_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=0.1 training.lr=0.00032 training.training_mode.energy_supervision_abs=False
        
      • CMT + LDR-L1
        python train.py -cd configs/paper/aldp/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=100_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=0.3 training.lr=0.00032
        
    • $1 \cdot 10^8$ target evals
      • CMT
        python train.py -cd configs/paper/aldp/ -cn cmt.yaml training.training_mode.n_samples_per_step=500_000 training.lr=0.0001
        
      • CMT + LDR-L2
        python train.py -cd configs/paper/aldp/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=500_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=0.9 training.lr=0.0001 training.training_mode.energy_supervision_abs=False
        
      • CMT + LDR-L1
        python train.py -cd configs/paper/aldp/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=500_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=1.1 training.lr=0.0001
        
  • Alanine hexapeptide

    • $1 \cdot 10^8$ target evals
      • CMT
        python train.py -cd configs/paper/hexa/ -cn cmt.yaml training.training_mode.n_samples_per_step=250_000 training.lr=0.00001
        
      • CMT + LDR-L2
        python train.py -cd configs/paper/hexa/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=250_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=0.1 training.lr=0.0001 training.training_mode.energy_supervision_abs=False
        
      • CMT + LDR-L1
        python train.py -cd configs/paper/hexa/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=250_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=0.1 training.lr=0.0001
        
    • $2 \cdot 10^8$ target evals
      • CMT
        python train.py -cd configs/paper/hexa/ -cn cmt.yaml training.training_mode.n_samples_per_step=500_000 training.lr=0.000032
        
      • CMT + LDR-L2
        python train.py -cd configs/paper/hexa/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=500_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=0.5 training.lr=0.0001 training.training_mode.energy_supervision_abs=False
        
      • CMT + LDR-L1
        python train.py -cd configs/paper/hexa/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=500_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=0.3 training.lr=0.0001
        
    • $4 \cdot 10^8$ target evals
      • CMT
        python train.py -cd configs/paper/hexa/ -cn cmt.yaml training.training_mode.n_samples_per_step=1_000_000 training.lr=3.2e-5
        
      • CMT + LDR-L2
        python train.py -cd configs/paper/hexa/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=1_000_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=0.9 training.lr=0.0001 training.training_mode.energy_supervision_abs=False
        
      • CMT + LDR-L1
        python train.py -cd configs/paper/hexa/ -cn cmt.yaml training.training_mode.loss_name=energy_supervision training.training_mode.n_samples_per_step=1_000_000 training.training_mode.energy_supervision_loss_weight_fwd_KL=0.7 training.lr=0.0001