From 9a60ef2fbdf3686edf602ce90f65d1e8d33736d9 Mon Sep 17 00:00:00 2001 From: Piotr Czapla Date: Thu, 15 Nov 2018 23:40:10 +0100 Subject: [PATCH] Add assertions to train_clas and ability to limit the dataset size --- fastai_contrib/utils.py | 9 +++++++++ ulmfit/train_clas.py | 22 ++++++++++++++++------ 2 files changed, 25 insertions(+), 6 deletions(-) diff --git a/fastai_contrib/utils.py b/fastai_contrib/utils.py index cbccac9..6a2ba70 100644 --- a/fastai_contrib/utils.py +++ b/fastai_contrib/utils.py @@ -43,6 +43,15 @@ def get_texts(path): return np.array(texts), np.array(labels) +def ensure_paths_exists(*paths): + error = False + for path in paths: + if not path.exists(): + print(f'Error: {path} does not exist.') + error = True + if error: + raise FileNotFoundError("One or more required files cannot be found.") + def prepare_imdb(file_path: str, prepare_lm = False): """ function to extract aclImdb and combine into fastai standard format of labels and then text diff --git a/ulmfit/train_clas.py b/ulmfit/train_clas.py index 4f721bc..7809fe4 100644 --- a/ulmfit/train_clas.py +++ b/ulmfit/train_clas.py @@ -8,7 +8,7 @@ import pickle import torch from fastai.text import TextLMDataBunch, TextClasDataBunch, language_model_learner, text_classifier_learner from fastai import fit_one_cycle -from fastai_contrib.utils import PAD, UNK, read_clas_data, PAD_TOKEN_ID, DATASETS, TRN, VAL, TST +from fastai_contrib.utils import PAD, UNK, read_clas_data, PAD_TOKEN_ID, DATASETS, TRN, VAL, TST, ensure_paths_exists from fastai.text.transform import Vocab import fire @@ -19,7 +19,7 @@ from pathlib import Path def new_train_clas(data_dir, lang='en', cuda_id=0, pretrain_name='wt103', model_dir='models', qrnn=False, fine_tune=True, max_vocab=30000, bs=20, bptt=70, name='imdb-clas', - dataset='imdb'): + dataset='imdb', ds_pct=1.0): """ :param data_dir: The path to the `data` directory :param lang: the language unicode @@ -52,9 +52,13 @@ def new_train_clas(data_dir, lang='en', cuda_id=0, pretrain_name='wt103', model_ f'Error: Name of data directory should be data, not {data_dir.name}.' dataset_dir = data_dir / dataset model_dir = Path(model_dir) - assert data_dir.exists(), f'Error: {data_dir} does not exist.' - assert dataset_dir.exists(), f'Error: {dataset_dir} does not exist.' - assert model_dir.exists(), f'Error: {model_dir} does not exist.' + pretrained_fname = (f'lstm_{pretrain_name}', f'itos_{pretrain_name}') + + ensure_paths_exists(data_dir, + dataset_dir, + model_dir, + model_dir/f"{pretrained_fname[0]}.pth", + model_dir/f"{pretrained_fname[1]}.pkl") if qrnn: print('Using QRNNs...') @@ -96,6 +100,12 @@ def new_train_clas(data_dir, lang='en', cuda_id=0, pretrain_name='wt103', model_ print(f'Train size: {len(ids[TRN])}. Valid size: {len(ids[VAL])}. ' f'Test size: {len(ids[TST])}.') + if ds_pct < 1.0: + print(f"Makeing the dataset smaller {ds_pct}") + for split in [TRN, VAL, TST]: + ids[split] = ids[split][:int(len(ids[split])*ds_pct)] + + data_lm = TextLMDataBunch.from_ids(path=tmp_dir, vocab=vocab, train_ids=ids[TRN], valid_ids=ids[VAL], bs=bs, bptt=bptt) @@ -111,7 +121,7 @@ def new_train_clas(data_dir, lang='en', cuda_id=0, pretrain_name='wt103', model_ learn = language_model_learner( data_lm, bptt=bptt, emb_sz=emb_sz, nh=nh, nl=nl, qrnn=qrnn, pad_token=PAD_TOKEN_ID, - pretrained_fnames=(f'lstm_{pretrain_name}', f'itos_{pretrain_name}'), + pretrained_fnames=pretrained_fname, path=model_dir.parent, model_dir=model_dir.name) if fine_tune and not (model_dir / "enc.pth").exists():