BAM-B2 / modeling_gpt_bert.py
Recor2d's picture
Upload BAM B2 from-scratch final model
b1a7488 verified
Raw
History Blame Contribute Delete
17.6 kB
import math
from dataclasses import dataclass
from typing import Optional, Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import PreTrainedModel
from transformers.modeling_outputs import BaseModelOutput, CausalLMOutput, MaskedLMOutput
from transformers.utils import ModelOutput
try:
from .configuration_gpt_bert import GPTBertConfig
except ImportError:
from configuration_gpt_bert import GPTBertConfig
@dataclass
class GPTBertTrainingOutput(ModelOutput):
loss: Optional[torch.Tensor] = None
logits: Optional[torch.Tensor] = None
ce_loss: Optional[torch.Tensor] = None
z_loss: Optional[torch.Tensor] = None
accuracy: Optional[torch.Tensor] = None
num_tokens: Optional[torch.Tensor] = None
class GeGLU(nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor:
value, gate = x.chunk(2, dim=-1)
return value * F.gelu(gate, approximate="tanh")
def _relative_position_buckets(
relative_position: torch.Tensor,
bucket_size: int,
max_position: int,
) -> torch.Tensor:
sign = torch.sign(relative_position)
mid = bucket_size // 2
abs_pos = torch.where(
(relative_position < mid) & (relative_position > -mid),
torch.full_like(relative_position, mid - 1),
torch.abs(relative_position).clamp(max=max_position - 1),
)
safe = abs_pos.clamp(min=mid)
log_pos = (
torch.ceil(
torch.log(safe.float() / mid)
/ math.log((max_position - 1) / mid)
* (mid - 1)
).long()
+ mid
)
bucket_pos = torch.where(abs_pos <= mid, relative_position, log_pos * sign)
return bucket_size - 1 + bucket_pos.long()
class GPTBertEmbeddings(nn.Module):
def __init__(self, config: GPTBertConfig):
super().__init__()
self.hidden_size = config.hidden_size
self.word_embeddings = nn.Embedding(
config.vocab_size,
config.hidden_size,
padding_idx=config.pad_token_id,
)
self.word_norm = nn.LayerNorm(
config.hidden_size,
eps=config.layer_norm_eps,
elementwise_affine=False,
)
self.relative_embeddings = nn.Parameter(
torch.empty(2 * config.position_bucket_size - 1, config.hidden_size)
)
self.relative_norm = nn.LayerNorm(
config.hidden_size,
eps=config.layer_norm_eps,
)
self.dropout = nn.Dropout(config.hidden_dropout_prob)
self.reset_parameters()
def reset_parameters(self):
std = math.sqrt(2.0 / (5.0 * self.hidden_size))
nn.init.trunc_normal_(
self.word_embeddings.weight, mean=0.0, std=std, a=-2 * std, b=2 * std
)
nn.init.trunc_normal_(
self.relative_embeddings, mean=0.0, std=std, a=-2 * std, b=2 * std
)
if self.word_embeddings.padding_idx is not None:
with torch.no_grad():
self.word_embeddings.weight[self.word_embeddings.padding_idx].zero_()
def forward(self, input_ids: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
x = self.dropout(self.word_norm(self.word_embeddings(input_ids)))
rel = self.relative_norm(self.relative_embeddings)
return x, rel
class GPTBertAttention(nn.Module):
def __init__(self, config: GPTBertConfig):
super().__init__()
if config.hidden_size % config.num_attention_heads != 0:
raise ValueError("hidden_size must be divisible by num_attention_heads")
self.config = config
self.hidden_size = config.hidden_size
self.num_heads = config.num_attention_heads
self.head_dim = config.hidden_size // config.num_attention_heads
self.qk_proj = nn.Linear(config.hidden_size, 2 * config.hidden_size)
self.vg_proj = nn.Linear(config.hidden_size, 2 * config.hidden_size)
self.out_proj = nn.Linear(config.hidden_size, config.hidden_size)
self.pre_norm = nn.LayerNorm(
config.hidden_size,
eps=config.layer_norm_eps,
elementwise_affine=False,
)
self.post_norm = nn.LayerNorm(
config.hidden_size,
eps=config.layer_norm_eps,
elementwise_affine=False,
)
self.dropout = nn.Dropout(config.attention_probs_dropout_prob)
self.out_dropout = nn.Dropout(config.hidden_dropout_prob)
self.scale = 1.0 / math.sqrt(3.0 * self.head_dim)
positions = (
torch.arange(config.max_position_embeddings).unsqueeze(1)
- torch.arange(config.max_position_embeddings).unsqueeze(0)
)
buckets = _relative_position_buckets(
positions,
config.position_bucket_size,
config.max_position_embeddings,
)
self.register_buffer("position_indices", buckets, persistent=False)
self.reset_parameters()
def reset_parameters(self):
std = math.sqrt(2.0 / (5.0 * self.hidden_size))
for layer in (self.qk_proj, self.vg_proj, self.out_proj):
nn.init.trunc_normal_(
layer.weight, mean=0.0, std=std, a=-2 * std, b=2 * std
)
if layer.bias is not None:
nn.init.zeros_(layer.bias)
def forward(
self,
hidden_states: torch.Tensor,
blocked_mask: torch.Tensor,
relative_embeddings: torch.Tensor,
) -> torch.Tensor:
batch_size, seq_len, _ = hidden_states.shape
x = self.pre_norm(hidden_states)
query, key = self.qk_proj(x).chunk(2, dim=-1)
value, gate = self.vg_proj(x).chunk(2, dim=-1)
gate = F.gelu(gate)
query = query.view(batch_size, seq_len, self.num_heads, self.head_dim)
key = key.view(batch_size, seq_len, self.num_heads, self.head_dim)
value = value.view(batch_size, seq_len, self.num_heads, self.head_dim)
query = query.permute(0, 2, 1, 3)
key = key.permute(0, 2, 1, 3)
value = value.permute(0, 2, 1, 3)
scores = torch.matmul(query, key.transpose(-1, -2)) * self.scale
rel_qk = self.qk_proj(self.dropout(relative_embeddings))
rel_q, rel_k = rel_qk.chunk(2, dim=-1)
indices = self.position_indices[:seq_len, :seq_len]
rel_q = F.embedding(indices, rel_q).view(
seq_len, seq_len, self.num_heads, self.head_dim
)
rel_k = F.embedding(indices, rel_k).view(
seq_len, seq_len, self.num_heads, self.head_dim
)
scores = scores + torch.einsum(
"bhqd,qkhd->bhqk", query, rel_k * self.scale
)
scores = scores + torch.einsum(
"bhkd,qkhd->bhqk", key * self.scale, rel_q
)
scores = scores.masked_fill(blocked_mask, torch.finfo(scores.dtype).min)
probs = torch.softmax(scores.float(), dim=-1).to(scores.dtype)
probs = self.dropout(probs)
context = torch.matmul(probs, value)
context = context.permute(0, 2, 1, 3).contiguous().view(
batch_size, seq_len, self.hidden_size
)
context = context * gate
context = self.post_norm(context)
context = self.out_proj(context)
return self.out_dropout(context)
class GPTBertFeedForward(nn.Module):
def __init__(self, config: GPTBertConfig, layer_index: int):
super().__init__()
self.norm1 = nn.LayerNorm(
config.hidden_size,
eps=config.layer_norm_eps,
elementwise_affine=False,
)
self.fc1 = nn.Linear(
config.hidden_size,
2 * config.intermediate_size,
bias=False,
)
self.act = GeGLU()
self.norm2 = nn.LayerNorm(
config.intermediate_size,
eps=config.layer_norm_eps,
elementwise_affine=False,
)
self.fc2 = nn.Linear(
config.intermediate_size,
config.hidden_size,
bias=False,
)
self.dropout = nn.Dropout(config.hidden_dropout_prob)
self.reset_parameters(layer_index)
def reset_parameters(self, layer_index: int):
std = math.sqrt(2.0 / (5.0 * self.fc2.out_features))
nn.init.trunc_normal_(self.fc1.weight, mean=0.0, std=std, a=-2 * std, b=2 * std)
nn.init.trunc_normal_(self.fc2.weight, mean=0.0, std=std, a=-2 * std, b=2 * std)
scale = math.sqrt(1.0 / (2.0 * (1 + layer_index)))
with torch.no_grad():
self.fc1.weight.mul_(scale)
self.fc2.weight.mul_(scale)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
x = self.norm1(hidden_states)
x = self.fc1(x)
x = self.act(x)
x = self.norm2(x)
x = self.fc2(x)
return self.dropout(x)
class GPTBertLayer(nn.Module):
def __init__(self, config: GPTBertConfig, layer_index: int):
super().__init__()
self.attention = GPTBertAttention(config)
self.ffn = GPTBertFeedForward(config, layer_index)
def forward(
self,
hidden_states: torch.Tensor,
blocked_mask: torch.Tensor,
relative_embeddings: torch.Tensor,
) -> torch.Tensor:
hidden_states = hidden_states + self.attention(
hidden_states, blocked_mask, relative_embeddings
)
hidden_states = hidden_states + self.ffn(hidden_states)
return hidden_states
class GPTBertPreTrainedModel(PreTrainedModel):
config_class = GPTBertConfig
base_model_prefix = "gpt_bert"
supports_gradient_checkpointing = False
def _init_weights(self, module):
# Components initialize themselves to match the LTG/GPT-BERT recipe.
return
class GPTBertModel(GPTBertPreTrainedModel):
def __init__(self, config: GPTBertConfig):
super().__init__(config)
self.embeddings = GPTBertEmbeddings(config)
self.layers = nn.ModuleList(
[GPTBertLayer(config, i) for i in range(config.num_hidden_layers)]
)
self.post_init()
def get_input_embeddings(self):
return self.embeddings.word_embeddings
def set_input_embeddings(self, value):
self.embeddings.word_embeddings = value
def _build_blocked_mask(
self,
input_ids: torch.Tensor,
attention_mask: Optional[torch.Tensor],
is_causal: bool,
) -> torch.Tensor:
batch_size, seq_len = input_ids.shape
if attention_mask is None:
attention_mask = torch.ones(
batch_size, seq_len, device=input_ids.device, dtype=torch.long
)
key_padding = attention_mask.eq(0)[:, None, None, :]
if is_causal:
causal = torch.ones(
seq_len, seq_len, device=input_ids.device, dtype=torch.bool
).triu(diagonal=1)[None, None, :, :]
return key_padding | causal
return key_padding.expand(batch_size, 1, seq_len, seq_len)
def forward(
self,
input_ids: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
is_causal: bool = False,
return_dict: bool = True,
**kwargs,
):
blocked_mask = self._build_blocked_mask(
input_ids, attention_mask, is_causal=is_causal
)
hidden_states, relative_embeddings = self.embeddings(input_ids)
for layer in self.layers:
hidden_states = layer(
hidden_states, blocked_mask, relative_embeddings
)
if not return_dict:
return (hidden_states,)
return BaseModelOutput(last_hidden_state=hidden_states)
class GPTBertLMHead(nn.Module):
def __init__(self, config: GPTBertConfig, embedding_weight: nn.Parameter):
super().__init__()
self.dense = nn.Linear(config.hidden_size, config.hidden_size)
self.activation = nn.GELU()
self.norm = nn.LayerNorm(
config.hidden_size,
eps=config.layer_norm_eps,
elementwise_affine=False,
)
self.decoder = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
self.decoder.weight = embedding_weight
self.bias = nn.Parameter(torch.zeros(config.vocab_size))
self.reset_parameters(config.hidden_size)
def reset_parameters(self, hidden_size: int):
std = math.sqrt(2.0 / (5.0 * hidden_size))
nn.init.trunc_normal_(
self.dense.weight, mean=0.0, std=std, a=-2 * std, b=2 * std
)
nn.init.zeros_(self.dense.bias)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
x = self.dense(hidden_states)
x = self.activation(x)
x = self.norm(x)
return self.decoder(x) + self.bias
class GPTBertForMaskedLM(GPTBertPreTrainedModel):
_tied_weights_keys = ["lm_head.decoder.weight"]
def __init__(self, config: GPTBertConfig):
super().__init__(config)
self.gpt_bert = GPTBertModel(config)
self.lm_head = GPTBertLMHead(
config, self.gpt_bert.embeddings.word_embeddings.weight
)
self.post_init()
def get_input_embeddings(self):
return self.gpt_bert.get_input_embeddings()
def set_input_embeddings(self, value):
self.gpt_bert.set_input_embeddings(value)
self.lm_head.decoder.weight = value.weight
def get_output_embeddings(self):
return self.lm_head.decoder
def set_output_embeddings(self, new_embeddings):
self.lm_head.decoder = new_embeddings
def _selected_stats(
self,
logits: torch.Tensor,
labels: torch.Tensor,
):
ce_loss = F.cross_entropy(logits.float(), labels)
z_loss = torch.logsumexp(logits.float(), dim=-1).pow(2).mean()
accuracy = (logits.argmax(dim=-1) == labels).float().mean()
return ce_loss, z_loss, accuracy
def forward(
self,
input_ids: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
labels: Optional[torch.Tensor] = None,
mode: str = "mntp",
z_loss_weight: float = 0.0,
return_dict: bool = True,
**kwargs,
):
if mode not in {"mntp", "masked", "causal"}:
raise ValueError(f"Unsupported mode: {mode}")
is_causal = mode == "causal"
hidden = self.gpt_bert(
input_ids=input_ids,
attention_mask=attention_mask,
is_causal=is_causal,
return_dict=True,
).last_hidden_state
if labels is not None and mode in {"mntp", "masked"}:
# MNTP: the hidden state at position i-1 predicts a masked token at i.
valid = labels[:, 1:].ne(-100)
selected_hidden = hidden[:, :-1][valid]
selected_labels = labels[:, 1:][valid]
if selected_labels.numel() == 0:
raise RuntimeError("MNTP batch contains no prediction targets")
selected_logits = self.lm_head(selected_hidden)
ce_loss, z_loss, accuracy = self._selected_stats(
selected_logits, selected_labels
)
loss = ce_loss + z_loss_weight * z_loss
return GPTBertTrainingOutput(
loss=loss,
logits=None,
ce_loss=ce_loss.detach(),
z_loss=z_loss.detach(),
accuracy=accuracy.detach(),
num_tokens=torch.tensor(
selected_labels.numel(), device=input_ids.device
),
)
if labels is not None and mode == "causal":
# Training input is [BOS] + tokens[:-1]; labels are tokens.
logits = self.lm_head(hidden)
valid = labels.ne(-100)
selected_logits = logits[valid]
selected_labels = labels[valid]
ce_loss, z_loss, accuracy = self._selected_stats(
selected_logits, selected_labels
)
loss = ce_loss + z_loss_weight * z_loss
return GPTBertTrainingOutput(
loss=loss,
logits=None,
ce_loss=ce_loss.detach(),
z_loss=z_loss.detach(),
accuracy=accuracy.detach(),
num_tokens=torch.tensor(
selected_labels.numel(), device=input_ids.device
),
)
raw_logits = self.lm_head(hidden)
# The BabyLM MNTP backend selects target_position - 1 itself.
if not return_dict:
return (raw_logits,)
return MaskedLMOutput(logits=raw_logits)
class GPTBertForCausalLM(GPTBertForMaskedLM):
def forward(
self,
input_ids: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
labels: Optional[torch.Tensor] = None,
return_dict: bool = True,
**kwargs,
):
hidden = self.gpt_bert(
input_ids=input_ids,
attention_mask=attention_mask,
is_causal=True,
return_dict=True,
).last_hidden_state
logits = self.lm_head(hidden)
loss = None
if labels is not None:
shift_logits = logits[:, :-1].contiguous()
shift_labels = labels[:, 1:].contiguous()
loss = F.cross_entropy(
shift_logits.view(-1, shift_logits.size(-1)).float(),
shift_labels.view(-1),
ignore_index=-100,
)
if not return_dict:
return (loss, logits) if loss is not None else (logits,)
return CausalLMOutput(loss=loss, logits=logits)