mirror of
https://github.com/wassname/multifit.git
synced 2026-09-09 11:27:26 +08:00
Add Kappa and Mathew score calcualtion + ls command to main
This commit is contained in:
+42
-34
@@ -30,7 +30,6 @@ def get_dataset_path(p, dataset_template):
|
||||
ds = [x for x in p.parents if x.name == "models"][0].parent
|
||||
lang = get_lang_from_dataset_path(ds)
|
||||
pattern = Template(dataset_template).substitute(lang=lang, ds_name=ds.name)
|
||||
print(f"Searching for {pattern}, {ds.parent}")
|
||||
for ds_path in ds.parent.glob(pattern):
|
||||
yield lang, ds_path
|
||||
|
||||
@@ -140,6 +139,17 @@ class ULMFiT:
|
||||
lr_sched=lr_sched,
|
||||
**kwargs)
|
||||
|
||||
def ls(self, glob, dataset_template='${ds_name}'):
|
||||
data_dir = Path("data").absolute()
|
||||
glob = str(glob)
|
||||
if "data" not in glob and not glob.startswith("/"):
|
||||
glob = "data/" + glob
|
||||
results = []
|
||||
for base_model in sorted(data_dir.parent.glob(glob)):
|
||||
for lang, dataset_path in sorted(get_dataset_path(base_model, dataset_template)):
|
||||
results.append((base_model, lang, dataset_path))
|
||||
return results
|
||||
|
||||
def eval(self, glob="data/mldoc/*-1/models/sp30k/lstm_nl4.m", dataset_template='${ds_name}', name=None,
|
||||
num_lm_epochs=0, train=True, to_csv=None, return_df=False, label_smoothing_eps=0.0,
|
||||
lmseed=None, ftseed=None, clsweightseed=None, clstrainseed=None, save_name="cls_best",
|
||||
@@ -159,39 +169,37 @@ class ULMFiT:
|
||||
glob=str(glob)
|
||||
if "data" not in glob and not glob.startswith("/"):
|
||||
glob = "data/"+glob
|
||||
for base_model in sorted(data_dir.parent.glob(glob)):
|
||||
print("Processing", base_model)
|
||||
for lang, dataset_path in sorted(get_dataset_path(base_model, dataset_template)):
|
||||
try:
|
||||
_name = name
|
||||
if name is None:
|
||||
_name = folder_name_to_model_name(base_model.name)
|
||||
params = CLSHyperParams.from_lm(dataset_path, base_model, lang=lang, name=_name, **model_args)
|
||||
last_model_dir = params.model_dir.relative_to(data_dir.parent)
|
||||
if (params.model_dir/"cls_best.pth").exists():
|
||||
print("Evaluating previously trained model")
|
||||
d_tst = params.validate_cls(save_name=save_name, label_smoothing_eps=label_smoothing_eps, use_cache=True, mode="test")
|
||||
d_val = params.validate_cls(save_name=save_name, label_smoothing_eps=label_smoothing_eps, use_cache=True, mode="valid")
|
||||
d={}
|
||||
d.update(d_val)
|
||||
d.update(d_tst)
|
||||
elif train:
|
||||
print("Training")
|
||||
d = params.train_cls(num_lm_epochs=num_lm_epochs, label_smoothing_eps=label_smoothing_eps, **trn_params)
|
||||
else:
|
||||
print("Skipping", (params.model_dir/"cls_best.pth"))
|
||||
d = None
|
||||
if d is not None:
|
||||
d['model_dir_parent'] = params.model_dir.relative_to(data_dir.parent).parent
|
||||
d['model_name'] = params.model_name
|
||||
np.save(params.model_dir / "results.npy", d)
|
||||
results.append(d)
|
||||
del params
|
||||
except Exception as e:
|
||||
print("Error", e)
|
||||
if not skip_on_error:
|
||||
raise e
|
||||
gc.collect()
|
||||
for base_model, lang, dataset_path in self.ls(glob, dataset_template):
|
||||
try:
|
||||
_name = name
|
||||
if name is None:
|
||||
_name = folder_name_to_model_name(base_model.name)
|
||||
params = CLSHyperParams.from_lm(dataset_path, base_model, lang=lang, name=_name, **model_args)
|
||||
last_model_dir = params.model_dir.relative_to(data_dir.parent)
|
||||
if (params.model_dir/"cls_best.pth").exists():
|
||||
print("Evaluating previously trained model")
|
||||
d_tst = params.validate_cls(save_name=save_name, label_smoothing_eps=label_smoothing_eps, use_cache=True, mode="test")
|
||||
d_val = params.validate_cls(save_name=save_name, label_smoothing_eps=label_smoothing_eps, use_cache=True, mode="valid")
|
||||
d={}
|
||||
d.update(d_val)
|
||||
d.update(d_tst)
|
||||
elif train:
|
||||
print("Training")
|
||||
d = params.train_cls(num_lm_epochs=num_lm_epochs, label_smoothing_eps=label_smoothing_eps, **trn_params)
|
||||
else:
|
||||
print("Skipping", (params.model_dir/"cls_best.pth"))
|
||||
d = None
|
||||
if d is not None:
|
||||
d['model_dir_parent'] = params.model_dir.relative_to(data_dir.parent).parent
|
||||
d['model_name'] = params.model_name
|
||||
np.save(params.model_dir / "results.npy", d)
|
||||
results.append(d)
|
||||
del params
|
||||
except Exception as e:
|
||||
print("Error", e)
|
||||
if not skip_on_error:
|
||||
raise e
|
||||
gc.collect()
|
||||
df = pd.DataFrame.from_records(results)
|
||||
print(df)
|
||||
if to_csv is not None:
|
||||
|
||||
+29
-10
@@ -66,6 +66,19 @@ class CLSHyperParams(LMHyperParams):
|
||||
learn.unfreeze()
|
||||
learn.fit_one_cycle(num_cls_epochs, slice(1e-2 / (2.6 ** 4), 2e-2), moms=(0.8, 0.7))
|
||||
|
||||
def lr_schedule_reverse_2cycle(self, learn, num_cls_epochs):
|
||||
print("Reverse 2cycle ")
|
||||
learn.unfreeze()
|
||||
for g in learn.layer_groups[-1:]:
|
||||
for l in g:
|
||||
if not learn.train_bn or not isinstance(l, bn_types): requires_grad(l, False)
|
||||
learn.create_opt(defaults.lr)
|
||||
print("training LM")
|
||||
learn.fit_one_cycle(num_cls_epochs, slice(1e-2 / (2.6 ** 4), 2e-2), moms=(0.8, 0.7))
|
||||
learn.unfreeze()
|
||||
print("training ALL")
|
||||
learn.fit_one_cycle(num_cls_epochs, slice(1e-3 / (2.6 ** 4), 2e-3), moms=(0.8, 0.7))
|
||||
|
||||
def lr_schedule_false_wd(self, learn, num_cls_epochs):
|
||||
learn.true_wd = False
|
||||
print("Starting classifier training")
|
||||
@@ -83,8 +96,9 @@ class CLSHyperParams(LMHyperParams):
|
||||
f1_score = FBeta(beta=1.0)
|
||||
precision = Precision()
|
||||
recall = Recall()
|
||||
metrics = [f1_score, precision, recall]
|
||||
|
||||
kappa_lin = KappaScore('linear')
|
||||
matthews_correff = MatthewsCorreff()
|
||||
metrics = [f1_score, precision, recall, kappa_lin, matthews_correff]
|
||||
# TODO: fix this in fast.ai
|
||||
if init:
|
||||
for metric in metrics: metric.on_train_begin()
|
||||
@@ -96,17 +110,19 @@ class CLSHyperParams(LMHyperParams):
|
||||
print(f"Loss: {results[0]}")
|
||||
print(f"Precision: {results[2].item()}")
|
||||
print(f"Recall: {results[3].item()}")
|
||||
print(f"Accuracy: {results[4].item()}")
|
||||
print(f"Accuracy: {results[6].item()}")
|
||||
d = {f"{mode} F1 score bin": results[1].item(),
|
||||
f"{mode} Loss": results[0],
|
||||
f"{mode} Precision": results[2].item(),
|
||||
f"{mode} Recall": results[3].item(),
|
||||
f"{mode} Accuracy": results[4].item()}
|
||||
f"{mode} Kappa Linear": results[4].item(),
|
||||
f"{mode} Matthews Correff": results[5].item(),
|
||||
f"{mode} Accuracy": results[6].item()}
|
||||
return {k:float(str(v)) for k,v in d.items()} # float(str(x)) to avoid float32 -> float64 conversion isssues
|
||||
|
||||
def train_cls(self, num_lm_epochs, unfreeze=True, num_cls_frozen_epochs=1, bs=40, drop_mul_lm=0.3, drop_mul_cls=0.5,
|
||||
use_test_for_validation=False, num_cls_epochs=2, limit=None, noise=0.0, cls_max_len=20*70, lr_sched='layered',
|
||||
label_smoothing_eps=0.0, random_init=False, dump_preds=None):
|
||||
label_smoothing_eps=0.0, random_init=False, dump_preds=None, early_stopping=True):
|
||||
assert use_test_for_validation == False, "use_test_for_validation=True is not supported"
|
||||
self.model_dir.mkdir(exist_ok=True, parents=True)
|
||||
|
||||
@@ -127,7 +143,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
learn = self.create_cls_learner(data_clas, drop_mult=drop_mul_cls, max_len=cls_max_len,
|
||||
label_smoothing_eps=label_smoothing_eps, random_init=random_init,
|
||||
metrics=self.get_metrics(),
|
||||
loss_func=loss_func)
|
||||
loss_func=loss_func, early_stopping=early_stopping)
|
||||
|
||||
if not random_init:
|
||||
try:
|
||||
@@ -189,7 +205,7 @@ class CLSHyperParams(LMHyperParams):
|
||||
|
||||
return labeled_results
|
||||
|
||||
def create_cls_learner(self, data_clas, dps=None, label_smoothing_eps=0.0, random_init=False, **kwargs):
|
||||
def create_cls_learner(self, data_clas, dps=None, label_smoothing_eps=0.0, random_init=False, early_stopping=True, **kwargs):
|
||||
assert self.bidir == False, "bidirectional model is not yet supported"
|
||||
config = dict(emb_sz=self.emb_sz, n_hid=self.nh, n_layers=self.nl, pad_token=PAD_TOKEN_ID, qrnn=self.qrnn)
|
||||
config.update(dps or self.dps)
|
||||
@@ -205,9 +221,12 @@ class CLSHyperParams(LMHyperParams):
|
||||
learn.load_pretrained(*fnames, strict=False)
|
||||
learn.freeze()
|
||||
|
||||
learn.callback_fns += [partial(CSVLogger, filename=f"{learn.model_dir}/cls-history"),
|
||||
partial(SaveModelCallback, every='improvement', name='cls_best_tmp', monitor="f_beta")
|
||||
]
|
||||
learn.callback_fns += [partial(CSVLogger, filename=f"{learn.model_dir}/cls-history")]
|
||||
if early_stopping:
|
||||
learn.callback_fns += [partial(SaveModelCallback, every='improvement',
|
||||
name='cls_best_tmp',
|
||||
monitor="f_beta")]
|
||||
|
||||
if label_smoothing_eps > 0.0:
|
||||
learn.loss_func = FlattenedLoss(LabelSmoothingCrossEntropy, eps=label_smoothing_eps)
|
||||
return learn
|
||||
|
||||
Reference in New Issue
Block a user