mirror of
https://github.com/wassname/multifit.git
synced 2026-09-10 12:12:50 +08:00
WIP Working Backward LM using our new LangaugeModelLoader
This commit is contained in:
@@ -24,7 +24,7 @@ class LanguageModelLoader(): # copy of the original LanguageModelLoader
|
||||
if getattr(self.dataset, 'item', None) is not None:
|
||||
yield LongTensor(getattr(self.dataset, 'item')).unsqueeze(1),LongTensor([0])
|
||||
idx = np.random.permutation(len(self.dataset)) if self.shuffle else range(len(self.dataset))
|
||||
data = self.batchify(np.concatenate([self.dataset.x[i] for i in idx]))
|
||||
data = self.batchify(np.concatenate([self.dataset.x.items[i] for i in idx]))
|
||||
|
||||
pos, itr = 0,0
|
||||
while pos < self.n-1 and itr<len(self):
|
||||
@@ -51,7 +51,7 @@ class LanguageModelLoader(): # copy of the original LanguageModelLoader
|
||||
|
||||
def get_batch(self, data:LongTensor, i:int, seq_len:int) -> Tuple[LongTensor, LongTensor]:
|
||||
"Create a batch at `i` of a given `seq_len`."
|
||||
seq_len = min(seq_len, len(self.data) - 1 - i)
|
||||
seq_len = min(seq_len, len(data) - 1 - i)
|
||||
x = data[i:i+seq_len]
|
||||
y = data[i+1:i+1+seq_len].contiguous() # x & y has 2 elements on the last dimension
|
||||
y = y.view(-1, 2) if self.lm_type == LanguageModelType.BiLM else y.view(-1)
|
||||
|
||||
@@ -10,7 +10,7 @@ def bilm_learner(data:DataBunch, bptt:int=70, emb_sz:int=400, nh:int=1150, nl:in
|
||||
pretrained_fnames:OptStrTuple=None, **kwargs) -> 'LanguageLearner':
|
||||
"Create a `Learner` with a language model."
|
||||
dps = default_dropout['language'] * drop_mult
|
||||
vocab_size = data.train_ds.vocab_size
|
||||
vocab_size = len(data.vocab.itos)
|
||||
model = get_bilm(vocab_size, emb_sz, nh, nl, pad_token, input_p=dps[0], output_p=dps[1],
|
||||
weight_p=dps[2], embed_p=dps[3], hidden_p=dps[4], tie_weights=tie_weights, bias=bias, qrnn=qrnn)
|
||||
learn = LanguageLearner(data, model, bptt, split_func=bilm_split, **kwargs)
|
||||
|
||||
@@ -44,7 +44,10 @@ class BiLMCore(nn.Module):
|
||||
self.hidden_dps = nn.ModuleList([RNNDropout(hidden_p) for l in range(n_layers)])
|
||||
|
||||
def forward(self, input:LongTensor)->Tuple[Tensor,Tensor]:
|
||||
sl,bs = input.size()
|
||||
sl,bs,tracks = input.size()
|
||||
assert tracks == 2, "It should have two tracks for forward and backward pass"
|
||||
|
||||
input = input[...,0] # Select forward pass only
|
||||
if bs!=self.bs:
|
||||
self.bs=bs
|
||||
self.reset()
|
||||
@@ -59,6 +62,8 @@ class BiLMCore(nn.Module):
|
||||
if l != self.n_layers - 1: raw_output = hid_dp(raw_output)
|
||||
outputs.append(raw_output)
|
||||
self.hidden = to_detach(new_hidden)
|
||||
|
||||
#bi_raw_outputs = torch.stack((outputs, outputs), dim=2)
|
||||
return raw_outputs, outputs
|
||||
|
||||
def _one_hidden(self, l:int)->Tensor:
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
import pytest
|
||||
from fastai import *
|
||||
from fastai.text import *
|
||||
|
||||
pytestmark = pytest.mark.integration
|
||||
|
||||
import fastai_contrib.data as contrib_data
|
||||
|
||||
from fastai_contrib.learner import bilm_learner
|
||||
|
||||
def read_file(fname):
|
||||
texts = []
|
||||
with open(fname, 'r') as f:
|
||||
texts = f.readlines()
|
||||
labels = [0] * len(texts)
|
||||
df = pd.DataFrame({'labels':labels, 'texts':texts}, columns = ['labels', 'texts'])
|
||||
return df
|
||||
|
||||
def prep_human_numbers():
|
||||
path = untar_data(URLs.HUMAN_NUMBERS)
|
||||
df_trn = read_file(path/'train.txt')
|
||||
df_val = read_file(path/'valid.txt')
|
||||
return path, df_trn, df_val
|
||||
|
||||
def manual_seed(seed=42):
|
||||
torch.manual_seed(seed)
|
||||
np.random.seed(seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def learn():
|
||||
path, df_trn, df_val = prep_human_numbers()
|
||||
data = TextLMDataBunch.from_df(path, df_trn, df_val, tokenizer=Tokenizer(BaseTokenizer))
|
||||
learn = language_model_learner(data, emb_sz=100, nl=1, drop_mult=0.1)
|
||||
learn.fit_one_cycle(4, 5e-3)
|
||||
return learn
|
||||
|
||||
###################### NEW CODE
|
||||
|
||||
def test_val_loss(learn):
|
||||
assert learn.validate()[1] > 0.5
|
||||
|
||||
def test_bwdlm_lstm_can_be_trained():
|
||||
manual_seed()
|
||||
path, df_trn, df_val = prep_human_numbers()
|
||||
data = TextLMDataBunch.from_df(path, df_trn, df_val, tokenizer=Tokenizer(BaseTokenizer),
|
||||
lm_type = contrib_data.LanguageModelType.BiLM,
|
||||
ld_cls = contrib_data.LanguageModelLoader)
|
||||
|
||||
learn = bilm_learner(data, emb_sz=100, nl=1, drop_mult=0.1, qrnn=False)
|
||||
learn.fit_one_cycle(4, 5e-3)
|
||||
assert learn.validate()[1] > 0.5
|
||||
|
||||
def test_bilm_lstm_can_be_trained():
|
||||
manual_seed()
|
||||
path, df_trn, df_val = prep_human_numbers()
|
||||
data = TextLMDataBunch.from_df(path, df_trn, df_val, tokenizer=Tokenizer(BaseTokenizer),
|
||||
lm_type = contrib_data.LanguageModelType.BwdLM,
|
||||
ld_cls = contrib_data.LanguageModelLoader)
|
||||
|
||||
learn = language_model_learner(data, emb_sz=100, nl=1, drop_mult=0.1, qrnn=False)
|
||||
learn.fit_one_cycle(4, 5e-3)
|
||||
assert learn.validate()[1] > 0.5
|
||||
Reference in New Issue
Block a user