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, )