ab_ext_binary16 / modeling_attn_ext.py
E6E831728's picture
Upload folder using huggingface_hub
2791ac5 verified
Raw
History Blame Contribute Delete
18.5 kB
import math
from typing import Optional
import torch
import torch.utils.checkpoint
import torch.nn as nn
import torch.nn.functional as F
from transformers import PreTrainedModel
from transformers.generation import GenerationMixin
from transformers.modeling_outputs import CausalLMOutputWithPast
from .configuration_attn_ext import AttnExtConfig
def round_up(value: int, multiple: int) -> int:
return multiple * math.ceil(value / multiple)
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x):
dtype = x.dtype
xf = x.float()
xf = xf * torch.rsqrt(
xf.pow(2).mean(dim=-1, keepdim=True) + self.eps
)
return (xf * self.weight.float()).to(dtype)
def rotate_half(x):
x1 = x[..., ::2]
x2 = x[..., 1::2]
return torch.stack((-x2, x1), dim=-1).flatten(-2)
class RotaryEmbedding(nn.Module):
def __init__(self, dim, max_position, theta):
super().__init__()
inv_freq = 1.0 / (
theta
** (
torch.arange(0, dim, 2, dtype=torch.float32)
/ dim
)
)
positions = torch.arange(
max_position,
dtype=torch.float32,
)
frequencies = torch.outer(positions, inv_freq)
embedding = torch.repeat_interleave(
frequencies,
repeats=2,
dim=-1,
)
self.register_buffer(
"cos_cached",
embedding.cos(),
persistent=False,
)
self.register_buffer(
"sin_cached",
embedding.sin(),
persistent=False,
)
def forward(self, q, k, position_ids=None):
sequence_length = q.shape[-2]
if position_ids is None:
cos = self.cos_cached[:sequence_length][
None, None, :, :
]
sin = self.sin_cached[:sequence_length][
None, None, :, :
]
else:
cos = self.cos_cached[position_ids][:, None, :, :]
sin = self.sin_cached[position_ids][:, None, :, :]
cos = cos.to(device=q.device, dtype=q.dtype)
sin = sin.to(device=q.device, dtype=q.dtype)
q = q * cos + rotate_half(q) * sin
k = k * cos + rotate_half(k) * sin
return q, k
class CausalSelfAttention(nn.Module):
def __init__(self, config):
super().__init__()
self.d_model = config.d_model
self.n_head = config.n_head
self.head_dim = config.head_dim
self.dropout_p = config.dropout
self.q_proj = nn.Linear(
config.d_model,
config.d_model,
bias=config.attention_bias,
)
self.k_proj = nn.Linear(
config.d_model,
config.d_model,
bias=config.attention_bias,
)
self.v_proj = nn.Linear(
config.d_model,
config.d_model,
bias=config.attention_bias,
)
self.o_proj = nn.Linear(
config.d_model,
config.d_model,
bias=config.attention_bias,
)
self.rope = RotaryEmbedding(
config.head_dim,
config.block_size,
config.rope_theta,
)
def forward(
self,
x,
attention_mask=None,
position_ids=None,
):
batch_size, sequence_length, channels = x.shape
q = self.q_proj(x).view(
batch_size,
sequence_length,
self.n_head,
self.head_dim,
).transpose(1, 2)
k = self.k_proj(x).view(
batch_size,
sequence_length,
self.n_head,
self.head_dim,
).transpose(1, 2)
v = self.v_proj(x).view(
batch_size,
sequence_length,
self.n_head,
self.head_dim,
).transpose(1, 2)
q, k = self.rope(
q,
k,
position_ids=position_ids,
)
dropout_p = self.dropout_p if self.training else 0.0
if attention_mask is None or bool(attention_mask.all()):
output = F.scaled_dot_product_attention(
q,
k,
v,
attn_mask=None,
dropout_p=dropout_p,
is_causal=True,
)
else:
if attention_mask.shape != (
batch_size,
sequence_length,
):
raise ValueError(
"attention_mask must have shape "
f"{(batch_size, sequence_length)}"
)
causal = torch.ones(
sequence_length,
sequence_length,
device=x.device,
dtype=torch.bool,
).tril()
allowed = (
causal[None, None, :, :]
& attention_mask[:, None, None, :].bool()
)
output = F.scaled_dot_product_attention(
q,
k,
v,
attn_mask=allowed,
dropout_p=dropout_p,
is_causal=False,
)
output = output.transpose(1, 2).contiguous().view(
batch_size,
sequence_length,
channels,
)
return self.o_proj(output)
class SwiGLU(nn.Module):
def __init__(self, config):
super().__init__()
hidden_dim = round_up(
int(config.ffn_multiplier * config.d_model),
config.multiple_of,
)
self.gate_proj = nn.Linear(
config.d_model,
hidden_dim,
bias=config.mlp_bias,
)
self.up_proj = nn.Linear(
config.d_model,
hidden_dim,
bias=config.mlp_bias,
)
self.down_proj = nn.Linear(
hidden_dim,
config.d_model,
bias=config.mlp_bias,
)
self.dropout = nn.Dropout(config.dropout)
def forward(self, x):
x = F.silu(self.gate_proj(x)) * self.up_proj(x)
return self.dropout(self.down_proj(x))
class TransformerBlock(nn.Module):
def __init__(self, config):
super().__init__()
self.input_norm = RMSNorm(
config.d_model,
config.rms_norm_eps,
)
self.post_attention_norm = RMSNorm(
config.d_model,
config.rms_norm_eps,
)
self.attention = CausalSelfAttention(config)
self.mlp = SwiGLU(config)
def forward(
self,
x,
attention_mask=None,
position_ids=None,
):
x = x + self.attention(
self.input_norm(x),
attention_mask=attention_mask,
position_ids=position_ids,
)
x = x + self.mlp(
self.post_attention_norm(x)
)
return x
def canonical_binary_codebook(
vocab_size,
bits,
encoding,
):
token_ids = torch.arange(
vocab_size,
dtype=torch.int64,
)
shifts = torch.arange(
bits,
dtype=torch.int64,
)
codebook = (
(token_ids[:, None] >> shifts[None, :]) & 1
).to(torch.float32)
if encoding == "bipolar":
codebook = codebook.mul(2.0).sub(1.0)
return codebook.contiguous()
def gf2_rank(matrix):
matrix = matrix.detach().cpu().to(
torch.uint8
).clone()
matrix &= 1
rows, columns = matrix.shape
rank = 0
for column in range(columns):
pivot = None
for row in range(rank, rows):
if int(matrix[row, column]) == 1:
pivot = row
break
if pivot is None:
continue
if pivot != rank:
temporary = matrix[rank].clone()
matrix[rank] = matrix[pivot]
matrix[pivot] = temporary
for row in range(rows):
if row != rank and int(
matrix[row, column]
) == 1:
matrix[row] ^= matrix[rank]
rank += 1
if rank == rows:
break
return rank
def make_invertible_gf2_matrix(
bits,
seed,
min_row_weight,
min_col_weight,
):
generator = torch.Generator(device="cpu")
generator.manual_seed(seed)
for _ in range(1_000_000):
matrix = torch.randint(
0,
2,
(bits, bits),
generator=generator,
dtype=torch.uint8,
)
if bool(
torch.any(
matrix.sum(dim=1) < min_row_weight
)
):
continue
if bool(
torch.any(
matrix.sum(dim=0) < min_col_weight
)
):
continue
if gf2_rank(matrix) == bits:
return matrix.contiguous()
raise RuntimeError(
"Could not construct an invertible GF(2) matrix"
)
def gf2_binary_codebook(config):
source = canonical_binary_codebook(
config.vocab_size,
config.binary_dim,
"zero_one",
).to(torch.uint8)
matrix = make_invertible_gf2_matrix(
bits=config.binary_dim,
seed=config.code_seed,
min_row_weight=config.min_row_weight,
min_col_weight=config.min_col_weight,
)
shift = torch.zeros(
config.binary_dim,
dtype=torch.uint8,
)
codebook = (
source.to(torch.int16)
@ matrix.to(torch.int16).T
).remainder(2).to(torch.uint8)
codebook = codebook ^ shift
if config.binary_encoding == "bipolar":
codebook = (
codebook.float().mul(2.0).sub(1.0)
)
else:
codebook = codebook.float()
return (
codebook.contiguous(),
matrix.contiguous(),
shift.contiguous(),
)
class FixedBinaryEmbedding(nn.Module):
def __init__(self, config):
super().__init__()
if config.input_mode == "binary16":
codebook = canonical_binary_codebook(
config.vocab_size,
config.binary_dim,
config.binary_encoding,
)
matrix = None
shift = None
elif config.input_mode == "gf2":
codebook, matrix, shift = (
gf2_binary_codebook(config)
)
else:
raise ValueError(
"FixedBinaryEmbedding requires a "
"frozen-code input mode"
)
self.register_buffer(
"codebook",
codebook,
persistent=True,
)
if matrix is not None:
self.register_buffer(
"A_gf2",
matrix,
persistent=True,
)
self.register_buffer(
"b_gf2",
shift,
persistent=True,
)
self.repeat = config.binary_repeat
self.binary_scale = config.binary_scale
@property
def weight(self):
return self.codebook
def forward(self, input_ids):
code = self.codebook[input_ids.long()]
output = code.repeat(
*([1] * (code.ndim - 1)),
self.repeat,
)
if self.binary_scale != 1.0:
output = output * self.binary_scale
return output
class AttnExtPreTrainedModel(PreTrainedModel):
config_class = AttnExtConfig
base_model_prefix = "attn_ext"
supports_gradient_checkpointing = True
_supports_sdpa = True
_no_split_modules = ["TransformerBlock"]
def _init_weights(self, module):
if isinstance(module, nn.Linear):
nn.init.normal_(
module.weight,
mean=0.0,
std=self.config.initializer_range,
)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
nn.init.normal_(
module.weight,
mean=0.0,
std=self.config.initializer_range,
)
class AttnExtForCausalLM(
AttnExtPreTrainedModel,
GenerationMixin,
):
main_input_name = "input_ids"
def __init__(self, config):
super().__init__(config)
if config.input_mode == "learned":
self.token_embeddings = nn.Embedding(
config.vocab_size,
config.d_model,
)
else:
self.token_embeddings = FixedBinaryEmbedding(
config
)
self.layers = nn.ModuleList(
[
TransformerBlock(config)
for _ in range(config.n_layer)
]
)
self.final_norm = RMSNorm(
config.d_model,
config.rms_norm_eps,
)
self.lm_head = nn.Linear(
config.d_model,
config.vocab_size,
bias=False,
)
self.gradient_checkpointing = False
self.post_init()
residual_std = (
config.initializer_range
/ math.sqrt(2 * config.n_layer)
)
for layer in self.layers:
nn.init.normal_(
layer.attention.o_proj.weight,
mean=0.0,
std=residual_std,
)
nn.init.normal_(
layer.mlp.down_proj.weight,
mean=0.0,
std=residual_std,
)
def get_input_embeddings(self):
return self.token_embeddings
def set_input_embeddings(self, value):
if self.config.input_mode != "learned":
raise RuntimeError(
"Frozen input codes cannot be replaced "
"through set_input_embeddings"
)
self.token_embeddings = value
def get_output_embeddings(self):
return self.lm_head
def set_output_embeddings(self, value):
self.lm_head = value
def prepare_inputs_for_generation(
self,
input_ids,
attention_mask=None,
**kwargs,
):
if input_ids.shape[1] > self.config.block_size:
input_ids = input_ids[
:, -self.config.block_size:
]
if attention_mask is not None:
attention_mask = attention_mask[
:, -self.config.block_size:
]
position_ids = None
if attention_mask is not None:
position_ids = (
attention_mask.long().cumsum(-1) - 1
)
position_ids.masked_fill_(
attention_mask == 0,
0,
)
return {
"input_ids": input_ids,
"attention_mask": attention_mask,
"position_ids": position_ids,
"use_cache": False,
}
def forward(
self,
input_ids=None,
attention_mask=None,
labels=None,
position_ids=None,
inputs_embeds=None,
use_cache=None,
return_dict=None,
**kwargs,
):
if input_ids is None and inputs_embeds is None:
raise ValueError(
"input_ids or inputs_embeds is required"
)
if inputs_embeds is not None:
x = inputs_embeds
batch_size, sequence_length, _ = x.shape
else:
batch_size, sequence_length = input_ids.shape
x = self.token_embeddings(input_ids)
if sequence_length > self.config.block_size:
raise ValueError(
f"Sequence length {sequence_length} exceeds "
f"block_size={self.config.block_size}"
)
if attention_mask is not None:
expected = (batch_size, sequence_length)
if attention_mask.shape != expected:
raise ValueError(
f"attention_mask must have shape {expected}"
)
# HF_EXPORT_INPUT_DTYPE_FIX
# Frozen floating-point buffers may remain FP32 after loading.
# Match the residual stream to the backbone parameter dtype.
x = x.to(dtype=self.layers[0].attention.q_proj.weight.dtype)
for layer in self.layers:
if self.gradient_checkpointing and self.training:
def custom_forward(hidden_states, current_layer=layer):
return current_layer(
hidden_states,
attention_mask=attention_mask,
position_ids=position_ids,
)
x = torch.utils.checkpoint.checkpoint(
custom_forward,
x,
use_reentrant=False,
)
else:
x = layer(
x,
attention_mask=attention_mask,
position_ids=position_ids,
)
x = self.final_norm(x)
logits = self.lm_head(x)
loss = None
if labels is not None:
if labels.shape != (
batch_size,
sequence_length,
):
raise ValueError(
"labels must have the same shape as input_ids"
)
shift_logits = logits[:, :-1, :].contiguous()
shift_labels = labels[:, 1:].contiguous().clone()
if attention_mask is not None:
shift_labels.masked_fill_(
attention_mask[:, 1:].eq(0),
-100,
)
loss = F.cross_entropy(
shift_logits.float().view(
-1,
self.config.vocab_size,
),
shift_labels.view(-1),
ignore_index=-100,
)
return_dict = (
self.config.use_return_dict
if return_dict is None
else return_dict
)
if not return_dict:
output = (logits,)
return ((loss,) + output) if loss is not None else output
return CausalLMOutputWithPast(
loss=loss,
logits=logits,
past_key_values=None,
)