mirror of
https://github.com/wassname/multifit.git
synced 2026-09-09 11:27:26 +08:00
Make models use bptt parameter
This commit is contained in:
+6
-4
@@ -50,7 +50,8 @@ class ULMFiT:
|
||||
return FireView(train=params.train_cls, validate_cls=params.validate_cls)
|
||||
|
||||
|
||||
def eval_noise_resistance(self, lang="de", size=1, prefix_name="", model="sp15k/qrnn_nl4.m"):
|
||||
def eval_noise_resistance(self, lang="de", size=1, prefix_name="", model="sp15k/qrnn_nl4.m",
|
||||
num_cls_epochs=8, bs=18, lr_sched="1cycle", label_smoothing_eps=0.0):
|
||||
def first_or_default(l, default=None):
|
||||
l = list(l)
|
||||
if l:
|
||||
@@ -63,9 +64,10 @@ class ULMFiT:
|
||||
name=f"nl4_{prefix_name}{noise}",
|
||||
noise=noise/100,
|
||||
dataset_template='${lang}-'+str(size),
|
||||
num_cls_epochs=8,
|
||||
bs=18,
|
||||
lr_sched="1cycle")
|
||||
num_cls_epochs=num_cls_epochs,
|
||||
bs=bs,
|
||||
lr_sched=lr_sched,
|
||||
label_smoothing_eps=label_smoothing_eps)
|
||||
val = first_or_default(d.values(), default=-1)
|
||||
results.append((noise/100, val))
|
||||
df = pd.DataFrame(results, columns=["noise", "accuracy"])
|
||||
|
||||
@@ -175,7 +175,7 @@ class LMHyperParams:
|
||||
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, label_smoothing_eps=label_smoothing_eps)
|
||||
|
||||
print("Bptt", data_lm.bptt)
|
||||
learn.true_wd = true_wd
|
||||
if num_epochs > 0:
|
||||
if self.pretrained_fnames or self.pretrained_model:
|
||||
@@ -255,6 +255,7 @@ class LMHyperParams:
|
||||
classes=None,
|
||||
bs=bs,
|
||||
text_cols='texts',
|
||||
bptt=self.bptt,
|
||||
**args)
|
||||
|
||||
itos, stoi, trn_path = data_lm.vocab.itos, data_lm.vocab.stoi, data_lm.path
|
||||
|
||||
@@ -112,7 +112,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
trn_args=dict(bptt=self.bptt, clip=self.clip)
|
||||
trn_args.update(kwargs)
|
||||
learn = text_classifier_learner(data_clas, AWD_LSTM, config=config,
|
||||
pretrained=False, path=self.model_dir.parent, model_dir=self.model_dir.name, **trn_args)
|
||||
pretrained=False, path=self.model_dir.parent, model_dir=self.model_dir.name, bptt=self.bptt, **trn_args)
|
||||
|
||||
if self.pretrained_model is not None:
|
||||
print("Loading pretrained model")
|
||||
@@ -151,7 +151,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
def merge_cols(self, df):
|
||||
if len(df.columns) <= 2:
|
||||
return df
|
||||
ndf = df[[0,1]].copy()
|
||||
ndf = df[[0,1]].copy().fillna(" ")
|
||||
for i in range(2, len(df.columns)):
|
||||
ndf[1] += ("\n" + FLD + "\n") + df[i].fillna(" ")
|
||||
|
||||
@@ -213,7 +213,8 @@ class CLSHyperParams(LMHyperParams):
|
||||
cls_name = f'{cls_name}noise{noise}tv'
|
||||
|
||||
args = self.tokenizer_to_fastai_args(sp_data_func=lambda: trn_df[1], use_moses=use_moses)
|
||||
data_lm = self.lm_databunch('lm', train_df=lm_trn_df, valid_df=lm_val_df, bs=bs, force=force, **args)
|
||||
lm_suffix = self.bptt if self.bptt != 70 else ""
|
||||
data_lm = self.lm_databunch(f'lm{lm_suffix}', train_df=lm_trn_df, valid_df=lm_val_df, bs=bs, force=force, bptt=self.bptt, **args)
|
||||
args['vocab'] = data_lm.vocab
|
||||
data_cls = self.cls_databunch(cls_name, train_df=trn_df, valid_df=val_df, bs=bs, force=force, **args)
|
||||
data_tst = self.cls_databunch('tst', train_df=val_df, valid_df=tst_df, bs=bs, force=force, **args) # Hack to load test dataset with labels
|
||||
|
||||
Reference in New Issue
Block a user