mirror of
https://github.com/wassname/multifit.git
synced 2026-09-10 12:12:50 +08:00
Fix BiLM training after update to newest fastai
This commit is contained in:
@@ -59,7 +59,7 @@ class LanguageModelLoader(): # copy of the original LanguageModelLoader
|
||||
seq_len = min(seq_len, data.shape[1] - 1 - i)
|
||||
x = data[:,i:i+seq_len]
|
||||
y = data[:,i+1:i+1+seq_len]
|
||||
#y = y.view(-1, 2) if self.lm_type == LanguageModelType.BiLM else y.view(-1)
|
||||
y = y.contiguous().view(-1, 2) if self.lm_type == LanguageModelType.BiLM else y.contiguous().view(-1)
|
||||
return x,y
|
||||
|
||||
#endregion
|
||||
|
||||
+36
-26
@@ -6,10 +6,11 @@ from fastai.text.models import *
|
||||
|
||||
class BiLMModel(nn.Module):
|
||||
|
||||
def __init__(self, fwd_lm:nn.Module, bwd_lm:nn.Module):
|
||||
def __init__(self, fwd_lm:nn.Module, bwd_lm:nn.Module, squash_bs_sl=False):
|
||||
super().__init__()
|
||||
self.fwd_lm = fwd_lm
|
||||
self.bwd_lm = bwd_lm
|
||||
self.squash_bs_sl = squash_bs_sl
|
||||
|
||||
def __getitem__(self, idx):
|
||||
return BiLMModel(self.fwd_lm[idx], self.bwd_lm[idx])
|
||||
@@ -35,49 +36,57 @@ class BiLMModel(nn.Module):
|
||||
fwd_o = self.fwd_lm(f)
|
||||
bwd_o = self.bwd_lm(b)
|
||||
|
||||
return self.stack(fwd_o, bwd_o)
|
||||
outs = self.stack(fwd_o, bwd_o)
|
||||
if self.squash_bs_sl:
|
||||
o = outs[0]
|
||||
o = o.view(o.shape[0]*o.shape[1],o.shape[2],o.shape[3])
|
||||
outs[0] = o
|
||||
return outs
|
||||
|
||||
def reset(self):
|
||||
"Reset the hidden states of underlaying lms."
|
||||
self.fwd_lm.reset()
|
||||
self.bwd_lm.reset()
|
||||
|
||||
class MultiBatchBiLMModel(BiLMModel):
|
||||
"Create a RNNCore module that can process a full sentence."
|
||||
|
||||
class BiPoolingLinearClassifier(nn.Module):
|
||||
def __init__(self, bptt:int, max_seq:int, *args, **kwargs):
|
||||
self.max_seq,self.bptt = max_seq,bptt
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def concat(self, arrs:Collection[Tensor])->Tensor:
|
||||
"Concatenate the `arrs` along the batch dimension."
|
||||
return [torch.cat([l[si] for l in arrs], dim=1) for si in range_of(arrs[0])]
|
||||
|
||||
def forward(self, input:LongTensor)->Tuple[Tensor,Tensor]:
|
||||
bs,sl = input.size()
|
||||
self.reset()
|
||||
raw_outputs, outputs = [],[]
|
||||
for i in range(0, sl, self.bptt):
|
||||
r, o = super().forward(input[:,i: min(i+self.bptt, sl)])
|
||||
if i>(sl-self.max_seq):
|
||||
raw_outputs.append(r)
|
||||
outputs.append(o)
|
||||
return self.concat(raw_outputs), self.concat(outputs)
|
||||
|
||||
class BiPoolingLinearClassifier(PoolingLinearClassifier):
|
||||
"Create a linear classifier with pooling."
|
||||
|
||||
def __init__(self, layers:Collection[int], drops:Collection[float]):
|
||||
super().__init__()
|
||||
mod_layers = []
|
||||
activs = [nn.ReLU(inplace=True)] * (len(layers) - 2) + [None]
|
||||
for n_in,n_out,p,actn in zip(layers[:-1],layers[1:], drops, activs):
|
||||
mod_layers += bn_drop_lin(n_in, n_out, p=p, actn=actn)
|
||||
self.layers = nn.Sequential(*mod_layers)
|
||||
|
||||
def pool(self, x:Tensor, bs:int, is_max:bool):
|
||||
"Pool the tensor along the seq_len dimension."
|
||||
f = F.adaptive_max_pool1d if is_max else F.adaptive_avg_pool1d
|
||||
return f(x.permute(1,2,0), (1,)).view(bs,-1)
|
||||
|
||||
def forward(self, input:Tuple[Tensor,Tensor])->Tuple[Tensor,Tensor,Tensor]:
|
||||
raw_outputs, outputs = input
|
||||
output = outputs[-1]
|
||||
if len(output.size()) == 3:
|
||||
sl,bs,_ = output.size()
|
||||
avgpool = self.pool(output, bs, False)
|
||||
mxpool = self.pool(output, bs, True)
|
||||
x = torch.cat([output[-1], mxpool, avgpool], 1)
|
||||
x = self.layers(x)
|
||||
return x, raw_outputs, outputs
|
||||
return super().forward(input)
|
||||
elif len(output.size()) == 4:
|
||||
sl, bs, em_sz, passes = output.size()
|
||||
bs, sl, em_sz, passes = output.size()
|
||||
|
||||
f_avgpool = self.pool(output[...,0], bs, False)
|
||||
f_mxpool = self.pool(output[...,0], bs, True)
|
||||
b_avgpool = self.pool(output[..., 1], bs, False)
|
||||
b_mxpool = self.pool(output[..., 1], bs, True)
|
||||
x = torch.cat([output[-1][..., 0], f_mxpool, f_avgpool,
|
||||
output[-1][..., 1], b_mxpool, b_avgpool,], 1)
|
||||
x = torch.cat([output[:,-1,..., 0], f_mxpool, f_avgpool,
|
||||
output[:,-1,..., 1], b_mxpool, b_avgpool,], 1)
|
||||
x = self.layers(x)
|
||||
return x, raw_outputs, outputs
|
||||
|
||||
@@ -134,7 +143,8 @@ def get_bilm(vocab_sz:int, emb_sz:int, n_hid:int, n_layers:int, pad_token:int, t
|
||||
|
||||
return BiLMModel(
|
||||
fwd_lm=SequentialRNN(fwd_rnn_enc, LinearDecoder(vocab_sz, emb_sz, output_p, tie_encoder=enc, bias=bias)),
|
||||
bwd_lm=SequentialRNN(bwd_rnn_enc, LinearDecoder(vocab_sz, emb_sz, output_p, tie_encoder=enc, bias=bias)))
|
||||
bwd_lm=SequentialRNN(bwd_rnn_enc, LinearDecoder(vocab_sz, emb_sz, output_p, tie_encoder=enc, bias=bias)),
|
||||
squash_bs_sl=True)
|
||||
|
||||
def get_birnn_classifier(bptt:int, max_seq:int, n_class:int, vocab_sz:int, emb_sz:int, n_hid:int, n_layers:int,
|
||||
pad_token:int, layers:Collection[int], drops:Collection[float], bidir:bool=False, qrnn:bool=False,
|
||||
|
||||
@@ -31,6 +31,8 @@ CLASSES = ['neg', 'pos', 'unsup']
|
||||
number_match_re = re.compile(r'^([0-9]+[,.]?)+$')
|
||||
number_split_re = re.compile(r'([,.])')
|
||||
|
||||
# FIXME: coping of tokens from one sentencepiece model to another does not work for 50% of tokens
|
||||
# FIXME: tokens in sentencepiece are uppercase eventhough post-transformation will convert them to lowercase
|
||||
class MosesTokenizerFunc(BaseTokenizer):
|
||||
"Wrapper around a MosesTokenizer to make it a `BaseTokenizer`."
|
||||
def __init__(self, lang:str):
|
||||
@@ -278,7 +280,6 @@ def read_clas_data(dir_path, dataset, lang) -> Tuple[Dict[str, List[List[str]]],
|
||||
2. a dictionary mapping splits to a list of labels
|
||||
"""
|
||||
processors = {
|
||||
'imdb': read_imdb,
|
||||
'xnli': read_xnli
|
||||
}
|
||||
processor = processors[dataset]
|
||||
|
||||
@@ -56,7 +56,7 @@ def test_ulmfit_works_with_relative_paths():
|
||||
exp = ulmfit.pretrain_lm.LMHyperParams(
|
||||
dataset_path=wt2.relative_to(Path.cwd()),
|
||||
lang='en',
|
||||
qrnn=True,
|
||||
qrnn=False,
|
||||
max_vocab=1000,
|
||||
name=lm_name,
|
||||
cuda_id=cuda_id)
|
||||
@@ -83,7 +83,7 @@ def test_ulmfit_default_end_to_end():
|
||||
exp = ulmfit.pretrain_lm.LMHyperParams(
|
||||
dataset_path=wt2,
|
||||
lang='en',
|
||||
qrnn=True,
|
||||
qrnn=False,
|
||||
max_vocab=1000,
|
||||
name=lm_name,
|
||||
cuda_id=cuda_id)
|
||||
@@ -105,7 +105,7 @@ def test_ulmfit_fastai_end_to_end():
|
||||
dataset_path=wt2,
|
||||
lang='en',
|
||||
cuda_id=cuda_id,
|
||||
qrnn=True,
|
||||
qrnn=False,
|
||||
tokenizer='f',
|
||||
max_vocab=100,
|
||||
name=lm_name,
|
||||
@@ -124,7 +124,7 @@ def test_ulmfit_fastai_bidir_end_to_end():
|
||||
dataset_path=wt2,
|
||||
lang='en',
|
||||
cuda_id=cuda_id,
|
||||
qrnn=True,
|
||||
qrnn=False,
|
||||
bidir=True,
|
||||
tokenizer='f',
|
||||
max_vocab=100,
|
||||
@@ -144,7 +144,7 @@ def test_ulmfit_moses_fa_bidir_end_to_end():
|
||||
dataset_path=wt2,
|
||||
lang='en',
|
||||
cuda_id=cuda_id,
|
||||
qrnn=True,
|
||||
qrnn=False,
|
||||
bidir=True,
|
||||
tokenizer='vf',
|
||||
max_vocab=100,
|
||||
@@ -168,15 +168,15 @@ def test_ulmfit_sentencepiece_end_to_end():
|
||||
dataset_path=wt2,
|
||||
lang='en',
|
||||
cuda_id=cuda_id,
|
||||
qrnn=True,
|
||||
qrnn=False,
|
||||
tokenizer=ulmfit.pretrain_lm.Tokenizers.SUBWORD,
|
||||
max_vocab=100,
|
||||
max_vocab=200,
|
||||
name=lm_name,
|
||||
)
|
||||
exp.train_lm(num_epochs=1, bs=2)
|
||||
# not supported yet
|
||||
# exp2 = ulmfit.train_clas.CLSHyperParams.from_lm(test_data / 'imdb', exp.model_dir)
|
||||
# exp2.train_cls(num_lm_epochs=0, unfreeze=False, bs=4, )
|
||||
exp2 = ulmfit.train_clas.CLSHyperParams.from_lm(test_data / 'imdb', exp.model_dir)
|
||||
exp2.train_cls(num_lm_epochs=0, unfreeze=False, bs=4, )
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user