lightning-30m-ft / modeling_lightning.py
Aobangaming's picture
Update modeling_lightning.py
c2f0c02 verified
Raw History Blame Contribute Delete
5.89 kB
import torch
import torch.nn as nn
from transformers import PreTrainedModel, PretrainedConfig, GenerationMixin
from transformers.modeling_outputs import CausalLMOutput
import torch
import torch.nn.functional as F
import torch.nn as nn
import math
embedding = 256
heads = 4
layers = 4
dropout = 0.1
msl = 160
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=5000):
super(PositionalEncoding, self).__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0).transpose(0, 1)
self.register_buffer('pe', pe)
def forward(self, x):
return x + self.pe[:x.size(1), :].transpose(0, 1)
class CausalSelfAttention(nn.Module):
def __init__(self, d_model, nhead, dropout=0.1):
super().__init__()
assert d_model % nhead == 0
self.nhead = nhead
self.head_dim = d_model // nhead
self.dropout = dropout
self.qkv = nn.Linear(d_model, d_model * 3)
self.out_proj = nn.Linear(d_model, d_model)
def forward(self, x):
B, T, C = x.shape
# Create Q, K, V
q, k, v = self.qkv(x).chunk(3, dim=-1)
# [B, T, C] -> [B, heads, T, head_dim]
q = q.view(B, T, self.nhead, self.head_dim).transpose(1, 2)
k = k.view(B, T, self.nhead, self.head_dim).transpose(1, 2)
v = v.view(B, T, self.nhead, self.head_dim).transpose(1, 2)
y = F.scaled_dot_product_attention(
q,
k,
v,
attn_mask=None,
dropout_p=self.dropout if self.training else 0.0,
is_causal=True
)
y = y.transpose(1, 2).contiguous().view(B, T, C)
return self.out_proj(y)
class TransformerBlock(nn.Module):
def __init__(self, d_model, nhead, dropout=0.1):
super().__init__()
self.norm1 = nn.LayerNorm(d_model)
self.attention = CausalSelfAttention(
d_model,
nhead,
dropout
)
self.norm2 = nn.LayerNorm(d_model)
self.ffn = nn.Sequential(
nn.Linear(d_model, d_model * 4),
nn.GELU(),
nn.Linear(d_model * 4, d_model),
nn.Dropout(dropout)
)
def forward(self, x):
x = x + self.attention(self.norm1(x))
x = x + self.ffn(self.norm2(x))
return x
class TransformerLanguageModel(nn.Module):
def __init__(
self,
vocab_size,
d_model=512,
nhead=8,
num_layers=8,
dropout=0.1,
max_seq_len=160
):
super().__init__()
self.d_model = d_model
self.max_seq_len = max_seq_len
self.token_embedding = nn.Embedding(
vocab_size,
d_model
)
self.positional_encoding = PositionalEncoding(
d_model,
max_seq_len
)
self.transformer = nn.ModuleList([
TransformerBlock(
d_model,
nhead,
dropout
)
for _ in range(num_layers)
])
self.final_norm = nn.LayerNorm(d_model)
self.output_layer = nn.Linear(
d_model,
vocab_size,
bias=False
)
def forward(self, src):
x = self.token_embedding(src)
x = self.positional_encoding(x)
for layer in self.transformer:
x = layer(x)
x = self.final_norm(x)
return self.output_layer(x)
class LightningConfig(PretrainedConfig):
model_type = "lightning"
def __init__(
self,
vocab_size=50000,
d_model=256,
nhead=4,
num_layers=4,
dropout=0.1,
max_seq_len=160,
**kwargs
):
super().__init__(
tie_word_embeddings=False,
**kwargs
)
self.vocab_size = vocab_size
self.d_model = d_model
self.nhead = nhead
self.num_layers = num_layers
self.dropout = dropout
self.num_hidden_layers = num_layers
self.num_attention_heads = nhead
self.hidden_size = d_model
self.max_seq_len = max_seq_len
class LightningForCausalLM(PreTrainedModel, GenerationMixin):
config_class = LightningConfig
base_model_prefix = "lightning"
def __init__(self, config):
super().__init__(config)
self.lightning = TransformerLanguageModel(
vocab_size=config.vocab_size,
d_model=config.d_model,
nhead=config.nhead,
num_layers=config.num_layers,
dropout=config.dropout,
max_seq_len=config.max_seq_len
)
self.post_init()
def forward(self, input_ids=None, labels=None, **kwargs):
logits = self.lightning(input_ids)
loss = None
if labels is not None:
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
loss_fn = nn.CrossEntropyLoss()
loss = loss_fn(
shift_logits.view(-1, shift_logits.size(-1)),
shift_labels.view(-1)
)
return CausalLMOutput(
loss=loss,
logits=logits
)
def get_input_embeddings(self):
return self.lightning.token_embedding
def set_input_embeddings(self, value):
self.lightning.token_embedding = value
def get_output_embeddings(self):
return self.lightning.output_layer
def set_output_embeddings(self, new_embeddings):
self.lightning.output_layer = new_embeddings