mirror of
https://github.com/wassname/multifit.git
synced 2026-09-09 11:27:26 +08:00
Add command line convenience function & ability limit data set
So one can pretrain a language model from commandline The limit was added to support quick tests
This commit is contained in:
@@ -9,9 +9,6 @@ class Experiment:
|
|||||||
def new(self):
|
def new(self):
|
||||||
return {n: getattr(multifit.configurations,n) for n in multifit.configurations.__all__}
|
return {n: getattr(multifit.configurations,n) for n in multifit.configurations.__all__}
|
||||||
|
|
||||||
def load(self, model_path):
|
|
||||||
return multifit.ULMFiT().load_(Path(model_path))
|
|
||||||
|
|
||||||
def from_pretrained(self):
|
def from_pretrained(self):
|
||||||
return multifit.from_pretrained
|
return multifit.from_pretrained
|
||||||
|
|
||||||
|
|||||||
+24
-10
@@ -12,6 +12,8 @@ __all__ = [
|
|||||||
'multifit_lstm',
|
'multifit_lstm',
|
||||||
'multifit1152_lstm_nl3',
|
'multifit1152_lstm_nl3',
|
||||||
'multifit1152_lstm_nl3_fp16_large',
|
'multifit1152_lstm_nl3_fp16_large',
|
||||||
|
|
||||||
|
'multifit_mini_test',
|
||||||
]
|
]
|
||||||
|
|
||||||
def multifit1552_fp32(bs=64):
|
def multifit1552_fp32(bs=64):
|
||||||
@@ -25,7 +27,7 @@ def multifit1552_fp32(bs=64):
|
|||||||
bs=bs,
|
bs=bs,
|
||||||
use_adam_08=False,
|
use_adam_08=False,
|
||||||
early_stopping=None,
|
early_stopping=None,
|
||||||
name=_use_caller_name()
|
config_name=_use_caller_name()
|
||||||
)
|
)
|
||||||
self.arch.replace_(
|
self.arch.replace_(
|
||||||
tokenizer_type='fsp',
|
tokenizer_type='fsp',
|
||||||
@@ -42,29 +44,29 @@ def multifit1552_fp32(bs=64):
|
|||||||
multifit_fp32 = multifit1552_fp32
|
multifit_fp32 = multifit1552_fp32
|
||||||
|
|
||||||
def multifit_fp32_nl3():
|
def multifit_fp32_nl3():
|
||||||
return multifit1552_fp32().replace_(n_layers=3, name=_use_caller_name())
|
return multifit1552_fp32().replace_(n_layers=3, config_name=_use_caller_name())
|
||||||
|
|
||||||
# FP16
|
# FP16
|
||||||
|
|
||||||
def multifit1552_fp16():
|
def multifit1552_fp16():
|
||||||
return multifit1552_fp32(bs=128).replace_(fp16=True, name=_use_caller_name())
|
return multifit1552_fp32(bs=128).replace_(fp16=True, config_name=_use_caller_name())
|
||||||
|
|
||||||
def multifit1552_fp16_nl3_large():
|
def multifit1552_fp16_nl3_large():
|
||||||
return multifit1552_fp32(bs=448).replace_(fp16=True, n_layers=3, num_epochs=20, name=_use_caller_name())
|
return multifit1552_fp32(bs=448).replace_(fp16=True, n_layers=3, num_epochs=20, config_name=_use_caller_name())
|
||||||
|
|
||||||
multifit_fp16 = multifit1552_fp16
|
multifit_fp16 = multifit1552_fp16
|
||||||
|
|
||||||
def multifit_lstm():
|
def multifit_lstm():
|
||||||
return multifit1552_fp32(bs=128).replace_(qrnn=False, n_hid=1552, name=_use_caller_name())
|
return multifit1552_fp32(bs=128).replace_(qrnn=False, n_hid=1552, config_name=_use_caller_name())
|
||||||
|
|
||||||
def multifit1152_lstm_nl3(bs=128):
|
def multifit1152_lstm_nl3(bs=128):
|
||||||
return multifit1552_fp32(bs).replace_(qrnn=False, n_hid=1152, n_layers=3, name=_use_caller_name())
|
return multifit1552_fp32(bs).replace_(qrnn=False, n_hid=1152, n_layers=3, config_name=_use_caller_name())
|
||||||
|
|
||||||
def multifit1152_lstm_nl3_fp16_large():
|
def multifit1152_lstm_nl3_fp16_large():
|
||||||
return multifit1152_lstm_nl3(bs=448).replace_(fp16=True, num_epochs=20, name=_use_caller_name())
|
return multifit1152_lstm_nl3(bs=448).replace_(fp16=True, num_epochs=20, config_name=_use_caller_name())
|
||||||
|
|
||||||
def multifit_fp16_nl3():
|
def multifit_fp16_nl3():
|
||||||
return multifit1552_fp16().replace_(n_layers=3, name=_use_caller_name())
|
return multifit1552_fp16().replace_(n_layers=3, config_name=_use_caller_name())
|
||||||
|
|
||||||
def multifit_paper_version():
|
def multifit_paper_version():
|
||||||
self = ULMFiT()
|
self = ULMFiT()
|
||||||
@@ -81,7 +83,7 @@ def multifit_paper_version():
|
|||||||
early_stopping=None,
|
early_stopping=None,
|
||||||
clip=0.12,
|
clip=0.12,
|
||||||
dropout_values=dps,
|
dropout_values=dps,
|
||||||
name=_use_caller_name()
|
config_name=_use_caller_name()
|
||||||
)
|
)
|
||||||
self.arch.replace_(
|
self.arch.replace_(
|
||||||
tokenizer_type='sp',
|
tokenizer_type='sp',
|
||||||
@@ -99,7 +101,7 @@ def ulmfit_orig():
|
|||||||
self = multifit_paper_version()
|
self = multifit_paper_version()
|
||||||
self.replace_(
|
self.replace_(
|
||||||
seed=None,
|
seed=None,
|
||||||
name=_use_caller_name()
|
config_name=_use_caller_name()
|
||||||
)
|
)
|
||||||
self.arch.replace_(
|
self.arch.replace_(
|
||||||
tokenizer_type='f',
|
tokenizer_type='f',
|
||||||
@@ -110,6 +112,18 @@ def ulmfit_orig():
|
|||||||
)
|
)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
|
def multifit_mini_test():
|
||||||
|
self = multifit_paper_version()
|
||||||
|
self.replace_(
|
||||||
|
config_name=_use_caller_name(),
|
||||||
|
n_hid=240,
|
||||||
|
n_layers=2,
|
||||||
|
bs=40,
|
||||||
|
num_epochs=1,
|
||||||
|
fp16=False,
|
||||||
|
limit=100
|
||||||
|
)
|
||||||
|
return self
|
||||||
|
|
||||||
def _use_caller_name():
|
def _use_caller_name():
|
||||||
return inspect.stack()[1].function
|
return inspect.stack()[1].function
|
||||||
@@ -36,7 +36,6 @@ class Dataset:
|
|||||||
dataset_path: Path
|
dataset_path: Path
|
||||||
|
|
||||||
noise: float = 0.0
|
noise: float = 0.0
|
||||||
limit: int = None
|
|
||||||
|
|
||||||
ds_type: str = None
|
ds_type: str = None
|
||||||
lang: str = None
|
lang: str = None
|
||||||
@@ -96,6 +95,8 @@ class Dataset:
|
|||||||
use_lang_as_prefix=True)
|
use_lang_as_prefix=True)
|
||||||
else:
|
else:
|
||||||
self.read_data = read_clas_csv
|
self.read_data = read_clas_csv
|
||||||
|
self.lang = self._language_from_dataset_path()
|
||||||
|
self.uses_moses = False
|
||||||
self.trn_path = self.dataset_path / self.trn_name
|
self.trn_path = self.dataset_path / self.trn_name
|
||||||
self.val_path = self.dataset_path / self.val_name
|
self.val_path = self.dataset_path / self.val_name
|
||||||
self.tst_path = self.dataset_path / self.tst_name
|
self.tst_path = self.dataset_path / self.tst_name
|
||||||
@@ -161,11 +162,6 @@ class Dataset:
|
|||||||
trn_df = self._add_noise(trn_df, self.noise)
|
trn_df = self._add_noise(trn_df, self.noise)
|
||||||
val_df = self._add_noise(val_df, self.noise)
|
val_df = self._add_noise(val_df, self.noise)
|
||||||
|
|
||||||
if self.limit is not None:
|
|
||||||
print("Limiting data set to:", self.limit)
|
|
||||||
trn_df = trn_df[:self.limit]
|
|
||||||
val_df = val_df[:self.limit]
|
|
||||||
|
|
||||||
return trn_df, val_df, tst_df
|
return trn_df, val_df, tst_df
|
||||||
|
|
||||||
def load_unsupervised_data(self):
|
def load_unsupervised_data(self):
|
||||||
@@ -200,14 +196,16 @@ class ULMFiTDataset(Dataset):
|
|||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
self._vocab = None
|
self._vocab = None
|
||||||
|
|
||||||
def load_lm_databunch(self, bs, bptt):
|
def load_lm_databunch(self, bs, bptt, limit=None):
|
||||||
lm_suffix = str(bptt) if bptt != 70 else ""
|
lm_suffix = str(bptt) if bptt != 70 else ""
|
||||||
lm_suffix += "" if self.use_tst_for_lm else "-notst"
|
lm_suffix += "" if self.use_tst_for_lm else "-notst"
|
||||||
|
lm_suffix += "" if limit is None else f"-{limit}"
|
||||||
data_lm = self.load_n_cache_databunch(f"lm{lm_suffix}",
|
data_lm = self.load_n_cache_databunch(f"lm{lm_suffix}",
|
||||||
bunch_class=TextLMDataBunch,
|
bunch_class=TextLMDataBunch,
|
||||||
data_loader=self.load_unsupervised_data,
|
data_loader=self.load_unsupervised_data,
|
||||||
bptt=bptt,
|
bptt=bptt,
|
||||||
bs=bs)
|
bs=bs,
|
||||||
|
limit=limit)
|
||||||
|
|
||||||
with (self.cache_path / "itos.pkl").open('wb') as f:
|
with (self.cache_path / "itos.pkl").open('wb') as f:
|
||||||
pickle.dump(data_lm.vocab.itos, f)
|
pickle.dump(data_lm.vocab.itos, f)
|
||||||
@@ -223,24 +221,26 @@ class ULMFiTDataset(Dataset):
|
|||||||
self._vocab = self.load_lm_databunch(bs=20, bptt=70).vocab
|
self._vocab = self.load_lm_databunch(bs=20, bptt=70).vocab
|
||||||
return self._vocab
|
return self._vocab
|
||||||
|
|
||||||
def load_clas_databunch(self, bs):
|
def load_clas_databunch(self, bs, limit=None):
|
||||||
vocab = self._load_vocab()
|
vocab = self._load_vocab()
|
||||||
|
|
||||||
cls_name = "cls"
|
cls_name = "cls"
|
||||||
if self.limit is not None:
|
if limit is not None:
|
||||||
cls_name = f'{cls_name}limit{self.limit}'
|
cls_name = f'{cls_name}limit{limit}'
|
||||||
if self.noise > 0.0:
|
if self.noise > 0.0:
|
||||||
cls_name = f'{cls_name}noise{self.noise}'
|
cls_name = f'{cls_name}noise{self.noise}'
|
||||||
|
|
||||||
args = dict(vocab=vocab, bunch_class=TextClasDataBunch, bs=bs)
|
args = dict(vocab=vocab, bunch_class=TextClasDataBunch, bs=bs)
|
||||||
data_cls = self.load_n_cache_databunch(cls_name, data_loader=lambda: self.load_supervised_data()[:2], **args)
|
trn_val_dl = lambda: self.load_supervised_data()[:2]
|
||||||
|
data_cls = self.load_n_cache_databunch(cls_name, data_loader=trn_val_dl, limit=limit, **args)
|
||||||
# Hack to load test dataset with labels
|
# Hack to load test dataset with labels
|
||||||
data_tst = self.load_n_cache_databunch('tst', data_loader=lambda: self.load_supervised_data()[1:], **args)
|
val_tst_dl = lambda: self.load_supervised_data()[1:]
|
||||||
|
data_tst = self.load_n_cache_databunch('tst', data_loader=val_tst_dl, **args)
|
||||||
data_cls.test_dl = data_tst.valid_dl # data_tst.valid_dl holds test data
|
data_cls.test_dl = data_tst.valid_dl # data_tst.valid_dl holds test data
|
||||||
data_cls.lang = self.lang
|
data_cls.lang = self.lang
|
||||||
return data_cls
|
return data_cls
|
||||||
|
|
||||||
def load_n_cache_databunch(self, name, bunch_class, data_loader, bs, **args):
|
def load_n_cache_databunch(self, name, bunch_class, data_loader, bs, limit=None, **args):
|
||||||
bunch_path = self.cache_path / name
|
bunch_path = self.cache_path / name
|
||||||
databunch = None
|
databunch = None
|
||||||
if bunch_path.exists():
|
if bunch_path.exists():
|
||||||
@@ -251,6 +251,10 @@ class ULMFiTDataset(Dataset):
|
|||||||
if databunch is None:
|
if databunch is None:
|
||||||
print(f"Running tokenization: '{name}' ...")
|
print(f"Running tokenization: '{name}' ...")
|
||||||
train_df, valid_df = data_loader()
|
train_df, valid_df = data_loader()
|
||||||
|
if limit is not None:
|
||||||
|
print(f"Limiting number of examples in train and valid sets to: {limit}")
|
||||||
|
train_df = train_df[:limit]
|
||||||
|
valid_df = valid_df[:limit]
|
||||||
databunch = self.databunch_from_df(bunch_class, train_df, valid_df, **args)
|
databunch = self.databunch_from_df(bunch_class, train_df, valid_df, **args)
|
||||||
databunch.save(name)
|
databunch.save(name)
|
||||||
print(f"Data {name}, trn: {len(databunch.train_ds)}, val: {len(databunch.valid_ds)}")
|
print(f"Data {name}, trn: {len(databunch.train_ds)}, val: {len(databunch.valid_ds)}")
|
||||||
|
|||||||
+27
-7
@@ -41,6 +41,7 @@ class ULMFiTArchitecture(Params):
|
|||||||
tokenizer_type: str = "f"
|
tokenizer_type: str = "f"
|
||||||
max_vocab: int = 60000
|
max_vocab: int = 60000
|
||||||
lang: str = None
|
lang: str = None
|
||||||
|
config_name: str = None
|
||||||
|
|
||||||
emb_sz: int = awd_lstm_lm_config['emb_sz']
|
emb_sz: int = awd_lstm_lm_config['emb_sz']
|
||||||
n_hid: int = awd_lstm_lm_config['n_hid']
|
n_hid: int = awd_lstm_lm_config['n_hid']
|
||||||
@@ -130,7 +131,7 @@ class ULMFiTTrainingCommand(Params):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def model_name(self):
|
def model_name(self):
|
||||||
return (self.name or self.arch.model_name()) + (
|
return (self.name or self.arch.config_name or self.arch.model_name()) + (
|
||||||
"" if self.seed is None or self.seed == 0 or "seed" in self.name else f"seed{self.seed}")
|
"" if self.seed is None or self.seed == 0 or "seed" in self.name else f"seed{self.seed}")
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -161,7 +162,7 @@ class ULMFiTTrainingCommand(Params):
|
|||||||
exp_path = params.get('experiment_path', None)
|
exp_path = params.get('experiment_path', None)
|
||||||
if exp_path:
|
if exp_path:
|
||||||
fn = self.info_json
|
fn = self.info_json
|
||||||
print("Saving dump to", exp_path / fn)
|
print("Saving args to", exp_path / fn)
|
||||||
json_str = json.dumps(to_json_serializable(params), indent=2)
|
json_str = json.dumps(to_json_serializable(params), indent=2)
|
||||||
with (exp_path / fn).open("w") as f:
|
with (exp_path / fn).open("w") as f:
|
||||||
f.write(json_str)
|
f.write(json_str)
|
||||||
@@ -194,6 +195,12 @@ class ULMFiTTrainingCommand(Params):
|
|||||||
self.replace_(_verbose_diff=not silent, **d)
|
self.replace_(_verbose_diff=not silent, **d)
|
||||||
return arch
|
return arch
|
||||||
|
|
||||||
|
def train_(self, dataset_or_path, **kwargs):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def validate(self, **kwargs):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ULMFiTPretraining(ULMFiTTrainingCommand):
|
class ULMFiTPretraining(ULMFiTTrainingCommand):
|
||||||
@@ -210,6 +217,7 @@ class ULMFiTPretraining(ULMFiTTrainingCommand):
|
|||||||
clip: float = None
|
clip: float = None
|
||||||
fp16: bool = False
|
fp16: bool = False
|
||||||
lr: float = 5e-3
|
lr: float = 5e-3
|
||||||
|
limit: int = None
|
||||||
|
|
||||||
def get_learner(self, data_lm, **additional_trn_args):
|
def get_learner(self, data_lm, **additional_trn_args):
|
||||||
config = awd_lstm_lm_config.copy()
|
config = awd_lstm_lm_config.copy()
|
||||||
@@ -264,7 +272,7 @@ class ULMFiTPretraining(ULMFiTTrainingCommand):
|
|||||||
tokenizer = self.arch.new_tokenizer()
|
tokenizer = self.arch.new_tokenizer()
|
||||||
|
|
||||||
dataset = self._set_dataset_(dataset_or_path, tokenizer)
|
dataset = self._set_dataset_(dataset_or_path, tokenizer)
|
||||||
learn = self.get_learner(data_lm=dataset.load_lm_databunch(bs=self.bs, bptt=self.bptt))
|
learn = self.get_learner(data_lm=dataset.load_lm_databunch(bs=self.bs, bptt=self.bptt, limit=self.limit))
|
||||||
experiment_path = learn.path / learn.model_dir
|
experiment_path = learn.path / learn.model_dir
|
||||||
print("Experiment", experiment_path)
|
print("Experiment", experiment_path)
|
||||||
if self.num_epochs > 0:
|
if self.num_epochs > 0:
|
||||||
@@ -280,7 +288,7 @@ class ULMFiTPretraining(ULMFiTTrainingCommand):
|
|||||||
print("Language model saved to", self.experiment_path)
|
print("Language model saved to", self.experiment_path)
|
||||||
|
|
||||||
def validate(self):
|
def validate(self):
|
||||||
raise NotImplementedError("The validation on the language model is not implemented.")
|
return "not implemented"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def model_fnames(self):
|
def model_fnames(self):
|
||||||
@@ -346,6 +354,7 @@ class ULMFiTClassifier(ULMFiTTrainingCommand):
|
|||||||
seed: int = 0
|
seed: int = 0
|
||||||
bptt: int = 70
|
bptt: int = 70
|
||||||
fp16: bool = False
|
fp16: bool = False
|
||||||
|
limit: int = None
|
||||||
arch: ULMFiTArchitecture = None
|
arch: ULMFiTArchitecture = None
|
||||||
|
|
||||||
def get_learner(self, data_clas, eval_only=False, **additional_trn_args):
|
def get_learner(self, data_clas, eval_only=False, **additional_trn_args):
|
||||||
@@ -402,7 +411,7 @@ class ULMFiTClassifier(ULMFiTTrainingCommand):
|
|||||||
|
|
||||||
base_tokenizer = self.base.tokenizer
|
base_tokenizer = self.base.tokenizer
|
||||||
dataset = self._set_dataset_(dataset_or_path, base_tokenizer)
|
dataset = self._set_dataset_(dataset_or_path, base_tokenizer)
|
||||||
data_clas = dataset.load_clas_databunch(bs=self.bs)
|
data_clas = dataset.load_clas_databunch(bs=self.bs, limit=self.limit)
|
||||||
learn = self.get_learner(data_clas=data_clas)
|
learn = self.get_learner(data_clas=data_clas)
|
||||||
print(f"Training: {learn.path / learn.model_dir}")
|
print(f"Training: {learn.path / learn.model_dir}")
|
||||||
learn.unfreeze()
|
learn.unfreeze()
|
||||||
@@ -415,7 +424,6 @@ class ULMFiTClassifier(ULMFiTTrainingCommand):
|
|||||||
print("Classifier model saved to", self.experiment_path)
|
print("Classifier model saved to", self.experiment_path)
|
||||||
self.save_paramters()
|
self.save_paramters()
|
||||||
learn.destroy()
|
learn.destroy()
|
||||||
return
|
|
||||||
|
|
||||||
def _validate(self, learn, ds_type):
|
def _validate(self, learn, ds_type):
|
||||||
ds_name = ds_type.name.lower()
|
ds_name = ds_type.name.lower()
|
||||||
@@ -438,7 +446,7 @@ class ULMFiTClassifier(ULMFiTTrainingCommand):
|
|||||||
return json.load(fp)
|
return json.load(fp)
|
||||||
|
|
||||||
if data_cls is None:
|
if data_cls is None:
|
||||||
data_cls = self.dataset.load_clas_databunch(bs=self.bs)
|
data_cls = self.dataset.load_clas_databunch(bs=self.bs, limit=self.limit)
|
||||||
|
|
||||||
learn = self.get_learner(data_cls, eval_only=True)
|
learn = self.get_learner(data_cls, eval_only=True)
|
||||||
# avg = 'binary' if learn.data.c == 2 else 'macro'
|
# avg = 'binary' if learn.data.c == 2 else 'macro'
|
||||||
@@ -572,6 +580,18 @@ class ULMFiT:
|
|||||||
{self.classifier},
|
{self.classifier},
|
||||||
)""")
|
)""")
|
||||||
|
|
||||||
|
def train_(self, pretrain_dataset=None, clas_dataset=None):
|
||||||
|
results = {}
|
||||||
|
if pretrain_dataset is not None:
|
||||||
|
self.pretrain_lm.train_(pretrain_dataset)
|
||||||
|
results['pretrain_lm'] = self.pretrain_lm.validate()
|
||||||
|
if clas_dataset is not None:
|
||||||
|
self.finetune_lm.train_(clas_dataset)
|
||||||
|
results['finetune_lm'] = self.finetune_lm.validate(use_cache=None)
|
||||||
|
self.classifier.train_(clas_dataset)
|
||||||
|
results['classifier'] = self.classifier.validate(use_cache=None)
|
||||||
|
return results
|
||||||
|
|
||||||
def from_pretrained_(self, name, repo="n-waves/multifit-models"):
|
def from_pretrained_(self, name, repo="n-waves/multifit-models"):
|
||||||
name = name.rstrip(".tgz") # incase someone put's tgz name the name
|
name = name.rstrip(".tgz") # incase someone put's tgz name the name
|
||||||
url = f"https://github.com/{repo}/releases/download/{name}/{name}.tgz"
|
url = f"https://github.com/{repo}/releases/download/{name}/{name}.tgz"
|
||||||
|
|||||||
Reference in New Issue
Block a user