babybabellm-gpt_bert-deu-causal / modeling_gpt_bert.py
haznitrama's picture
Rehost from suchirsalhan/babybabellm-mono-deu
bfc11a9 verified
Raw
History Blame
5.73 kB
import math, torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import PreTrainedModel
from transformers.modeling_outputs import MaskedLMOutput, CausalLMOutputWithCrossAttentions
from .configuration_gpt_bert import GPTBertConfig
class GeGLU(nn.Module):
def forward(self, x):
x, gate = x.chunk(2, dim=-1)
return x * F.gelu(gate, approximate='tanh')
class FeedForward(nn.Module):
def __init__(self, config):
super().__init__()
self.mlp = nn.Sequential(
nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps, elementwise_affine=False)
,nn.Linear(config.hidden_size, 2*config.intermediate_size, bias=False)
,GeGLU()
,nn.LayerNorm(config.intermediate_size, eps=config.layer_norm_eps, elementwise_affine=False)
,nn.Linear(config.intermediate_size, config.hidden_size, bias=False)
,nn.Dropout(config.hidden_dropout_prob)
)
self._init(config.hidden_size)
def _init(self, h):
std = math.sqrt(2.0 / (5.0 * h))
nn.init.trunc_normal_(self.mlp[1].weight, mean=0.0, std=std, a=-2*std, b=2*std)
nn.init.trunc_normal_(self.mlp[-2].weight, mean=0.0, std=std, a=-2*std, b=2*std)
def forward(self, x): return self.mlp(x)
class Attention(nn.Module):
def __init__(self, config):
super().__init__()
if config.hidden_size % config.num_attention_heads != 0:
raise ValueError('hidden not divisible by heads')
self.nh = config.num_attention_heads
self.dh = config.hidden_size // config.num_attention_heads
self.qkv = nn.Linear(config.hidden_size, 3*config.hidden_size, bias=False)
self.o = nn.Linear(config.hidden_size, config.hidden_size, bias=False)
self.drop = nn.Dropout(config.hidden_dropout_prob)
def forward(self, x, attn_mask, rel):
B,S,H = x.shape
qkv = self.qkv(x).view(B,S,3,self.nh,self.dh).permute(2,0,3,1,4)
q,k,v = qkv[0],qkv[1],qkv[2]
attn = (q @ k.transpose(-1,-2)) / math.sqrt(self.dh)
if attn_mask is not None: attn = attn.masked_fill(attn_mask[:,None,:, :]==0, float('-inf'))
attn = torch.softmax(attn, dim=-1)
attn = self.drop(attn)
y = attn @ v
y = y.transpose(1,2).contiguous().view(B,S,H)
return self.o(y)
class Block(nn.Module):
def __init__(self, config):
super().__init__()
self.attn = Attention(config)
self.ff = FeedForward(config)
def forward(self, x, attn_mask, rel):
x = x + self.attn(x, attn_mask, rel)
x = x + self.ff(x)
return x
class Encoder(nn.Module):
def __init__(self, config):
super().__init__()
self.layers = nn.ModuleList([Block(config) for _ in range(config.num_hidden_layers)])
def forward(self, x, attn_mask, rel):
for layer in self.layers:
x = layer(x, attn_mask, rel)
return x
class Embedding(nn.Module):
def __init__(self, config):
super().__init__()
self.word_embedding = nn.Embedding(config.vocab_size, config.hidden_size)
self.pos_embedding = nn.Embedding(config.max_position_embeddings, config.hidden_size)
self.dropout = nn.Dropout(config.hidden_dropout_prob)
def forward(self, input_ids):
B,S = input_ids.shape
pos = torch.arange(0,S, device=input_ids.device).unsqueeze(0).expand(B,S)
x = self.word_embedding(input_ids) + self.pos_embedding(pos)
return self.dropout(x), None
class CoreModel(nn.Module):
def __init__(self, config):
super().__init__()
self.embedding = Embedding(config)
self.transformer = Encoder(config)
self.layer_norm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps, elementwise_affine=False)
self.head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
self.head.weight = self.embedding.word_embedding.weight
def forward(self, input_ids, attention_mask=None):
x,_ = self.embedding(input_ids)
if attention_mask is not None: attn = attention_mask.unsqueeze(1)
else: attn = None
x = self.transformer(x, attn, None)
x = self.layer_norm(x)
return self.head(x)
class GPTBertForMaskedLM(PreTrainedModel):
config_class = GPTBertConfig
base_model_prefix = 'gpt_bert'
def __init__(self, config: GPTBertConfig):
super().__init__(config)
self.model = CoreModel(config)
def forward(self, input_ids, attention_mask=None, labels=None):
logits = self.model(input_ids, attention_mask)
loss=None
if labels is not None:
loss_fct = nn.CrossEntropyLoss(ignore_index=-100)
loss = loss_fct(logits.view(-1, logits.size(-1)), labels.view(-1))
return MaskedLMOutput(loss=loss, logits=logits)
class GPTBertForCausalLM(PreTrainedModel):
config_class = GPTBertConfig
base_model_prefix = 'gpt_bert'
def __init__(self, config: GPTBertConfig):
super().__init__(config)
self.model = CoreModel(config)
def prepare_inputs_for_generation(self, input_ids, **kwargs):
return {'input_ids': input_ids}
def forward(self, input_ids, attention_mask=None, labels=None):
logits = self.model(input_ids, attention_mask)
loss=None
if labels is not None:
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
loss_fct = nn.CrossEntropyLoss(ignore_index=-100)
loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))
return CausalLMOutputWithCrossAttentions(loss=loss, logits=logits)