Add code to test & train mldoc classifier

This commit is contained in:
Piotr Czapla
2019-02-10 09:52:58 +01:00
parent 0085c18ae0
commit 26736d95de
2 changed files with 85 additions and 77 deletions
+27
View File
@@ -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
View File
@@ -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)