mirror of
https://github.com/wassname/multifit.git
synced 2026-09-09 11:27:26 +08:00
Merge branch 'master' into use-configs
This commit is contained in:
@@ -1,14 +1,26 @@
|
||||
import gc
|
||||
import shutil
|
||||
from functools import wraps
|
||||
|
||||
import fire
|
||||
from .pretrain_lm import LMHyperParams
|
||||
from .train_clas import CLSHyperParams
|
||||
from pathlib import Path
|
||||
|
||||
class FireView:
|
||||
def __init__(self, **kwargs):
|
||||
for k,v in kwargs.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
def get_dataset_path(p):
|
||||
return [x for x in p.parents if x.name == "models"][0].parent
|
||||
|
||||
def get_lang_from_dataset_path(ds):
|
||||
lang,*_ = ds.name.split("-")
|
||||
if len(lang) == 2:
|
||||
return lang
|
||||
return "en"
|
||||
|
||||
class ULMFiT:
|
||||
@wraps(LMHyperParams)
|
||||
def lm(self, dataset_path, **changes):
|
||||
@@ -22,5 +34,23 @@ class ULMFiT:
|
||||
params = CLSHyperParams.from_lm(dataset_path, base_lm_path, **changes)
|
||||
return FireView(train=params.train_cls, validate_cls=params.validate_cls)
|
||||
|
||||
def eval(self, glob="mldoc/*-1/models/sp30k/lstm_nl4.m", name="tmp-100", cuda_id=0, **trn_params):
|
||||
results={}
|
||||
for base_model in Path("data").glob(glob):
|
||||
dataset_path = get_dataset_path(base_model)
|
||||
lang = get_lang_from_dataset_path(dataset_path)
|
||||
params = CLSHyperParams.from_lm(dataset_path, base_model, lang=lang, name=name, cuda_id=cuda_id)
|
||||
key = str(params.model_dir.relative_to(Path.cwd()))
|
||||
if params.model_dir.exists():
|
||||
print("Evaluating previously trained model")
|
||||
results[key] = params.validate_cls()[1]
|
||||
else:
|
||||
print("Training")
|
||||
results[key] = params.train_cls(num_lm_epochs=0, **trn_params)[1]
|
||||
params = None
|
||||
gc.collect()
|
||||
|
||||
print(list(sorted(results.items())))
|
||||
|
||||
if __name__ == '__main__':
|
||||
fire.Fire(ULMFiT())
|
||||
|
||||
@@ -53,7 +53,7 @@ def limit_vocab(unk_path, vocab):
|
||||
tokens = [''] + tokens
|
||||
line = ' '.join(tokens)
|
||||
f_out.write(line)
|
||||
print(f'{unk_path.name}. # of tokens: {total_num_tokens}')
|
||||
print(f'{unk_path.name}. # of tokens: {total_num_tokens}')
|
||||
temp_file_path.replace(unk_path)
|
||||
|
||||
|
||||
@@ -101,5 +101,6 @@ def postprocess_wikitext(path, lang):
|
||||
unk_path = dest_path / f'{lang}.wiki.{split}.tokens'
|
||||
limit_vocab(unk_path, vocab)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
fire.Fire(postprocess_wikitext)
|
||||
fire.Fire(postprocess_wikitext)
|
||||
|
||||
+28
-7
@@ -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
|
||||
@@ -71,12 +72,17 @@ class LMHyperParams:
|
||||
dps = dict(output_p=0.25, hidden_p=0.1, input_p=0.2, embed_p=0.02, weight_p=0.15) # consider removing dps & clip from the default hyperparams and put them to train
|
||||
clip: float = 0.12
|
||||
bptt: int = 70
|
||||
# alpha and beta - defaults like in fastai/text/learner.py:RNNLearner()
|
||||
rnn_alpha: float = 2 # activation regularization (AR)
|
||||
rnn_beta: float = 1 # temporal activation regularization (TAR)
|
||||
|
||||
lang: str = 'en'
|
||||
name: str = None
|
||||
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
|
||||
@@ -89,7 +95,6 @@ class LMHyperParams:
|
||||
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('Max vocab:', self.max_vocab)
|
||||
print('Cache dir:', self.cache_dir)
|
||||
print('Model dir:', self.model_dir)
|
||||
@@ -100,7 +105,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"
|
||||
@@ -110,9 +124,14 @@ 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, sp_data_func, use_moses):
|
||||
def tokenizer_to_fastai_args(self, sp_data_func, use_moses):
|
||||
tok_func = MosesTokenizerFunc if use_moses else BaseTokenizer
|
||||
if self.tokenizer is Tokenizers.SUBWORD:
|
||||
if self.base_lm_path and not(self.cache_dir/"spm.model").exists(): # ensure we are using the same sentence piece model
|
||||
@@ -146,6 +165,7 @@ class LMHyperParams:
|
||||
print("Saving info", self.model_dir / 'info.json')
|
||||
|
||||
def train_lm(self, num_epochs=20, data_lm=None, bs=70, true_wd=False, drop_mult=0.0, lr=5e-3):
|
||||
self.model_dir.mkdir(exist_ok=True, parents=True)
|
||||
data_lm = self.load_wiki_data(bs=bs) if data_lm is None else data_lm
|
||||
learn = self.create_lm_learner(data_lm, drop_mult=drop_mult)
|
||||
|
||||
@@ -168,7 +188,7 @@ class LMHyperParams:
|
||||
learn.unfreeze()
|
||||
if not learn.true_wd: learn.fit_one_cycle(num_epochs, lr, (0.8, 0.7), wd=1e-7)
|
||||
else: learn.fit_one_cycle(num_epochs, lr, (0.8, 0.7)) # TODO find proper values
|
||||
learn.save("lm_best_with_opt", with_opt=False)
|
||||
learn.save("lm_best_with_opt", with_opt=True)
|
||||
learn.save_encoder(ENC_BEST)
|
||||
learn.save(LM_BEST, with_opt=False)
|
||||
print(learn.path)
|
||||
@@ -182,7 +202,7 @@ class LMHyperParams:
|
||||
config = dict(emb_sz=self.emb_sz, n_hid=self.nh, n_layers=self.nl, pad_token=PAD_TOKEN_ID, qrnn=self.qrnn, bidir=self.bidir,
|
||||
tie_weights=True, out_bias=True)
|
||||
config.update(dps or self.dps)
|
||||
trn_args = dict(clip=self.clip)
|
||||
trn_args = dict(clip=self.clip, alpha=self.rnn_alpha, beta=self.rnn_beta)
|
||||
trn_args.update(kwargs)
|
||||
print ("Training args: ", trn_args, "dps: ", dps or self.dps)
|
||||
learn = language_model_learner(data_lm, AWD_LSTM, config=config, model_dir=self.model_dir.relative_to(data_lm.path), pretrained=False, **trn_args)
|
||||
@@ -210,13 +230,14 @@ class LMHyperParams:
|
||||
return [line.rstrip('\n') for line in f]
|
||||
|
||||
def load_wiki_data(self, bs=70):
|
||||
self.model_dir.mkdir(exist_ok=True, parents=True)
|
||||
trn_path = self.dataset_path / f'{self.lang}.wiki.train.tokens'
|
||||
val_path = self.dataset_path / f'{self.lang}.wiki.valid.tokens'
|
||||
tst_path = self.dataset_path / f'{self.lang}.wiki.test.tokens'
|
||||
for path_ in [trn_path, val_path, tst_path]:
|
||||
assert path_.exists(), f'Error: {path_} does not exist.'
|
||||
|
||||
args = self.tokenzier_to_fastai_args(sp_data_func=self.load_train_text, use_moses=False)
|
||||
args = self.tokenizer_to_fastai_args(sp_data_func=self.load_train_text, use_moses=False)
|
||||
try:
|
||||
data_lm = TextLMDataBunch.load(self.cache_dir, '.',
|
||||
bs=bs)
|
||||
|
||||
+12
-9
@@ -38,9 +38,10 @@ class CLSHyperParams(LMHyperParams):
|
||||
@property
|
||||
def need_fine_tune_lm(self): return not (self.model_dir/f"enc_best.pth").exists()
|
||||
|
||||
def train_cls(self, num_lm_epochs, unfreeze=True, bs=40, true_wd=True, drop_mul_lm=0.3, drop_mul_cls=0.5,
|
||||
def train_cls(self, num_lm_epochs, unfreeze=True, num_cls_frozen_epochs=1, bs=40, true_wd=True, drop_mul_lm=0.3, drop_mul_cls=0.5,
|
||||
use_test_for_validation=False, num_cls_epochs=2, limit=None, noise=0.0, cls_max_len=20*70):
|
||||
assert use_test_for_validation == False, "use_test_for_validation=True is not supported"
|
||||
self.model_dir.mkdir(exist_ok=True, parents=True)
|
||||
|
||||
data_clas, data_lm, data_tst = self.load_cls_data(bs, limit=limit, noise=noise)
|
||||
|
||||
@@ -55,7 +56,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
learn.true_wd = True
|
||||
print("Starting classifier training")
|
||||
learn.freeze_to(-1)
|
||||
learn.fit_one_cycle(1, 2e-2, moms=(0.8, 0.7))
|
||||
learn.fit_one_cycle(num_cls_frozen_epochs, 2e-2, moms=(0.8, 0.7))
|
||||
if unfreeze:
|
||||
learn.freeze_to(-2)
|
||||
learn.fit_one_cycle(1, slice(1e-2 / (2.6 ** 4), 1e-2), moms=(0.8, 0.7))
|
||||
@@ -66,7 +67,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
else:
|
||||
learn.true_wd = False
|
||||
print("Starting classifier training")
|
||||
learn.fit_one_cycle(1, 5e-2, moms=(0.8, 0.7), wd=1e-7)
|
||||
learn.fit_one_cycle(num_cls_frozen_epochs, 5e-2, moms=(0.8, 0.7), wd=1e-7)
|
||||
if unfreeze:
|
||||
learn.freeze_to(-2)
|
||||
learn.fit_one_cycle(1, slice(5e-2 / (2.6 ** 4), 5e-2), moms=(0.8, 0.7), wd=1e-7)
|
||||
@@ -77,17 +78,18 @@ class CLSHyperParams(LMHyperParams):
|
||||
print(f"Saving models at {learn.path / learn.model_dir}")
|
||||
learn.save('cls_last', with_opt=False)
|
||||
|
||||
self.validate_cls('cls_best', bs=bs, limit=limit, data_tst=data_tst, learn=learn)
|
||||
return None
|
||||
return self.validate_cls('cls_best', bs=bs, data_tst=data_tst, learn=learn)
|
||||
|
||||
def validate_cls(self, save_name='cls_last', limit=None, bs=40, data_tst=None, learn=None):
|
||||
def validate_cls(self, save_name='cls_last', bs=40, data_tst=None, learn=None):
|
||||
if data_tst is None:
|
||||
_, _, data_tst = self.load_cls_data(bs, limit=limit)
|
||||
_, _, data_tst = self.load_cls_data(bs)
|
||||
if learn is None:
|
||||
learn = self.create_cls_learner(data_tst, drop_mult=0.3)
|
||||
learn.unfreeze()
|
||||
learn.load(save_name)
|
||||
print(f"Loss and accuracy using ({save_name}):", learn.validate(data_tst.valid_dl))
|
||||
results = learn.validate(data_tst.valid_dl)
|
||||
print(f"Loss and accuracy using ({save_name}):", results)
|
||||
return list(map(float, results))
|
||||
|
||||
def create_cls_learner(self, data_clas, dps=None, **kwargs):
|
||||
assert self.bidir == False, "bidirectional model is not yet supported"
|
||||
@@ -110,6 +112,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
return learn
|
||||
|
||||
def load_cls_data(self, bs, **kwargs):
|
||||
self.model_dir.mkdir(exist_ok=True, parents=True)
|
||||
add_trn_to_lm = True
|
||||
lang = self.lang
|
||||
use_moses = True
|
||||
@@ -155,7 +158,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
lm_trn_df = lm_trn_df[val_len:]
|
||||
lm_val_df = lm_trn_df[:val_len]
|
||||
|
||||
args = self.tokenzier_to_fastai_args(sp_data_func=lambda: trn_df[1], use_moses=use_moses)
|
||||
args = self.tokenizer_to_fastai_args(sp_data_func=lambda: trn_df[1], use_moses=use_moses)
|
||||
try:
|
||||
if force: raise FileNotFoundError("Forcing reloading of caches")
|
||||
data_lm = TextLMDataBunch.load(self.cache_dir, 'lm', bs=bs)
|
||||
|
||||
Reference in New Issue
Block a user