From e6407fe0c914a9626a012a3c863a4dd78875b457 Mon Sep 17 00:00:00 2001 From: Tomasz Pietruszka Date: Thu, 10 Jan 2019 00:45:22 +0100 Subject: [PATCH] Added the param and model type for BwdLM --- ulmfit/pretrain_lm.py | 21 +++++++++++++++++++-- 1 file changed, 19 insertions(+), 2 deletions(-) diff --git a/ulmfit/pretrain_lm.py b/ulmfit/pretrain_lm.py index c9d3b47..3f70319 100644 --- a/ulmfit/pretrain_lm.py +++ b/ulmfit/pretrain_lm.py @@ -56,6 +56,7 @@ class LMHyperParams: dataset_path: str # data_dir base_lm_path: str = None + backwards: str = False bidir: bool =False qrnn: bool = True max_vocab: int = 60000 @@ -77,6 +78,8 @@ class LMHyperParams: cuda_id: InitVar[int] = 0 def __post_init__(self, cuda_id): + if self.bidir and self.backwards: + raise ValueError('Both "backwards" and "bidir" options cannot be enabled at the same time') if not torch.cuda.is_available(): print('CUDA not available. Setting device=-1.') cuda_id = -1 @@ -101,7 +104,16 @@ class LMHyperParams: def tokenizer_prefix(self): return f"{self.tokenizer.value}{self.max_vocab // 1000}k" @property - def model_prefix(self): return ('bi' if self.bidir else '') + ('qrnn' if self.qrnn else 'lstm') + def model_direction(self): + if self.bidir: + return 'bi' + if self.backwards: + return 'bwd' + else: + return '' + + @property + def model_prefix(self): return self.model_direction + ('qrnn' if self.qrnn else 'lstm') @property def model_name(self): return f"{self.model_prefix}_{self.name}.m" @@ -111,7 +123,12 @@ class LMHyperParams: @property def lm_type(self): - return contrib_data.LanguageModelType.BiLM if self.bidir else contrib_data.LanguageModelType.FwdLM + if self.bidir: + return contrib_data.LanguageModelType.BiLM + if self.backwards: + return contrib_data.LanguageModelType.BwdLM + else: + return contrib_data.LanguageModelType.FwdLM def tokenzier_to_fastai_args(self, trn_data_loading_func, add_moses): tok_func = MosesTokenizerFunc if add_moses else BaseTokenizer