Official code for DuetDA: Decomposed and Dynamic Data Attribution with Model-State Gating for Accelerated Scientific Endeavors.
This repo supports:
- Building MatBench CGCNN splits (count-based OOD and difficulty-based OOD)
- Meta-training a DuetDA data valuator
- Training CGCNN/SchNet/ALIGNN with DuetDA-only data selection
Recommended: Python 3.10+.
python -m venv .venv
source .venv/bin/activate
pip install --upgrade pip
pip install torch numpy pandas scikit-learn pymatgen matbench wandb pyarrowpython prepare_cgcnn_count_based_ood.pyOutput (default):
cgcnn_data/matbench_log_kvrh_CountBasedOOD/
python prepare_cgcnn_difficulty.py \
--fit_csv fit_difficulty.csvOutput (default):
cgcnn_data/matbench_log_kvrh_difficultyOOD1/
python modules/meta_train_cgcnn.py \
--data_root cgcnn_data/matbench_log_kvrh_CountBasedOOD \
--checkpoint-dir checkpoints/meta_train/cgcnn \
--backbone cgcnn \
--fold 1 \
--num-models 3 \
--num-outer-steps 50 \
--truncation-steps 3 \
--inner-steps 5 \
--meta-lr 0.001 \
--batch-size 256Example checkpoint to use later:
checkpoints/meta_train/cgcnn/data_attributor_meta_step_*.pt
main.py now supports only --da-method duetda.
python main.py \
--data-root cgcnn_data/matbench_log_kvrh_difficultyOOD1 \
--data-name cgcnn_matbench \
--da-model-ckpt checkpoints/meta_train/cgcnn/data_attributor_meta_step_50.pt \
--model-name schnet \
--task-type regression \
--da-method duetda \
--fold 1 \
--epochs 2 \
--cuda \
--optim Adam \
--batch-size 256 \
--print-freq 5 \
--seed 42 \
--selection-ratio 0.5Optional W&B logging:
--use-wandb --wandb-entity <entity> --wandb-project <project> --wandb-name <run_name>- Best checkpoint:
checkpoints/<data>_duetda_<ratio>_model_best.pth.tar - Last checkpoint:
checkpoints/<data>_duetda_<ratio>_last.pth.tar - Test predictions:
<duetda>_<ratio>_<data>_test_results.csv(or--test-res-path)
- Use split
foldvalues consistent across data prep, meta-training, and final training. - If you train on CPU, remove
--cuda.