File size: 2,109 Bytes
bc45d7d 7279c87 bc45d7d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 | import torch
import torch.nn as nn
import torch.nn.functional as F
import transformers
class PositionEmbedding(nn.Module):
def __init__(self, config):
super().__init__()
self.position_embeddings = nn.Embedding(
config.max_position_embeddings, config.hidden_size
)
self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
self.dropout = nn.Dropout(config.hidden_dropout_prob)
self.register_buffer(
"position_ids", torch.arange(config.max_position_embeddings).expand((1, -1))
)
self.position_embedding_type = getattr(
config, "position_embedding_type", "absolute"
)
def forward(self, x):
input_shape = x.size()
seq_length = input_shape[1]
position_ids = self.position_ids[:, :seq_length]
position_embeddings = self.position_embeddings(position_ids)
embeddings = x + position_embeddings
embeddings = self.LayerNorm(embeddings)
embeddings = self.dropout(embeddings)
return embeddings
class Transformer(nn.Module):
def __init__(self, config, n_classes=50):
super().__init__()
self.l1 = nn.Linear(
in_features=config.input_size, out_features=config.hidden_size
)
self.embedding = PositionEmbedding(config)
setattr(config.model_config, "_attn_implementation", "eager")
self.layers = nn.ModuleList(
[
transformers.BertLayer(config.model_config)
for _ in range(config.num_hidden_layers)
]
)
self.l2 = nn.Linear(in_features=config.hidden_size, out_features=n_classes)
def forward(self, x):
x = self.l1(x)
x = self.embedding(x)
for layer in self.layers:
out = layer(x)
# Transformers <5 returned tuples; Transformers 5 returns a Tensor.
x = out[0] if isinstance(out, (tuple, list)) else out
x = torch.max(x, dim=1).values
x = F.dropout(x, p=0.2)
x = self.l2(x)
return x
|