diff --git a/prepare_imdb.sh b/prepare_imdb.sh index 6d2488c..89815b7 100644 --- a/prepare_imdb.sh +++ b/prepare_imdb.sh @@ -6,6 +6,6 @@ mkdir -p "${DATA_DIR}" echo "Saving data in $DATA_DIR" wget -c "http://files.fast.ai/data/aclImdb.tgz" -P "${DATA_DIR}" -echo "Imdb is raw text so we are tokenizing it with Moses" -python -m fastai_contrib.utils prepare_imdb "${DATA_DIR}/aclImdb.tgz" --prepare_lm==False +echo "Imdb is raw text no preparation is done" +python -m fastai_contrib.utils prepare_imdb "${DATA_DIR}/aclImdb.tgz" diff --git a/tests/test_end_to_end.py b/tests/test_end_to_end.py index a5de301..014b6f1 100644 --- a/tests/test_end_to_end.py +++ b/tests/test_end_to_end.py @@ -40,6 +40,7 @@ def get_test_data(): copy_head(wt / 'en.wiki.train.tokens', test_wt / 'en.wiki.test.tokens', n=6*sz) copy_head(imdb / 'train.csv', test_imdb / 'train.csv', n=10*sz) copy_head(imdb / 'train.csv', test_imdb / 'test.csv', n=6*sz) + copy_head(imdb / 'train.csv', test_imdb / 'unsup.csv', n=1*sz) return test_data, test_wt diff --git a/ulmfit/train_clas.py b/ulmfit/train_clas.py index 91d75d1..c9b5826 100644 --- a/ulmfit/train_clas.py +++ b/ulmfit/train_clas.py @@ -99,6 +99,12 @@ class CLSHyperParams(LMHyperParams): def load_cls_data_imdb(self, bs): trn_df = pd.read_csv(self.dataset_path / 'train.csv', header=None) tst_df = pd.read_csv(self.dataset_path / 'test.csv', header=None) + unsp_df = pd.read_csv(self.dataset_path / 'unsup.csv', header=None) + + lm_trn_df = pd.concat([unsp_df, trn_df, tst_df]) + val_len = max(int(len(lm_trn_df) * 0.1), 2) + lm_trn_df = lm_trn_df[val_len:] + lm_val_df = lm_trn_df[:val_len] if self.use_test_for_validation: val_len = max(int(len(tst_df) * 0.1), 2) @@ -127,12 +133,9 @@ class CLSHyperParams(LMHyperParams): print(f"Tokenized data loaded, lm.trn {len(data_lm.train_ds)}, lm.val {len(data_lm.valid_ds)}") except FileNotFoundError: print(f"Running tokenization...") - - # wikitext is pretokenized with Moses - - data_lm = TextLMDataBunch.from_df(path=self.cache_dir, train_df=pd.concat([trn_df,tst_df]), - valid_df=val_df, test_df=tst_df, - lm_type=self.lm_type, max_vocab=self.max_vocab, **args) + data_lm = TextLMDataBunch.from_df(path=self.cache_dir, train_df=lm_trn_df, valid_df=lm_val_df, + max_vocab=self.max_vocab, bs=bs, lm_type=self.lm_type, **args) + print(f"Saving tokenized: cls.trn {len(data_lm.train_ds)}, cls.val {len(data_lm.valid_ds)}") data_lm.save('lm') print(f" cls.trn {len(data_lm.train_ds)}, cls.val {len(data_lm.valid_ds)}")