mirror of
https://github.com/wassname/multifit.git
synced 2026-09-10 12:12:50 +08:00
Add code to test & train mldoc classifier
This commit is contained in:
@@ -0,0 +1,27 @@
|
||||
import fire
|
||||
import urllib.request
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
langs = ['english', 'spanish', 'german', 'chinese', 'french', 'russian', 'japanese', 'italian']
|
||||
lang_codes = ['en', 'es', 'de', 'zh', 'fr', 'ru', 'ja', 'it']
|
||||
|
||||
def fetch_mldoc(url_prefix, mldoc_path="data/mldoc"):
|
||||
""" Fetch mldoc from server using basic auth
|
||||
url_prefix should point to mldoc stored as follow
|
||||
"https://user:passwd@server/path/[english|spanish...].[dev|test|train.[1000|2000|5000|10000]].csv"
|
||||
"""
|
||||
def fetch(url, mldoc):
|
||||
mldoc.parent.mkdir(parents=True, exist_ok=True)
|
||||
print("fetching", url, mldoc)
|
||||
urllib.request.urlretrieve(url, mldoc)
|
||||
for lang,code in zip(langs, lang_codes):
|
||||
for size in [1000, 2000, 5000, 10000]:
|
||||
dir = Path(mldoc_path)/f"{code}-{size // 1000}"
|
||||
fetch(f"{url_prefix}/{lang}.dev.csv", dir / f"{code}.dev.csv")
|
||||
fetch(f"{url_prefix}/{lang}.test.csv", dir / f"{code}.test.csv")
|
||||
fetch(f"{url_prefix}/{lang}.train.{size}.csv", dir / f"{code}.train.csv")
|
||||
fetch(f"{url_prefix}/{lang}.train.10000.csv", dir / f"{code}.unsup.csv")
|
||||
|
||||
if __name__ == "__main__":
|
||||
fire.Fire(fetch_mldoc)
|
||||
+58
-77
@@ -37,10 +37,11 @@ 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=40, true_wd=True, drop_mul_lm=0.3, drop_mul_cls=0.5,
|
||||
use_test_for_validation=False):
|
||||
data_clas, data_lm = self.load_cls_data(bs, use_test_for_validation=use_test_for_validation)
|
||||
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)
|
||||
|
||||
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)
|
||||
@@ -74,15 +75,13 @@ class CLSHyperParams(LMHyperParams):
|
||||
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)
|
||||
self.validate_cls('cls_last', bs=bs)
|
||||
self.validate_cls('cls_best', bs=bs)
|
||||
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, save_name='cls_last', bs=40):
|
||||
data_clas, data_lm = self.load_cls_data(bs, use_test_for_validation=True)
|
||||
learn = self.create_cls_learner(data_clas, drop_mult=0.1)
|
||||
def validate_cls(self, learn, data_tst, save_name='cls_last', bs=40):
|
||||
learn.load(save_name)
|
||||
print(f"Loss and accuracy using ({save_name}):", learn.validate())
|
||||
print(f"Loss and accuracy using ({save_name}):", learn.validate(data_tst.test_dl))
|
||||
|
||||
def create_cls_learner(self, data_clas, dps=None, **kwargs):
|
||||
fastai.text.learner.default_dropout['language'] = dps or self.dps
|
||||
@@ -100,33 +99,48 @@ class CLSHyperParams(LMHyperParams):
|
||||
return learn
|
||||
|
||||
def load_cls_data(self, bs, **kwargs):
|
||||
add_trn_to_lm = True
|
||||
lang = self.lang
|
||||
if 'xnli' in str(self.dataset_dir):
|
||||
NotImplementedError("Support for Xnli is not implemented yet")
|
||||
if 'imdb' in self.dataset_dir.name:
|
||||
return self.load_cls_data_imdb(bs, **kwargs)
|
||||
add_trn_to_lm = True
|
||||
if 'mldoc' in str(self.dataset_dir):
|
||||
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)
|
||||
return self.databunches(bs, **data)
|
||||
|
||||
def load_data(self, lang='', **kwargs):
|
||||
prefix = '' if lang == '' else lang+'.'
|
||||
trn_df = pd.read_csv(self.dataset_path / f'{prefix}train.csv', header=None)
|
||||
tst_df = pd.read_csv(self.dataset_path / f'{prefix}test.csv', header=None)
|
||||
val_fn = self.dataset_path / f'{prefix}dev.csv'
|
||||
if val_fn.exists():
|
||||
print("Loading validation", val_fn)
|
||||
val_df = pd.read_csv(val_fn, header=None)
|
||||
else:
|
||||
assert self.tokenizer is Tokenizers.MOSES, "XNLI does not support other tokenizers than Moses"
|
||||
return self.load_cls_data_old_for_xnli(bs, **kwargs)
|
||||
val_df = None
|
||||
|
||||
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)
|
||||
unsup_df = pd.read_csv(self.dataset_path / f'{prefix}unsup.csv', header=None)
|
||||
|
||||
lm_trn_df = pd.concat([unsp_df, trn_df, tst_df])
|
||||
if val_df is None:
|
||||
print("Validation set not found using 10% of trn")
|
||||
val_len = max(int(len(trn_df) * 0.1), 2)
|
||||
trn_len = len(trn_df) - val_len
|
||||
trn_df, val_df = trn_df[:trn_len], trn_df[trn_len:]
|
||||
|
||||
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):
|
||||
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]
|
||||
|
||||
if use_test_for_validation:
|
||||
val_df = tst_df
|
||||
cls_cache = 'notst'
|
||||
else:
|
||||
val_len = max(int(len(trn_df) * 0.1), 2)
|
||||
trn_len = len(trn_df) - val_len
|
||||
trn_df, val_df = trn_df[:trn_len], trn_df[trn_len:]
|
||||
cls_cache = '.'
|
||||
|
||||
args = self.tokenzier_to_fastai_args(trn_data_loading_func=lambda: trn_df[1], add_moses=True)
|
||||
|
||||
try:
|
||||
if force: raise FileNotFoundError("Forcing reloading of caches")
|
||||
data_lm = TextLMDataBunch.load(self.cache_dir, 'lm', lm_type=self.lm_type, bs=bs)
|
||||
@@ -140,64 +154,31 @@ class CLSHyperParams(LMHyperParams):
|
||||
|
||||
try:
|
||||
if force: raise FileNotFoundError("Forcing reloading of caches")
|
||||
data_cls = TextClasDataBunch.load(self.cache_dir, cls_cache, bs=bs)
|
||||
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...")
|
||||
args['vocab'] = data_lm.vocab # make sure we use the same vocab for classifcation
|
||||
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)
|
||||
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(cls_cache)
|
||||
data_cls.save('.')
|
||||
|
||||
# Hack to load test dataset with labels
|
||||
try:
|
||||
if force: raise FileNotFoundError("Forcing reloading of caches")
|
||||
data_tst = TextClasDataBunch.load(self.cache_dir, 'tst', bs=bs)
|
||||
except FileNotFoundError:
|
||||
args['vocab'] = data_lm.vocab # make sure we use the same vocab for classifcation
|
||||
data_tst = TextClasDataBunch.from_df(path=self.cache_dir, train_df=val_df, valid_df=tst_df,
|
||||
max_vocab=self.max_vocab, bs=bs, **args)
|
||||
data_tst.save('tst')
|
||||
|
||||
#$data_cls.test_dl = data_tst.valid_dl
|
||||
#data_cls.test_ds = data_tst.valid_ds # AttributeError: can't set attribute
|
||||
print('Size of vocabulary:', len(data_lm.vocab.itos))
|
||||
print('First 20 words in vocab:', data_lm.vocab.itos[:20])
|
||||
return data_cls, data_lm
|
||||
|
||||
|
||||
def load_cls_data_old_for_xnli(self, bs):
|
||||
tmp_dir = self.cache_dir
|
||||
tmp_dir.mkdir(exist_ok=True)
|
||||
vocab_file = tmp_dir / f'vocab_{self.lang}.pkl'
|
||||
if not (tmp_dir / f'{TRN}_{self.lang}_ids.npy').exists():
|
||||
print('Reading the data...')
|
||||
toks, lbls = read_clas_data(self.dataset_dir, self.dataset_dir.name, self.lang)
|
||||
# create the vocabulary
|
||||
counter = Counter(word for example in toks[TRN] + toks[TST] + toks[VAL] for word in example)
|
||||
itos = [word for word, count in counter.most_common(n=self.max_vocab)]
|
||||
itos.insert(0, PAD)
|
||||
itos.insert(0, UNK)
|
||||
vocab = Vocab(itos)
|
||||
stoi = vocab.stoi
|
||||
with open(vocab_file, 'wb') as f:
|
||||
pickle.dump(vocab, f)
|
||||
ids = {}
|
||||
for split in [TRN, VAL, TST]:
|
||||
ids[split] = np.array([([stoi.get(w, stoi[UNK]) for w in s])
|
||||
for s in toks[split]])
|
||||
np.save(tmp_dir / f'{split}_{self.lang}_ids.npy', ids[split])
|
||||
np.save(tmp_dir / f'{split}_{self.lang}_lbl.npy', lbls[split])
|
||||
else:
|
||||
print('Loading the pickled data...')
|
||||
ids, lbls = {}, {}
|
||||
for split in [TRN, VAL, TST]:
|
||||
ids[split] = np.load(tmp_dir / f'{split}_{self.lang}_ids.npy')
|
||||
lbls[split] = np.load(tmp_dir / f'{split}_{self.lang}_lbl.npy')
|
||||
with open(vocab_file, 'rb') as f:
|
||||
vocab = pickle.load(f)
|
||||
print(f'Train size: {len(ids[TRN])}. Valid size: {len(ids[VAL])}. '
|
||||
f'Test size: {len(ids[TST])}.')
|
||||
for split in [TRN, VAL, TST]:
|
||||
ids[split] = np.array([np.array(e, dtype=np.int) for e in ids[split]])
|
||||
lbls[split] = np.array([np.array(e, dtype=np.int) for e in lbls[split]])
|
||||
data_lm = TextLMDataBunch.from_ids(path=tmp_dir, vocab=vocab, train_ids=np.concatenate([ids[TRN], ids[TST]]),
|
||||
valid_ids=ids[VAL], bs=bs, bptt=self.bptt, lm_type=self.lm_type)
|
||||
# TODO TextClasDataBunch allows tst_ids as input, but not tst_lbls?
|
||||
data_clas = TextClasDataBunch.from_ids(
|
||||
path=tmp_dir, vocab=vocab, train_ids=ids[TRN], valid_ids=ids[VAL],
|
||||
train_lbls=lbls[TRN], valid_lbls=lbls[VAL], bs=bs, classes={l: l for l in lbls[TRN]})
|
||||
|
||||
print(f"Sizes of train_ds {len(data_clas.train_ds)}, valid_ds {len(data_clas.valid_ds)}")
|
||||
return data_clas, data_lm
|
||||
return data_cls, data_lm, data_tst
|
||||
|
||||
if __name__ == '__main__':
|
||||
fire.Fire(CLSHyperParams)
|
||||
|
||||
Reference in New Issue
Block a user