Skip to content

Instantly share code, notes, and snippets.

@lake-hope
Last active March 22, 2018 03:43
Show Gist options
  • Select an option

  • Save lake-hope/25c537fe72084b8a1ef34525e31b46e6 to your computer and use it in GitHub Desktop.

Select an option

Save lake-hope/25c537fe72084b8a1ef34525e31b46e6 to your computer and use it in GitHub Desktop.
Word-level attention from Yang et. al 2016 in PyTorch
import torch
from torch import nn
import torch.nn.functional as F
class WordLevelAttention(nn.Module):
# this follows the word-level attention from Yang et al. 2016
# https://www.cs.cmu.edu/~diyiy/docs/naacl16.pdf
def __init__(self, n_hidden, *, batch_first=False):
super().__init__()
self.mlp = nn.Linear(n_hidden, n_hidden)
# word context vector
self.u_w = nn.Parameter(torch.rand(n_hidden))
self.batch_first = batch_first
def forward(self, X):
if not self.batch_first:
# make the input (batch_size, timesteps, features)
X = X.transpose(1, 0)
# get the hidden representation of the sequence
u_it = F.tanh(self.mlp(X))
# get attention weights for each timestep
alpha = F.softmax(torch.matmul(u_it, self.u_w), dim=1)
# get the weighted representation of the sequence
# and then get the sum
# (add a size 1 dimension to alpha so each time step's features could be scaled)
weighted_sequence = X * alpha.unsqueeze(2)
out = torch.sum(weighted_sequence, dim=1)
return out, alpha
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment