- DTSemNet Top-k Regression
- RL experiments with continous actions.
Install package in same directory (installs gym, torch, stablebaselines3 etc.). Please install Anaconda if you don't have it already.
conda env create -f environment.yml # creates conda environment
conda activate dtsemnet # activates conda environmentOne datasets is given with the repo due to size constraint, others needs to downloaded and put into the /dataset directory. Please download it from HuggingFace.
Common flags: . --model dtregnet_topk or dtregnet_ste or dgt or cart . --dataset name under ./dataset/ . -s Number of simulations (different seeds, used 10 seeds for the experiments) . --output_prefix prefix for logs . --verbose print training logs . -g use GPU if available
python ./train/regression_train_topk.py --model dtregnet_topk --dataset ctslice -s 10 --output_prefix leaves_ct --verbose True -gpython ./train/regression_train_ste.py --model dtregnet_ste --dataset ctslice -s 10 --output_prefix leaves_ct --verbose True -g python ./train/regression_train_ste.py --model dgt --dataset ctslice -s 10 --output_prefix leaves_ct --verbose True -g > dgt_ct_leaf.log python ./train/regression_train_cart.py --model cart --dataset ctslice -s 10 --output_prefix leaves_ct --verbose True -g > cart_ct_leaf.logpython train/clean_rl_sac_train.py --config configs/rl/dtsemnet_topk/lunar_cont.json --seed 11OR
./parallel_seed_train_local.sh