Aobangaming commited on
Commit
72f0568
·
verified ·
1 Parent(s): 18a8aca

Update modeling_lightning.py

Browse files
Files changed (1) hide show
  1. modeling_lightning.py +156 -1
modeling_lightning.py CHANGED
@@ -5,7 +5,162 @@ import torch.nn as nn
5
  from transformers import PreTrainedModel, PretrainedConfig
6
  from transformers.modeling_outputs import CausalLMOutput
7
 
8
- from model import TransformerLanguageModel
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9
 
10
 
11
  class LightningConfig(PretrainedConfig):
 
5
  from transformers import PreTrainedModel, PretrainedConfig
6
  from transformers.modeling_outputs import CausalLMOutput
7
 
8
+ import torch
9
+ import torch.nn.functional as F
10
+ import torch.nn as nn
11
+ import math
12
+
13
+ embedding = 256
14
+ heads = 4
15
+ layers = 4
16
+ dropout = 0.1
17
+ msl = 160
18
+
19
+ class PositionalEncoding(nn.Module):
20
+
21
+ def __init__(self, d_model, max_len=5000):
22
+
23
+ super(PositionalEncoding, self).__init__()
24
+
25
+ pe = torch.zeros(max_len, d_model)
26
+
27
+ position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
28
+
29
+ div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
30
+
31
+ pe[:, 0::2] = torch.sin(position * div_term)
32
+
33
+ pe[:, 1::2] = torch.cos(position * div_term)
34
+
35
+ pe = pe.unsqueeze(0).transpose(0, 1)
36
+
37
+ self.register_buffer('pe', pe)
38
+
39
+
40
+ def forward(self, x):
41
+
42
+ return x + self.pe[:x.size(1), :].transpose(0, 1)
43
+
44
+
45
+ class CausalSelfAttention(nn.Module):
46
+ def __init__(self, d_model, nhead, dropout=0.1):
47
+ super().__init__()
48
+
49
+ assert d_model % nhead == 0
50
+
51
+ self.nhead = nhead
52
+ self.head_dim = d_model // nhead
53
+ self.dropout = dropout
54
+
55
+ self.qkv = nn.Linear(d_model, d_model * 3)
56
+ self.out_proj = nn.Linear(d_model, d_model)
57
+
58
+ def forward(self, x):
59
+ B, T, C = x.shape
60
+
61
+ # Create Q, K, V
62
+ q, k, v = self.qkv(x).chunk(3, dim=-1)
63
+
64
+ # [B, T, C] -> [B, heads, T, head_dim]
65
+ q = q.view(B, T, self.nhead, self.head_dim).transpose(1, 2)
66
+ k = k.view(B, T, self.nhead, self.head_dim).transpose(1, 2)
67
+ v = v.view(B, T, self.nhead, self.head_dim).transpose(1, 2)
68
+
69
+ y = F.scaled_dot_product_attention(
70
+ q,
71
+ k,
72
+ v,
73
+ attn_mask=None,
74
+ dropout_p=self.dropout if self.training else 0.0,
75
+ is_causal=True
76
+ )
77
+
78
+ y = y.transpose(1, 2).contiguous().view(B, T, C)
79
+
80
+ return self.out_proj(y)
81
+
82
+
83
+ class TransformerBlock(nn.Module):
84
+ def __init__(self, d_model, nhead, dropout=0.1):
85
+ super().__init__()
86
+
87
+ self.norm1 = nn.LayerNorm(d_model)
88
+ self.attention = CausalSelfAttention(
89
+ d_model,
90
+ nhead,
91
+ dropout
92
+ )
93
+
94
+ self.norm2 = nn.LayerNorm(d_model)
95
+
96
+ self.ffn = nn.Sequential(
97
+ nn.Linear(d_model, d_model * 4),
98
+ nn.GELU(),
99
+ nn.Linear(d_model * 4, d_model),
100
+ nn.Dropout(dropout)
101
+ )
102
+
103
+ def forward(self, x):
104
+ x = x + self.attention(self.norm1(x))
105
+
106
+ x = x + self.ffn(self.norm2(x))
107
+
108
+ return x
109
+
110
+
111
+ class TransformerLanguageModel(nn.Module):
112
+ def __init__(
113
+ self,
114
+ vocab_size,
115
+ d_model=512,
116
+ nhead=8,
117
+ num_layers=8,
118
+ dropout=0.1,
119
+ max_seq_len=160
120
+ ):
121
+ super().__init__()
122
+
123
+ self.d_model = d_model
124
+ self.max_seq_len = max_seq_len
125
+
126
+ self.token_embedding = nn.Embedding(
127
+ vocab_size,
128
+ d_model
129
+ )
130
+
131
+ self.positional_encoding = PositionalEncoding(
132
+ d_model,
133
+ max_seq_len
134
+ )
135
+
136
+ self.transformer = nn.ModuleList([
137
+ TransformerBlock(
138
+ d_model,
139
+ nhead,
140
+ dropout
141
+ )
142
+ for _ in range(num_layers)
143
+ ])
144
+
145
+ self.final_norm = nn.LayerNorm(d_model)
146
+
147
+ self.output_layer = nn.Linear(
148
+ d_model,
149
+ vocab_size,
150
+ bias=False
151
+ )
152
+
153
+ def forward(self, src):
154
+ x = self.token_embedding(src)
155
+
156
+ x = self.positional_encoding(x)
157
+
158
+ for layer in self.transformer:
159
+ x = layer(x)
160
+
161
+ x = self.final_norm(x)
162
+
163
+ return self.output_layer(x)
164
 
165
 
166
  class LightningConfig(PretrainedConfig):