mirror of
https://github.com/wassname/multifit.git
synced 2026-09-09 11:27:26 +08:00
Fix validataion and add option to add noise to training labels
This commit is contained in:
@@ -113,17 +113,17 @@ class LMHyperParams:
|
||||
def lm_type(self):
|
||||
return contrib_data.LanguageModelType.BiLM if self.bidir else contrib_data.LanguageModelType.FwdLM
|
||||
|
||||
def tokenzier_to_fastai_args(self, trn_data_loading_func, add_moses):
|
||||
tok_func = MosesTokenizerFunc if add_moses else BaseTokenizer
|
||||
def tokenzier_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: # ensure we are using the same sentence piece model
|
||||
if self.base_lm_path and not(self.cache_dir/"spm.model").exists(): # ensure we are using the same sentence piece model
|
||||
shutil.copy(self.base_lm_path / '..' / 'itos.pkl', self.cache_dir)
|
||||
shutil.copy(self.base_lm_path / '..' / 'spm.model', self.cache_dir)
|
||||
shutil.copy(self.base_lm_path / '..' / 'spm.vocab', self.cache_dir)
|
||||
args = get_sentencepiece(self.cache_dir,
|
||||
trn_data_loading_func,
|
||||
sp_data_func,
|
||||
vocab_size=self.max_vocab,
|
||||
use_moses=add_moses,
|
||||
use_moses=use_moses,
|
||||
lang=self.lang)
|
||||
|
||||
elif self.tokenizer is Tokenizers.MOSES:
|
||||
@@ -207,7 +207,7 @@ class LMHyperParams:
|
||||
for path_ in [trn_path, val_path, tst_path]:
|
||||
assert path_.exists(), f'Error: {path_} does not exist.'
|
||||
|
||||
args = self.tokenzier_to_fastai_args(trn_data_loading_func=self.load_train_text, add_moses=False)
|
||||
args = self.tokenzier_to_fastai_args(sp_data_func=self.load_train_text, use_moses=False)
|
||||
try:
|
||||
data_lm = TextLMDataBunch.load(self.cache_dir, '.', lm_type=self.lm_type, bs=bs)
|
||||
print("Tokenized data loaded")
|
||||
|
||||
+38
-14
@@ -38,10 +38,10 @@ class CLSHyperParams(LMHyperParams):
|
||||
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,
|
||||
use_test_for_validation=False):
|
||||
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"
|
||||
|
||||
data_clas, data_lm, data_tst = self.load_cls_data(bs)
|
||||
data_clas, data_lm, data_tst = self.load_cls_data(bs, limit=limit, noise=noise)
|
||||
|
||||
if self.need_fine_tune_lm: self.train_lm(num_lm_epochs, data_lm=data_lm, true_wd=true_wd, drop_mult=drop_mul_lm)
|
||||
learn = self.create_cls_learner(data_clas, drop_mult=drop_mul_cls)
|
||||
@@ -61,7 +61,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
learn.freeze_to(-3)
|
||||
learn.fit_one_cycle(1, slice(5e-3 / (2.6 ** 4), 5e-3), moms=(0.8, 0.7))
|
||||
learn.unfreeze()
|
||||
learn.fit_one_cycle(2, slice(1e-3 / (2.6 ** 4), 1e-3), moms=(0.8, 0.7))
|
||||
learn.fit_one_cycle(num_cls_epochs, slice(1e-3 / (2.6 ** 4), 1e-3), moms=(0.8, 0.7))
|
||||
else:
|
||||
learn.true_wd = False
|
||||
print("Starting classifier training")
|
||||
@@ -72,16 +72,21 @@ class CLSHyperParams(LMHyperParams):
|
||||
learn.freeze_to(-3)
|
||||
learn.fit_one_cycle(1, slice(5e-4 / (2.6 ** 4), 5e-4), moms=(0.8, 0.7), wd=1e-7)
|
||||
learn.unfreeze()
|
||||
learn.fit_one_cycle(2, slice(1e-2 / (2.6 ** 4), 1e-2), moms=(0.8, 0.7), wd=1e-7)
|
||||
learn.fit_one_cycle(num_cls_epochs, slice(1e-2 / (2.6 ** 4), 1e-2), moms=(0.8, 0.7), wd=1e-7)
|
||||
print(f"Saving models at {learn.path / learn.model_dir}")
|
||||
learn.save('cls_last', with_opt=False)
|
||||
self.validate_cls(learn, data_tst, 'cls_last', bs=bs)
|
||||
self.validate_cls(learn, data_tst, 'cls_best', bs=bs)
|
||||
return learn
|
||||
|
||||
def validate_cls(self, learn, data_tst, save_name='cls_last', bs=40):
|
||||
self.validate_cls('cls_best', bs=bs, limit=limit, data_tst=data_tst, learn=learn)
|
||||
return None
|
||||
|
||||
def validate_cls(self, save_name='cls_last', limit=None, bs=40, data_tst=None, learn=None):
|
||||
if data_tst is None:
|
||||
_, _, data_tst = self.load_cls_data(bs, limit=limit)
|
||||
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.test_dl))
|
||||
print(f"Loss and accuracy using ({save_name}):", learn.validate(data_tst.valid_dl))
|
||||
|
||||
def create_cls_learner(self, data_clas, dps=None, **kwargs):
|
||||
fastai.text.learner.default_dropout['language'] = dps or self.dps
|
||||
@@ -101,6 +106,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
def load_cls_data(self, bs, **kwargs):
|
||||
add_trn_to_lm = True
|
||||
lang = self.lang
|
||||
use_moses = True
|
||||
if 'xnli' in str(self.dataset_dir):
|
||||
NotImplementedError("Support for Xnli is not implemented yet")
|
||||
if 'imdb' in self.dataset_dir.name:
|
||||
@@ -109,7 +115,10 @@ class CLSHyperParams(LMHyperParams):
|
||||
add_trn_to_lm = False # False as trn_df is contained in unsup already
|
||||
lang = self.lang
|
||||
|
||||
data = self.load_data(lang=lang, add_trn_to_lm=add_trn_to_lm,**kwargs)
|
||||
data = self.load_data(lang=lang,
|
||||
add_trn_to_lm=add_trn_to_lm,
|
||||
use_moses=use_moses,
|
||||
**kwargs)
|
||||
return self.databunches(bs, **data)
|
||||
|
||||
def load_data(self, lang='', **kwargs):
|
||||
@@ -134,13 +143,13 @@ class CLSHyperParams(LMHyperParams):
|
||||
kwargs.update(dict(trn_df=trn_df, val_df=val_df, tst_df=tst_df, unsup_df=unsup_df))
|
||||
return kwargs
|
||||
|
||||
def databunches(self, bs, trn_df, val_df, tst_df, unsup_df, add_trn_to_lm=True, force=False):
|
||||
def databunches(self, bs, trn_df, val_df, tst_df, unsup_df, add_trn_to_lm=True, use_moses=False, force=False, limit=None, noise=0.0):
|
||||
lm_trn_df = pd.concat([unsup_df, val_df, tst_df] + ([trn_df] if add_trn_to_lm else []))
|
||||
val_len = max(int(len(lm_trn_df) * 0.1), 2)
|
||||
lm_trn_df = lm_trn_df[val_len:]
|
||||
lm_val_df = lm_trn_df[:val_len]
|
||||
|
||||
args = self.tokenzier_to_fastai_args(trn_data_loading_func=lambda: trn_df[1], add_moses=True)
|
||||
args = self.tokenzier_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', lm_type=self.lm_type, bs=bs)
|
||||
@@ -152,9 +161,24 @@ class CLSHyperParams(LMHyperParams):
|
||||
print(f"Saving tokenized: cls.trn {len(data_lm.train_ds)}, cls.val {len(data_lm.valid_ds)}")
|
||||
data_lm.save('lm')
|
||||
|
||||
cls_name="."
|
||||
if limit is not None:
|
||||
print("Limiting data set to:", limit)
|
||||
trn_df = trn_df[:limit]
|
||||
val_df = val_df[:limit]
|
||||
cls_name=f'{cls_name}limit{limit}'
|
||||
if noise > 0.0:
|
||||
count = len(trn_df)
|
||||
labels = trn_df[0].unique()
|
||||
assert np.issubdtype(labels.dtype, np.integer), "noise only works on numerical numbers"
|
||||
modulo = labels.max()+1
|
||||
idx_to_distrub = np.random.permutation(count)[:int(count * noise)]
|
||||
trn_df.loc[idx_to_distrub, [0]] = (trn_df.loc[idx_to_distrub, [0]] + 1) % modulo
|
||||
print(f"Added noise to {len(idx_to_distrub)} examples, only {(count-len(idx_to_distrub))/count} have correct labels")
|
||||
cls_name = f'{cls_name}noise{noise}'
|
||||
try:
|
||||
if force: raise FileNotFoundError("Forcing reloading of caches")
|
||||
data_cls = TextClasDataBunch.load(self.cache_dir, '.', bs=bs)
|
||||
data_cls = TextClasDataBunch.load(self.cache_dir, cls_name, bs=bs)
|
||||
print(f"Tokenized data loaded, cls.trn {len(data_cls.train_ds)}, cls.val {len(data_cls.valid_ds)}")
|
||||
except FileNotFoundError:
|
||||
print(f"Running tokenization...")
|
||||
@@ -162,7 +186,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
data_cls = TextClasDataBunch.from_df(path=self.cache_dir, train_df=trn_df, valid_df=val_df,
|
||||
max_vocab=self.max_vocab, bs=bs, **args)
|
||||
print(f"Saving tokenized: cls.trn {len(data_cls.train_ds)}, cls.val {len(data_cls.valid_ds)}")
|
||||
data_cls.save('.')
|
||||
data_cls.save(cls_name)
|
||||
|
||||
# Hack to load test dataset with labels
|
||||
try:
|
||||
|
||||
Reference in New Issue
Block a user