Code accompanying the paper Reinforcement Learning for Topic Models in Findings of the Association for Computational Linguistics: ACL 2023.
conda create -n rl-for-topic-models python=3.9
conda install pytorch==1.11.0 torchvision==0.12.0 torchaudio==0.11.0 cudatoolkit=11.3 -c pytorch
pip install -r requirements.txt- update the configs in
model/decoder_network.pyandtrainer/config.pywith your training settings. - run
python train.py /path/to/data/pickle- add
--testif your data has a test subset
- add
- get raw data from https://trec.nist.gov/data/tweets
- put processed
Tweet.txtin thedata/raw/textsfolder - run
python data/dataset.py tweets2011
- get data from https://github.com/qiang2100/STTM
- put
StackOverflow.txtin thedata/raw/textsfolder - run
python data/dataset.py stackoverflow
- get data from https://github.com/qiang2100/STTM
- put
GoogleNews.txtin thedata/raw/textsfolder - run
python data/dataset.py googlenews
- get data from https://github.com/vinid/data
- put
dbpedia_sample_abstract_20k_unprep.txtin thedata/raw/texts folder - run
python data/dataset.py wiki20k
- run
python data/dataset.py 20ng
- get raw data from https://catalog.ldc.upenn.edu/LDC2008T19
- put tarball in
data/raw/nyt - run
pip install beautifulsoup4 lxml - run
python data/raw/nyt/nyt_untar.py - run
python data/dataset.py nytcorpus
- get raw data from https://github.com/nguyentthong/CLNTM
- put zipfile in
data/raw/clntm - run
python data/raw/clntm/scholar_unzip.py - run
python data/dataset.py contrastive
- get data from https://github.com/smutahoang/ntm
- put
*_preprocessed_data.txtindata/raw/texts
- run
python data/dataset.py 20ng --mwl 3
- run
python data/dataset.py custom --train_file /path/to/train/file --save_name /path/to/save/name- other arguments can be found at the bottom of
data/dataset.py
- other arguments can be found at the bottom of
- update search_dict in
run_experiments.pywith your hyperparameter search values. - run
python run_experiments.py experiment_name num_seeds --meta_seed meta_seed- meta seed will be randomly chosen from
random.randint(0, 2 ** 32)if not included as argument
- meta seed will be randomly chosen from
- run
python -m evals.compute_metrics topk num_experiments num_seeds /path/to/data/pickle /path/to/experiment
- run
python -m evals.dataset_stats /path/to/data/pickle
- run experiment with hyperparameters from paper
- run
python -m evals.dropout_sweep /path/to/dropout/experiment
- run
empirical_studies/examine_models.pyfrom https://github.com/smutahoang/ntm withtop_ks = [10] - move the
run.*.pklfiles into thentm_runsfolder - run experiments with hyperparameters from paper
- call them
ntm_20news_sweep,ntm_snippets_sweep,ntm_w2e_sweep, andntm_w2e_text_sweep
- call them
- run
python -m evals.plot_ntm_results
- run
python -m evals.save_topic_words topk /path/to/data/pickle /path/to/experiment/experiment_num/seeds/seed_num/seed_num_plotting_arrays.pkl