From 03c4a2cd1599f5c9d7bceb51ce3eb1d344d4936b Mon Sep 17 00:00:00 2001 From: aayush Date: Fri, 15 Feb 2019 02:10:18 +0530 Subject: [PATCH] Fixed self attention, local tests passed --- fastai_contrib/learner.py | 4 ++- fastai_contrib/models.py | 53 +++++++++++++++++++++------------------ 2 files changed, 32 insertions(+), 25 deletions(-) diff --git a/fastai_contrib/learner.py b/fastai_contrib/learner.py index 67e79cb..d761b35 100644 --- a/fastai_contrib/learner.py +++ b/fastai_contrib/learner.py @@ -39,7 +39,9 @@ def bilm_text_classifier_learner(data: DataBunch, bptt: int = 70, max_len: int = ds = data.train_ds vocab_size, n_class = len(data.vocab.itos), data.c if bicls_head == 'BiPoolingLinearClassifier': - count = 3*2 + count = 3 * 2 + elif if bicls_head == 'BiAttentionPoolingClassifier': + count = 5 else: count = 3 layers = [emb_sz * count] + lin_ftrs + [n_class] diff --git a/fastai_contrib/models.py b/fastai_contrib/models.py index 8cae758..89e8514 100644 --- a/fastai_contrib/models.py +++ b/fastai_contrib/models.py @@ -71,37 +71,36 @@ class MultiBatchBiLMModel(BiLMModel): return self.concat(raw_outputs), self.concat(outputs) class BiAttentionPoolingClassifier(nn.Module): - r" [WIP] BiLM Pooling with self attention" + r" BiLM Pooling with self attention" - def __init__(self, layers:Collection[int], drops:Collection[float]): + def __init__(self, layers:Collection[int], drops:Collection[float], emb_sz:int): 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.self_attn = MultiHeadAttention(n_head=8, d_model=1, d_k=64, d_v=64, dropout=0.1) + self.self_attn = MultiHeadAttention(n_head=8, d_model=emb_sz, d_k=64, d_v=64, dropout=0.1) 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) + return f(x.permute(2, 0, 1), (1,)).view(bs,-1) def forward(self, input:Tuple[Tensor,Tensor])->Tuple[Tensor,Tensor,Tensor]: raw_outputs, outputs = input output = outputs[-1] assert len(output.size()) == 4, 'Expected input dimension 4' - sl, bs, 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 = x.unsqueeze(-1) - x, _ = self.self_attn(x, x, x) + bs, sl, em_sz, passes = output.size() + x = torch.cat([output[..., 0], output[..., 1]], 1) + x, _ = self.self_attn(x, x, x) + + avgpool = self.pool(x, bs, False) + mxpool = self.pool(x, bs, True) + + x = torch.cat([output[:,-1,..., 0], x, mxpool, + avgpool, output[:,-1,..., 1]], 1) x = self.layers(x) return x, raw_outputs, outputs @@ -111,7 +110,7 @@ class ScaledDotProductAttention(nn.Module): based on: https://github.com/jadore801120/attention-is-all-you-need-pytorch """ - def __init__(self, temperature, attn_dropout=0.1): + def __init__(self, temperature:float, attn_dropout:float=0.1): super().__init__() self.temperature = temperature self.dropout = nn.Dropout(attn_dropout) @@ -134,7 +133,7 @@ class MultiHeadAttention(nn.Module): based on: https://github.com/jadore801120/attention-is-all-you-need-pytorch """ - def __init__(self, n_head, d_model, d_k, d_v, dropout=0.1): + def __init__(self, n_head:int, d_model:int, d_k:int, d_v:int, dropout:float=0.1): super().__init__() self.n_head = n_head @@ -157,7 +156,6 @@ class MultiHeadAttention(nn.Module): self.dropout = nn.Dropout(dropout) def forward(self, q, k, v): - d_k, d_v, n_head = self.d_k, self.d_v, self.n_head @@ -240,7 +238,6 @@ class AvgPoolingLinearClassifier(nn.Module): x = self.layers(x) return x, raw_outputs, outputs - 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: @@ -270,12 +267,20 @@ def get_birnn_classifier(bptt:int, max_seq:int, n_class:int, vocab_sz:int, emb_s qrnn=qrnn, hidden_p=hidden_p, input_p=input_p, embed_p=embed_p, weight_p=weight_p) head = BiPoolingLinearClassifier - if bicls_head == 'BiPoolingLinearClassifier': head = BiPoolingLinearClassifier - elif bicls_head == 'AvgPoolingLinearClassifier': head = AvgPoolingLinearClassifier - elif bicls_head == 'BiAttentionPoolingClassifier': head = BiAttentionPoolingClassifier - - model = SequentialRNN(BiLMModel(fwd_rnn_enc, bwd_rnn_enc), head(layers, drops)) - model.reset() + + if bicls_head == 'BiPoolingLinearClassifier': + head = BiPoolingLinearClassifier + model = SequentialRNN(BiLMModel(fwd_rnn_enc, bwd_rnn_enc), head(layers, drops)) + elif bicls_head == 'AvgPoolingLinearClassifier': + head = AvgPoolingLinearClassifier + model = SequentialRNN(BiLMModel(fwd_rnn_enc, bwd_rnn_enc), head(layers, drops)) + elif bicls_head == 'BiAttentionPoolingClassifier': + head = BiAttentionPoolingClassifier + # attention requires an additional argument + # maybe use kwargs for initialising classes + model = SequentialRNN(BiLMModel(fwd_rnn_enc, bwd_rnn_enc), head(layers, drops, emb_sz)) + + model.reset() return model #endregion