mirror of
https://github.com/wassname/multifit.git
synced 2026-09-09 11:27:26 +08:00
Add ability to evalulate multiple models at once
This commit is contained in:
Executable
+4
@@ -0,0 +1,4 @@
|
||||
#!/usr/bin/env bash
|
||||
|
||||
LANGS
|
||||
for
|
||||
@@ -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())
|
||||
@@ -89,7 +89,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)
|
||||
@@ -147,6 +146,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)
|
||||
|
||||
@@ -201,6 +201,7 @@ 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'
|
||||
|
||||
+11
-8
@@ -37,9 +37,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):
|
||||
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)
|
||||
|
||||
@@ -54,7 +55,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))
|
||||
@@ -65,7 +66,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)
|
||||
@@ -76,17 +77,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):
|
||||
fastai.text.learner.default_dropout['language'] = dps or self.dps
|
||||
@@ -104,6 +106,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
|
||||
|
||||
Reference in New Issue
Block a user