From ab9faa2ad9d1460cbe343dbb69134ac1ec06488d Mon Sep 17 00:00:00 2001 From: Piotr Czapla Date: Sat, 1 Dec 2018 13:42:10 +0100 Subject: [PATCH] Clean up Fire interface. --- tests/test_end_to_end.py | 8 ++--- ulmfit/__init__.py | 1 + ulmfit/__main__.py | 26 ++++++++++++++++ ulmfit/pretrain_lm.py | 65 +++++++++++++++++++--------------------- 4 files changed, 62 insertions(+), 38 deletions(-) create mode 100644 ulmfit/__main__.py diff --git a/tests/test_end_to_end.py b/tests/test_end_to_end.py index 1ac4e2b..7e4c56c 100644 --- a/tests/test_end_to_end.py +++ b/tests/test_end_to_end.py @@ -59,11 +59,11 @@ def test_ulmfit_default_end_to_end(): bs=2, name=lm_name) - exp.train_lm(num_lm_epochs=1) + exp.train_lm(num_epochs=1) #assert exp.results['accuracy'] > 0.02 - exp2 = ulmfit.train_clas.CLSHyperParams.based_on(exp.model_dir, test_data/'imdb') + exp2 = ulmfit.train_clas.CLSHyperParams.from_lm(test_data / 'imdb', exp.model_dir) exp2.train_cls(num_lm_epochs=0, unfreeze=False, bs=4,) def test_ulmfit_sentencepiece_end_to_end(): @@ -82,7 +82,7 @@ def test_ulmfit_sentencepiece_end_to_end(): bs=2, name=lm_name, ) - exp.train_lm(num_lm_epochs=1) + exp.train_lm(num_epochs=1) #assert exp.results['accuracy'] > 0.30 # NOTE: ds_pct is not available for sentencepiece -- tests are on the complete dataset @@ -90,5 +90,5 @@ def test_ulmfit_sentencepiece_end_to_end(): if __name__ == "__main__": - fire.Fire() # allows using all functions via CLI e.g. python utils.py prepare_imdb aclImdb.tgz + fire.Fire() # allows using all functions via CLI diff --git a/ulmfit/__init__.py b/ulmfit/__init__.py index e69de29..8b13789 100644 --- a/ulmfit/__init__.py +++ b/ulmfit/__init__.py @@ -0,0 +1 @@ + diff --git a/ulmfit/__main__.py b/ulmfit/__main__.py new file mode 100644 index 0000000..5c1684a --- /dev/null +++ b/ulmfit/__main__.py @@ -0,0 +1,26 @@ +from functools import wraps + +import fire +from .pretrain_lm import LMHyperParams +from .train_clas import CLSHyperParams + +class FireView: + def __init__(self, **kwargs): + for k,v in kwargs.items(): + setattr(self, k, v) + +class ULMFiT: + @wraps(LMHyperParams) + def lm(self, dataset_path, **changes): + changes['dataset_path'] = dataset_path + params = LMHyperParams(**changes) + return FireView(train=params.train_lm) + + lm2 = LMHyperParams + @wraps(CLSHyperParams) + def cls(self, dataset_path, baseon_path, **changes): + params = CLSHyperParams.from_lm(dataset_path, baseon_path, **changes) + return FireView(train=params.train_cls) + +if __name__ == '__main__': + fire.Fire(ULMFiT()) \ No newline at end of file diff --git a/ulmfit/pretrain_lm.py b/ulmfit/pretrain_lm.py index 3325ab5..eada9be 100644 --- a/ulmfit/pretrain_lm.py +++ b/ulmfit/pretrain_lm.py @@ -66,7 +66,7 @@ class LMHyperParams: bs: int = 70 lang: str = 'en' - name: str = '' + name: str = None cuda_id: InitVar[int] = 0 def __post_init__(self, cuda_id): @@ -78,8 +78,8 @@ class LMHyperParams: self.base_lm_path = Path(self.base_lm_path) if self.base_lm_path is not None else None assert self.dataset_path.exists() - self.cache_dir = self.dataset_path / 'models' / self.tok_name - self.model_dir = self.cache_dir / self.full_name + self.cache_dir = self.dataset_path / 'models' / self.tokenizer_prefix + self.model_dir = self.cache_dir / self.model_name self.model_dir.mkdir(exist_ok=True, parents=True) print('Batch size:', self.bs) @@ -88,14 +88,23 @@ class LMHyperParams: print('Model dir:', self.model_dir) self.dps = np.array(self.dps) if self.nh is None: self.nh = 1550 if self.qrnn else 1150 + if self.name is None: self.name = self.lang - @classmethod - def based_on(cls, base_lm_path, dataset_path, **kwargs) -> 'LMHyperParams': - with open(base_lm_path/'info.json', 'r') as f: d = json.load(f) - d['dataset_path'] = dataset_path - d['base_lm_path'] = base_lm_path - d.update(kwargs) - return cls(**d) + @property + def tokenizer_prefix(self): return f"{'sp' if self.subword else 'v'}{self.max_vocab // 1000}k" + + @property + def model_prefix(self): return ('bi' if self.bidir else '') + ('qrnn' if self.qrnn else 'lstm') + + @property + def model_name(self): return f"{self.model_prefix}_{self.name}.m" + + @property + def pretrained_fnames(self): return [self.base_lm_path / 'lm_best', self.base_lm_path / '../itos'] if self.base_lm_path else None + + @property + def lm_type(self): + return contrib_data.LanguageModelType.BiLM if self.bidir else contrib_data.LanguageModelType.FwdLM def save_info(self): from dataclasses import asdict @@ -105,38 +114,22 @@ class LMHyperParams: with (self.model_dir / 'info.json').open("w") as fp: json.dump(vals, fp) print("Saving info", self.model_dir / 'info.json') - @property - def tok_name(self): - pref = 'sp' if self.subword else 'v' - voc_size = self.max_vocab // 1000 - return f"{pref}{voc_size}k" - - @property - def full_name(self): return f"{self.model_name}_{self.name}.m" - - # todo rework - @property - def model_name(self): return ('bi' if self.bidir else '') + ('qrnn' if self.qrnn else 'lstm') - - @property - def pretrained_fnames(self): return [self.base_lm_path / 'lm_best', self.base_lm_path / '../itos'] if self.base_lm_path else None - - def train_lm(self, num_lm_epochs=10, data_lm=None): + def train_lm(self, num_epochs=10, data_lm=None): data_lm = self.load_wiki_data() if data_lm is None else data_lm learn = self.create_lm_learner(data_lm) - if num_lm_epochs > 0: + if num_epochs > 0: if self.pretrained_fnames : learn.fit_one_cycle(1, 1e-2, moms=(0.8, 0.7)) learn.unfreeze() - if num_lm_epochs > 0: learn.fit_one_cycle(num_lm_epochs, 1e-3, moms=(0.8, 0.7)) + if num_epochs > 0: learn.fit_one_cycle(num_epochs, 1e-3, moms=(0.8, 0.7)) else: try: learn.load("lm_best") print("Weights loaded") except FileNotFoundError: print("Starting from random weights") - learn.fit_one_cycle(num_lm_epochs, 5e-3, (0.8, 0.7), wd=1e-7) + learn.fit_one_cycle(num_epochs, 5e-3, (0.8, 0.7), wd=1e-7) opt_state_path = self.model_dir / 'opt_state.pth' print(f"Saving optimiser state at {opt_state_path}") torch.save(learn.opt.opt.state_dict(), opt_state_path) @@ -161,10 +154,6 @@ class LMHyperParams: learn.metrics = [accuracy_fwd, accuracy_bwd] if self.bidir else [accuracy] return learn - @property - def lm_type(self): - return contrib_data.LanguageModelType.BiLM if self.bidir else contrib_data.LanguageModelType.FwdLM - def load_wiki_data(self): trn_path = self.dataset_path / f'{self.lang}.wiki.train.tokens' val_path = self.dataset_path / f'{self.lang}.wiki.valid.tokens' @@ -222,6 +211,14 @@ class LMHyperParams: print('First 10 words in vocab:', ', '.join([itos[i] for i in range(10)])) return data_lm + @classmethod + def from_lm(cls, dataset_path, base_lm_path, **kwargs) -> 'LMHyperParams': + with open(base_lm_path/'info.json', 'r') as f: d = json.load(f) + d['dataset_path'] = dataset_path + d['base_lm_path'] = base_lm_path + d.update(kwargs) + return cls(**d) + def validate_lm(self): if not self.exp.subword and self.exp.max_vocab is None: raise NotImplementedError("figure out how to validate and save results")