From 38ca35c044effc285d5cf08de1deaa6750b7af96 Mon Sep 17 00:00:00 2001 From: Piotr Czapla Date: Wed, 14 Nov 2018 11:41:26 +0100 Subject: [PATCH] Fix pretrain_lm to work with yesterdays changes to fastai text API --- ulmfit/pretrain_lm.py | 11 ++++------- 1 file changed, 4 insertions(+), 7 deletions(-) diff --git a/ulmfit/pretrain_lm.py b/ulmfit/pretrain_lm.py index 1836708..ffa4503 100644 --- a/ulmfit/pretrain_lm.py +++ b/ulmfit/pretrain_lm.py @@ -8,9 +8,8 @@ import fastai import fire import numpy as np -from fastai import DataBunch, partial, optim, fit_one_cycle -from fastai.text import LanguageModelLoader, get_language_model, RNNLearner, TextLMDataBunch, NumericalizedDataset, \ - Vocab, language_model_learner +from fastai import * +from fastai.text import * import torch from fastai_contrib.utils import read_file, read_whitespace_file,\ DataStump, validate, PAD, UNK @@ -79,10 +78,8 @@ def pretrain_lm(dir_path, cuda_id=0, qrnn=True, clean=True, max_vocab=60000, val_ids = np.array([([stoi.get(w, stoi[UNK]) for w in s]) for s in val_tok]) # data_lm = TextLMDataBunch.from_ids(dir_path, trn_ids, [], val_ids, [], len(itos)) - data_lm = TextLMDataBunch.create( - train_ds=NumericalizedDataset(vocab, trn_ids, labels=np.zeros(len(trn_ids), dtype=np.int)), - valid_ds=NumericalizedDataset(vocab, val_ids, labels=np.zeros(len(val_ids), dtype=np.int)), - bs=bs, bptt=bptt) + data_lm = TextLMDataBunch.from_ids(path=dir_path, vocab=vocab, train_ids=trn_ids, + valid_ids=val_ids, bs=bs, bptt=bptt) else: # apply fastai preprocessing and tokenization read_file(trn_path, 'train')