mirror of
https://github.com/wassname/multifit.git
synced 2026-09-09 11:27:26 +08:00
Added initial classes and changes for BiLM implementation
This commit is contained in:
@@ -0,0 +1,34 @@
|
||||
from fastai.callbacks import *
|
||||
from fastai.basic_data import *
|
||||
from fastai.datasets import untar_data
|
||||
from fastai_contrib.models import get_bilm
|
||||
from fastai.text.learner import *
|
||||
|
||||
|
||||
def bilm_learner(data:DataBunch, bptt:int=70, emb_sz:int=400, nh:int=1150, nl:int=3, pad_token:int=1,
|
||||
drop_mult:float=1., tie_weights:bool=True, bias:bool=True, qrnn:bool=False, pretrained_model=None,
|
||||
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
|
||||
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)
|
||||
if pretrained_model is not None:
|
||||
model_path = untar_data(pretrained_model, data=False)
|
||||
fnames = [list(model_path.glob(f'*.{ext}'))[0] for ext in ['pth', 'pkl']]
|
||||
learn.load_pretrained(*fnames)
|
||||
learn.freeze()
|
||||
if pretrained_fnames is not None:
|
||||
fnames = [learn.path/learn.model_dir/f'{fn}.{ext}' for fn,ext in zip(pretrained_fnames, ['pth', 'pkl'])]
|
||||
learn.load_pretrained(*fnames)
|
||||
learn.freeze()
|
||||
return learn
|
||||
|
||||
|
||||
def bilm_split(model:nn.Module) -> List[nn.Module]:
|
||||
"Split a RNN `model` in groups for differential learning rates."
|
||||
groups = [[rnn, dp] for rnn, dp in zip(model[0].forward_rnns, model[0].hidden_dps)]
|
||||
groups += [[rnn, dp] for rnn, dp in zip(model[0].backward_rnns, model[0].hidden_dps)]
|
||||
groups.append([model[0].encoder, model[0].encoder_dp, model[1]])
|
||||
return groups
|
||||
@@ -0,0 +1,85 @@
|
||||
from fastai.torch_core import *
|
||||
from fastai.layers import *
|
||||
from fastai.text.models import *
|
||||
|
||||
|
||||
class BiLMCore(nn.Module):
|
||||
"""
|
||||
AWD-LSTM/QRNN inspired by https://arxiv.org/abs/1708.02182.
|
||||
Inspired by https://github.com/allenai/allennlp/blob/master/allennlp/models/bidirectional_lm.py#L65
|
||||
"""
|
||||
initrange=0.1
|
||||
|
||||
def __init__(self, vocab_sz:int, emb_sz:int, n_hid:int, n_layers:int, pad_token:int, bidir:bool=False,
|
||||
hidden_p:float=0.2, input_p:float=0.6, embed_p:float=0.1, weight_p:float=0.5, qrnn:bool=False):
|
||||
|
||||
super().__init__()
|
||||
self.bs,self.qrnn,self.ndir = 1, qrnn,(2 if bidir else 1)
|
||||
self.emb_sz,self.n_hid,self.n_layers = emb_sz,n_hid,n_layers
|
||||
# embeddings are shared between forward and backward LMs
|
||||
self.encoder = nn.Embedding(vocab_sz, emb_sz, padding_idx=pad_token)
|
||||
self.encoder_dp = EmbeddingDropout(self.encoder, embed_p)
|
||||
if self.qrnn:
|
||||
#Using QRNN requires cupy: https://github.com/cupy/cupy
|
||||
from fastai.text.qrnn.qrnn import QRNNLayer
|
||||
|
||||
def create_qrnn_layers():
|
||||
return [QRNNLayer(emb_sz if l == 0 else n_hid, (n_hid if l != n_layers - 1 else emb_sz)//self.ndir,
|
||||
save_prev_x=True, zoneout=0, window=2 if l == 0 else 1, output_gate=True,
|
||||
use_cuda=torch.cuda.is_available()) for l in range(n_layers)]
|
||||
self.forward_rnns = create_qrnn_layers()
|
||||
self.backward_rnns = create_qrnn_layers()
|
||||
for rnn in self.forward_rnns + self.backward_rnns:
|
||||
rnn.linear = WeightDropout(rnn.linear, weight_p, layer_names=['weight'])
|
||||
else:
|
||||
def create_lstm_layers():
|
||||
return [nn.LSTM(emb_sz if l == 0 else n_hid, (n_hid if l != n_layers - 1 else emb_sz)//self.ndir,
|
||||
1, bidirectional=False) for l in range(n_layers)]
|
||||
self.forward_rnns = [WeightDropout(rnn, weight_p) for rnn in create_lstm_layers()]
|
||||
self.backward_rnns = [WeightDropout(rnn, weight_p) for rnn in create_lstm_layers()]
|
||||
self.forward_rnns = torch.nn.ModuleList(self.forward_rnns)
|
||||
self.backward_rnns = torch.nn.ModuleList(self.backward_rnns)
|
||||
self.encoder.weight.data.uniform_(-self.initrange, self.initrange)
|
||||
self.input_dp = RNNDropout(input_p)
|
||||
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()
|
||||
if bs!=self.bs:
|
||||
self.bs=bs
|
||||
self.reset()
|
||||
raw_output = self.input_dp(self.encoder_dp(input))
|
||||
|
||||
# TODO get reverse input and compute backward representation
|
||||
new_hidden,raw_outputs,outputs = [],[],[]
|
||||
for l, (rnn,hid_dp) in enumerate(zip(self.forward_rnns, self.hidden_dps)):
|
||||
raw_output, new_h = rnn(raw_output, self.hidden[l])
|
||||
new_hidden.append(new_h)
|
||||
raw_outputs.append(raw_output)
|
||||
if l != self.n_layers - 1: raw_output = hid_dp(raw_output)
|
||||
outputs.append(raw_output)
|
||||
self.hidden = to_detach(new_hidden)
|
||||
return raw_outputs, outputs
|
||||
|
||||
def _one_hidden(self, l:int)->Tensor:
|
||||
"Return one hidden state."
|
||||
nh = (self.n_hid if l != self.n_layers - 1 else self.emb_sz)//self.ndir
|
||||
return self.weights.new(self.ndir, self.bs, nh).zero_()
|
||||
|
||||
def reset(self):
|
||||
"Reset the hidden states."
|
||||
[r.reset() for r in self.forward_rnns if hasattr(r, 'reset')]
|
||||
[r.reset() for r in self.backward_rnns if hasattr(r, 'reset')]
|
||||
self.weights = next(self.parameters()).data
|
||||
if self.qrnn: self.hidden = [self._one_hidden(l) for l in range(self.n_layers)]
|
||||
else: self.hidden = [(self._one_hidden(l), self._one_hidden(l)) for l in range(self.n_layers)]
|
||||
|
||||
|
||||
def get_bilm(vocab_sz:int, emb_sz:int, n_hid:int, n_layers:int, pad_token:int, tie_weights:bool=True,
|
||||
qrnn:bool=False, bias:bool=True, bidir:bool=False, output_p:float=0.4, hidden_p:float=0.2, input_p:float=0.6,
|
||||
embed_p:float=0.1, weight_p:float=0.5)->nn.Module:
|
||||
"Create a full AWD-LSTM."
|
||||
rnn_enc = BiLMCore(vocab_sz, emb_sz, n_hid=n_hid, n_layers=n_layers, pad_token=pad_token, qrnn=qrnn, bidir=bidir,
|
||||
hidden_p=hidden_p, input_p=input_p, embed_p=embed_p, weight_p=weight_p)
|
||||
enc = rnn_enc.encoder if tie_weights else None
|
||||
return SequentialRNN(rnn_enc, LinearDecoder(vocab_sz, emb_sz, output_p, tie_encoder=enc, bias=bias))
|
||||
@@ -14,6 +14,7 @@ from fastai.text import LanguageModelLoader, get_language_model, RNNLearner, Tex
|
||||
import torch
|
||||
from fastai_contrib.utils import read_file, read_whitespace_file,\
|
||||
DataStump, validate, PAD, UNK
|
||||
from fastai_contrib.learner import bilm_learner
|
||||
|
||||
import pickle
|
||||
|
||||
@@ -27,7 +28,8 @@ from collections import Counter
|
||||
|
||||
|
||||
def pretrain_lm(dir_path, cuda_id=0, qrnn=True, clean=True, max_vocab=60000,
|
||||
bs=70, bptt=70, name='wt-103', model_dir='models', num_epochs=10):
|
||||
bs=70, bptt=70, name='wt-103', model_dir='models', num_epochs=10,
|
||||
bidir=True):
|
||||
"""
|
||||
:param dir_path: The path to the directory of the file.
|
||||
:param cuda_id: The id of the GPU. Uses GPU 0 by default or no GPU when
|
||||
@@ -39,6 +41,7 @@ def pretrain_lm(dir_path, cuda_id=0, qrnn=True, clean=True, max_vocab=60000,
|
||||
:param bptt: The back-propagation-through-time sequence length.
|
||||
:param name: The name used for both the model and the vocabulary.
|
||||
:param model_dir: The path to the directory where the models should be saved
|
||||
:param bidir: whether the language model is bidirectional
|
||||
"""
|
||||
if not torch.cuda.is_available():
|
||||
print('CUDA not available. Setting device=-1.')
|
||||
@@ -109,9 +112,11 @@ def pretrain_lm(dir_path, cuda_id=0, qrnn=True, clean=True, max_vocab=60000,
|
||||
drop_mult = 0.1
|
||||
|
||||
fastai.text.learner.default_dropout['language'] = dps
|
||||
learn = language_model_learner(data_lm, bptt=bptt, emb_sz=emb_sz, nh=nh, nl=nl, pad_token=1,
|
||||
drop_mult=drop_mult, tie_weights=True,
|
||||
bias=True, qrnn=True, clip=0.12)
|
||||
|
||||
lm_learner = bilm_learner if bidir else language_model_learner
|
||||
learn = lm_learner(data_lm, bptt=bptt, emb_sz=emb_sz, nh=nh, nl=nl, pad_token=1,
|
||||
drop_mult=drop_mult, tie_weights=True,
|
||||
bias=True, qrnn=qrnn, clip=0.12)
|
||||
# compared to standard Adam, we set beta_1 to 0.8
|
||||
learn.opt_fn = partial(optim.Adam, betas=(0.8, 0.99))
|
||||
learn.true_wd = False
|
||||
|
||||
Reference in New Issue
Block a user