chen shangheng
deepxml
63239ac
Raw
History Blame Contribute Delete
11.4 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers.activations import ACT2FN
import numpy as np
simplify_dim = 500
class SelfAttention(nn.Module):
def __init__(
self,
config,
):
super().__init__()
self.self = BartAttention(config.hidden_size, config.num_attention_heads, config.vocab_size - 2, config.attention_probs_dropout_prob)
self.layer_norm = nn.LayerNorm(config.hidden_size)
self.dropout = nn.Dropout(config.hidden_dropout_prob)
def forward(self, hidden_states,
attention_mask=None, output_attentions=False, extra_attn=None,):
residual = hidden_states
hidden_states, attn_weights, _ = self.self(
hidden_states=hidden_states, attention_mask=attention_mask, output_attentions=output_attentions,
extra_attn=extra_attn,
)
hidden_states = self.dropout(hidden_states)
hidden_states = residual + hidden_states
hidden_states = self.layer_norm(hidden_states)
outputs = (hidden_states,)
if output_attentions:
outputs += (attn_weights,)
return outputs
class BartAttention(nn.Module):
"""Multi-headed attention from 'Attention Is All You Need' paper"""
def __init__(
self,
embed_dim: int,
num_heads: int,
num_labels: int,
dropout: float = 0.0,
is_decoder: bool = False,
bias: bool = True,
):
super().__init__()
self.embed_dim = embed_dim
self.num_heads = num_heads
self.dropout = dropout
self.head_dim = embed_dim // num_heads
assert (
self.head_dim * num_heads == self.embed_dim
), f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`: {num_heads})."
self.scaling = self.head_dim ** -0.5
self.is_decoder = is_decoder
self.k_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
self.v_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
self.q_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
self.k_simplify_proj = nn.Linear(num_labels, simplify_dim, bias=bias)
self.v_simplify_proj = nn.Linear(num_labels, simplify_dim, bias=bias)
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()
def forward(
self,
hidden_states: torch.Tensor,
key_value_states=None,
past_key_value=None,
attention_mask=None,
output_attentions: bool = False,
extra_attn=None,
only_attn=False,
):
"""Input shape: Batch x Time x Channel"""
# if key_value_states are provided this layer is used as a cross-attention layer
# for the decoder
is_cross_attention = key_value_states is not None
bsz, tgt_len, embed_dim = hidden_states.size()
# get query proj
query_states = self.q_proj(hidden_states) * self.scaling
# get key, value proj
if is_cross_attention and past_key_value is not None:
# reuse k,v, cross_attentions
key_states = past_key_value[0]
value_states = past_key_value[1]
elif is_cross_attention:
# cross_attentions
key_states = self._shape(self.k_proj(key_value_states), -1, bsz)
value_states = self._shape(self.v_proj(key_value_states), -1, bsz)
elif past_key_value is not None:
# reuse k, v, self_attention
key_states = self._shape(self.k_proj(hidden_states), -1, bsz)
value_states = self._shape(self.v_proj(hidden_states), -1, bsz)
key_states = torch.cat([past_key_value[0], key_states], dim=2)
value_states = torch.cat([past_key_value[1], value_states], dim=2)
else:
# self_attention
key_states = self._shape(self.k_proj(hidden_states), -1, bsz)
value_states = self._shape(self.v_proj(hidden_states), -1, bsz)
if self.is_decoder:
# if cross_attention save Tuple(torch.Tensor, torch.Tensor) of all cross attention key/value_states.
# Further calls to cross_attention layer can then reuse all cross-attention
# key/value_states (first "if" case)
# if uni-directional self-attention (decoder) save Tuple(torch.Tensor, torch.Tensor) of
# all previous decoder key/value_states. Further calls to uni-directional self-attention
# can concat previous decoder key/value_states to current projected key/value_states (third "elif" case)
# if encoder bi-directional self-attention `past_key_value` is always `None`
past_key_value = (key_states, value_states)
proj_shape = (bsz * self.num_heads, -1, self.head_dim)
query_states = self._shape(query_states, tgt_len, bsz).view(*proj_shape)
key_states = key_states.view(*proj_shape).transpose(1, 2)
value_states = value_states.view(*proj_shape).transpose(1, 2)
src_len = key_states.size(1)
key_states = self.k_simplify_proj(key_states)
value_states = self.v_simplify_proj(value_states).transpose(1, 2)
attn_weights = torch.bmm(query_states, key_states)
if extra_attn is not None:
# extra_attn = self.attn_simplify_proj(extra_attn)
attn_weights += extra_attn
# assert attn_weights.size() == (
# bsz * self.num_heads,
# tgt_len,
# src_len,
# ), f"Attention weights should be of size {(bsz * self.num_heads, tgt_len, src_len)}, but is {attn_weights.size()}"
if attention_mask is not None:
# assert attention_mask.size() == (
# bsz,
# 1,
# tgt_len,
# src_len,
# ), f"Attention mask should be of size {(bsz, 1, tgt_len, src_len)}, but is {attention_mask.size()}"
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len) + attention_mask
attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
attn_weights = F.softmax(attn_weights, dim=-1)
if output_attentions:
# this operation is a bit akward, but it's required to
# make sure that attn_weights keeps its gradient.
# In order to do so, attn_weights have to reshaped
# twice and have to be reused in the following
attn_weights_reshaped = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
attn_weights = attn_weights_reshaped.view(bsz * self.num_heads, tgt_len, src_len)
else:
attn_weights_reshaped = None
if only_attn:
return attn_weights_reshaped
attn_weights = F.dropout(attn_weights, p=self.dropout, training=self.training)
attn_output = torch.bmm(attn_weights, value_states)
# assert attn_output.size() == (
# bsz * self.num_heads,
# tgt_len,
# self.head_dim,
# ), f"`attn_output` should be of size {(bsz, self.num_heads, tgt_len, self.head_dim)}, but is {attn_output.size()}"
attn_output = (
attn_output.view(bsz, self.num_heads, tgt_len, self.head_dim)
.transpose(1, 2)
.reshape(bsz, tgt_len, embed_dim)
)
attn_output = self.out_proj(attn_output)
return attn_output, attn_weights_reshaped, past_key_value
class GraphLayer(nn.Module):
def __init__(self, config, last=False):
super(GraphLayer, self).__init__()
self.config = config
class _Actfn(nn.Module):
def __init__(self):
super(_Actfn, self).__init__()
if isinstance(config.hidden_act, str):
self.intermediate_act_fn = ACT2FN[config.hidden_act]
else:
self.intermediate_act_fn = config.hidden_act
def forward(self, x):
return self.intermediate_act_fn(x)
self.hir_attn = SelfAttention(config)
self.output_layer = nn.Sequential(nn.Linear(config.hidden_size, config.intermediate_size),
_Actfn(),
nn.Linear(config.intermediate_size, config.hidden_size),
)
self.output_layer_norm = nn.LayerNorm(config.hidden_size)
self.dropout = nn.Dropout(config.hidden_dropout_prob)
def forward(self, label_emb, extra_attn, self_attn_mask):
label_emb = self.hir_attn(label_emb,
attention_mask=self_attn_mask, extra_attn=extra_attn)[0]
label_emb = self.output_layer_norm(self.dropout(self.output_layer(label_emb)) + label_emb)
return label_emb
class GraphEncoder(nn.Module):
def __init__(self, config, layer=2, graph_hierarchy=None, label_emb_init=None,emb_trainable=True, **kwargs):
super(GraphEncoder, self).__init__()
config.num_attention_heads = 2
self.config = config
config.vocab_size = label_emb_init.shape[0]
self.hir_layers = nn.ModuleList([GraphLayer(config, last=i == layer - 1) for i in range(layer)])
# config.num_hidden_layers
# GRAPH
self.label_name = torch.tensor(graph_hierarchy["classes"]).unsqueeze(1)
from deepxml.match import BertEmbeddings
self.initializer_range = 0.02
self.label_embeddings = BertEmbeddings(config, label_emb_init, emb_trainable, pos_trainable=False)
# config.hidden_size = 1
# self.dist_embeddings = BertEmbeddings(config, pos_trainable=False)
config.max_position_embeddings = 5
self.edge_embeddings = BertEmbeddings(config, pos_trainable=True)
self.edge_encoding= nn.Linear(simplify_dim, simplify_dim)
self.dist_embeddings= nn.Linear(config.vocab_size - 2, simplify_dim)
self.extra = nn.Linear(simplify_dim, simplify_dim)
self.label_num = graph_hierarchy["label_num"]
self.distance = torch.tensor(graph_hierarchy["distance_matrix"], dtype=torch.float)
self.edge = torch.tensor(graph_hierarchy["edge_matrix"])
def forward(self):
label_emb = self.label_embeddings(self.label_name).sum(dim=1)
label_emb = label_emb.unsqueeze(0)
expand_size = label_emb.size(-2) // self.label_name.size(0)
extra_attn = None
edge_encodings = torch.nn.functional.elu(self.edge_embeddings(self.edge).view(self.label_num, -1))
edge_encodings = self.edge_encoding(edge_encodings)
extra_attn = self.dist_embeddings(self.distance) + edge_encodings
extra_attn = extra_attn.view(self.label_num, 1, simplify_dim, 1).expand(-1, expand_size, -1, expand_size)
extra_attn = extra_attn.reshape(1, self.label_num * expand_size, -1)
extra_attn = torch.relu(self.extra(extra_attn))
self_attn_mask = None
for hir_layer in self.hir_layers:
label_emb = hir_layer(label_emb, extra_attn, self_attn_mask)
return label_emb