mirror of
https://github.com/wassname/multifit.git
synced 2026-09-09 11:27:26 +08:00
Tweak hyper training params of cls (drop_mul, bs, true_wd=True)
I've set the same hyperparams as in lesson3
This commit is contained in:
+37
-25
@@ -41,6 +41,10 @@ import fastai_contrib.data as contrib_data
|
||||
# :param model_dir: The path to the directory where the models should be saved
|
||||
# :param bidir: whether the language model is bidirectional
|
||||
# """
|
||||
LM_BEST = "lm_best"
|
||||
ENC_BEST = "enc_best"
|
||||
|
||||
|
||||
class Tokenizers(Enum):
|
||||
SUBWORD='sb'
|
||||
MOSES='v'
|
||||
@@ -127,44 +131,52 @@ class LMHyperParams:
|
||||
with (self.model_dir / 'info.json').open("w") as fp: json.dump(vals, fp)
|
||||
print("Saving info", self.model_dir / 'info.json')
|
||||
|
||||
def train_lm(self, num_epochs=10, data_lm=None):
|
||||
def train_lm(self, num_epochs=20, data_lm=None, true_wd=False, drop_mult=0.1):
|
||||
data_lm = self.load_wiki_data() if data_lm is None else data_lm
|
||||
learn = self.create_lm_learner(data_lm)
|
||||
learn = self.create_lm_learner(data_lm, drop_mult=drop_mult)
|
||||
|
||||
learn.true_wd = true_wd
|
||||
try:
|
||||
learn.load("lm_best_with_opt")
|
||||
print("Continuing training")
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
if num_epochs > 0:
|
||||
if self.pretrained_fnames :
|
||||
learn.fit_one_cycle(1, 1e-2, moms=(0.8, 0.7)) # TODO Fix the learning rates
|
||||
learn.unfreeze()
|
||||
if num_epochs > 0: learn.fit_one_cycle(num_epochs, 1e-3, moms=(0.8, 0.7))
|
||||
if self.pretrained_fnames or self.pretrained_model:
|
||||
print("Training lm from: ", self.pretrained_fnames or self.pretrained_model)
|
||||
if learn.true_wd:
|
||||
learn.fit_one_cycle(1, 1e-2, moms=(0.8, 0.7))
|
||||
learn.unfreeze()
|
||||
learn.fit_one_cycle(num_epochs, 1e-3, moms=(0.8, 0.7))
|
||||
else:
|
||||
learn.fit_one_cycle(1, 1e-2, moms=(0.8, 0.7), wd=1e-7) # TODO Fix the learning rates
|
||||
learn.unfreeze()
|
||||
learn.fit_one_cycle(num_epochs, 1e-3, moms=(0.8, 0.7), wd=1e-7)
|
||||
else:
|
||||
try:
|
||||
learn.load("lm_best")
|
||||
print("Weights loaded")
|
||||
except FileNotFoundError:
|
||||
print("Starting from random weights")
|
||||
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)
|
||||
learn.save_encoder("enc_best")
|
||||
learn.save("lm_best", with_opt=False)
|
||||
print("Training lm from random weights")
|
||||
if not learn.true_wd: learn.fit_one_cycle(num_epochs, 5e-3, (0.8, 0.7), wd=1e-7)
|
||||
else: learn.fit_one_cycle(num_epochs, 5e-3, (0.8, 0.7)) # TODO find proper values
|
||||
learn.save("lm_best_with_opt", with_opt=False)
|
||||
learn.save_encoder(ENC_BEST)
|
||||
learn.save(LM_BEST, with_opt=False)
|
||||
print(learn.path)
|
||||
|
||||
self.save_info()
|
||||
return learn
|
||||
|
||||
def create_lm_learner(self, data_lm):
|
||||
fastai.text.learner.default_dropout['language'] = self.dps
|
||||
def create_lm_learner(self, data_lm, dps=None, **kwargs):
|
||||
fastai.text.learner.default_dropout['language'] = dps or self.dps
|
||||
lm_learner = bilm_learner if self.bidir else language_model_learner
|
||||
|
||||
learn = lm_learner(data_lm, bptt=self.bptt, emb_sz=self.emb_sz, nh=self.nh, nl=self.nl, pad_token=PAD_TOKEN_ID,
|
||||
drop_mult=self.drop_mult, tie_weights=True, model_dir= self.model_dir.relative_to(data_lm.path),
|
||||
bias=True, qrnn=self.qrnn, clip=self.clip, pretrained_fnames=self.pretrained_fnames,
|
||||
pretrained_model=self.pretrained_model)
|
||||
trn_args = dict(drop_mult=self.drop_mult, tie_weights=True, clip=self.clip, bptt=self.bptt,
|
||||
pretrained_fnames=self.pretrained_fnames,
|
||||
pretrained_model=self.pretrained_model)
|
||||
trn_args.update(kwargs)
|
||||
print ("Training args: ", trn_args, "dps: ", dps)
|
||||
learn = lm_learner(data_lm, emb_sz=self.emb_sz, nh=self.nh, nl=self.nl, pad_token=PAD_TOKEN_ID,
|
||||
bias=True, qrnn=self.qrnn, model_dir=self.model_dir.relative_to(data_lm.path), **trn_args)
|
||||
# compared to standard Adam, we set beta_1 to 0.8
|
||||
learn.opt_fn = partial(optim.Adam, betas=(0.8, 0.99))
|
||||
learn.true_wd = False
|
||||
print("true_wd: ", learn.true_wd)
|
||||
learn.metrics = [accuracy_fwd, accuracy_bwd] if self.bidir else [accuracy]
|
||||
return learn
|
||||
|
||||
|
||||
+49
-39
@@ -25,7 +25,8 @@ import fire
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
|
||||
from ulmfit.pretrain_lm import LMHyperParams, Tokenizers
|
||||
from ulmfit.pretrain_lm import LMHyperParams, Tokenizers, ENC_BEST
|
||||
|
||||
|
||||
class MosesTokenizerFunc(BaseTokenizer):
|
||||
"Wrapper around a MosesTokenizer to make it a `BaseTokenizer`."
|
||||
@@ -50,53 +51,61 @@ 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=70):
|
||||
def train_cls(self, num_lm_epochs, unfreeze=True, bs=40, true_wd=True, drop_mul_lm=0.3, drop_mul_cls=0.5):
|
||||
data_clas, data_lm = self.load_cls_data(bs)
|
||||
|
||||
if self.need_fine_tune_lm: self.train_lm(num_lm_epochs, data_lm=data_lm)
|
||||
learn = self.create_cls_learner(data_clas)
|
||||
|
||||
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)
|
||||
try:
|
||||
learn.load('cls_last')
|
||||
print("Loading last classifier")
|
||||
except FileNotFoundError:
|
||||
learn.load_encoder("enc_best")
|
||||
|
||||
learn.true_wd = False
|
||||
print("Starting classifier training")
|
||||
learn.fit_one_cycle(1, 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)
|
||||
|
||||
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.load_encoder(ENC_BEST)
|
||||
if true_wd:
|
||||
learn.true_wd = True
|
||||
print("Starting classifier training")
|
||||
learn.freeze_to(-1)
|
||||
learn.fit_one_cycle(1, 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))
|
||||
learn.freeze_to(-3)
|
||||
learn.fit_one_cycle(2, slice(1e-3 / (2.6 ** 4), 1e-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))
|
||||
else:
|
||||
learn.true_wd = False
|
||||
print("Starting classifier training")
|
||||
learn.fit_one_cycle(1, 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)
|
||||
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)
|
||||
print(f"Saving models at {learn.path / learn.model_dir}")
|
||||
learn.save('cls_last', with_opt=False)
|
||||
return learn
|
||||
|
||||
def create_cls_learner(self, data_clas):
|
||||
fastai.text.learner.default_dropout['language'] = self.dps
|
||||
def create_cls_learner(self, data_clas, dps=None, **kwargs):
|
||||
fastai.text.learner.default_dropout['language'] = dps or self.dps
|
||||
trn_args=dict(drop_mult=self.drop_mult, bptt=self.bptt, clip=self.clip,)
|
||||
trn_args.update(kwargs)
|
||||
classifier_learner = bilm_text_classifier_learner if self.bidir else text_classifier_learner
|
||||
learn = classifier_learner(data_clas, bptt=self.bptt, pad_token=PAD_TOKEN_ID,
|
||||
learn = classifier_learner(data_clas, pad_token=PAD_TOKEN_ID,
|
||||
path=self.model_dir.parent, model_dir=self.model_dir.name,
|
||||
qrnn=self.qrnn, emb_sz=self.emb_sz, nh=self.nh, nl=self.nl, drop_mult=self.drop_mult)
|
||||
learn.true_wd = False
|
||||
print("true_wd: ", learn.true_wd)
|
||||
qrnn=self.qrnn, emb_sz=self.emb_sz, nh=self.nh, nl=self.nl, **trn_args)
|
||||
return learn
|
||||
|
||||
def load_cls_data(self, bs):
|
||||
def load_cls_data(self, bs, **kwargs):
|
||||
if self.dataset_dir.name == 'imdb':
|
||||
return self.load_cls_data_imdb(bs)
|
||||
return self.load_cls_data_imdb(bs, **kwargs)
|
||||
else:
|
||||
assert self.tokenizer is Tokenizers.MOSES, "XNLI does not support other tokenizers than Moses"
|
||||
return self.load_cls_data_old_for_xnli(bs)
|
||||
return self.load_cls_data_old_for_xnli(bs, **kwargs)
|
||||
|
||||
def load_cls_data_imdb(self, bs):
|
||||
def load_cls_data_imdb(self, bs, force=False, use_test_for_validation=False):
|
||||
trn_df = pd.read_csv(self.dataset_path / 'train.csv', header=None)
|
||||
tst_df = pd.read_csv(self.dataset_path / 'test.csv', header=None)
|
||||
unsp_df = pd.read_csv(self.dataset_path / 'unsup.csv', header=None)
|
||||
@@ -106,7 +115,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
lm_trn_df = lm_trn_df[val_len:]
|
||||
lm_val_df = lm_trn_df[:val_len]
|
||||
|
||||
if self.use_test_for_validation:
|
||||
if use_test_for_validation:
|
||||
val_len = max(int(len(tst_df) * 0.1), 2)
|
||||
tst_len = len(tst_df) - val_len
|
||||
val_df = trn_df[:tst_len]
|
||||
@@ -115,6 +124,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
trn_len = len(trn_df) - val_len
|
||||
trn_df, val_df = trn_df[:trn_len], trn_df[trn_len:]
|
||||
|
||||
|
||||
if self.tokenizer is Tokenizers.SUBWORD:
|
||||
#TODO Fix me to make sure it trains correct dictionary
|
||||
args = get_sentencepiece(self.dataset_path, self.dataset_path / 'train.csv', self.name, vocab_size=self.max_vocab)
|
||||
@@ -129,7 +139,8 @@ class CLSHyperParams(LMHyperParams):
|
||||
f"self.tokenizer has wrong value {self.tokenizer}, Allowed values are taken from {Tokenizers}")
|
||||
|
||||
try:
|
||||
data_lm = TextLMDataBunch.load(self.cache_dir, 'lm', lm_type=self.lm_type)
|
||||
if force: raise FileNotFoundError("Forcing reloading of caches")
|
||||
data_lm = TextLMDataBunch.load(self.cache_dir, 'lm', lm_type=self.lm_type, bs=bs)
|
||||
print(f"Tokenized data loaded, lm.trn {len(data_lm.train_ds)}, lm.val {len(data_lm.valid_ds)}")
|
||||
except FileNotFoundError:
|
||||
print(f"Running tokenization...")
|
||||
@@ -137,18 +148,17 @@ class CLSHyperParams(LMHyperParams):
|
||||
max_vocab=self.max_vocab, bs=bs, lm_type=self.lm_type, **args)
|
||||
print(f"Saving tokenized: cls.trn {len(data_lm.train_ds)}, cls.val {len(data_lm.valid_ds)}")
|
||||
data_lm.save('lm')
|
||||
print(f" cls.trn {len(data_lm.train_ds)}, cls.val {len(data_lm.valid_ds)}")
|
||||
|
||||
args['vocab'] = data_lm.vocab # make sure we use the same vocab for classifcation
|
||||
try:
|
||||
data_cls = TextClasDataBunch.load(self.cache_dir, '.')
|
||||
if force: raise FileNotFoundError("Forcing reloading of caches")
|
||||
data_cls = TextClasDataBunch.load(self.cache_dir, '.', bs=bs)
|
||||
print(f"Tokenized data loaded, cls.trn {len(data_cls.train_ds)}, cls.val {len(data_cls.valid_ds)}")
|
||||
except FileNotFoundError:
|
||||
args['vocab'] = data_lm.vocab # make sure we use the same vocab for classifcation
|
||||
print(f"Running tokenization...")
|
||||
data_cls = TextClasDataBunch.from_df(path=self.cache_dir, train_df=trn_df,
|
||||
valid_df=val_df, test_df=tst_df, max_vocab=self.max_vocab,
|
||||
**args)
|
||||
print(f" cls.trn {len(data_cls.train_ds)}, cls.val {len(data_cls.valid_ds)}")
|
||||
data_cls = TextClasDataBunch.from_df(path=self.cache_dir, train_df=trn_df, valid_df=val_df,
|
||||
test_df=tst_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('.')
|
||||
print('Size of vocabulary:', len(data_lm.vocab.itos))
|
||||
print('First 20 words in vocab:', data_lm.vocab.itos[:20])
|
||||
|
||||
Reference in New Issue
Block a user