Spaces:
Sleeping
Sleeping
| 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 | |