E6E831728/ab_ext_binary16
0469
1import math2from typing import Optional3 4import torch5import torch.utils.checkpoint6import torch.nn as nn7import torch.nn.functional as F8 9from transformers import PreTrainedModel10from transformers.generation import GenerationMixin11from transformers.modeling_outputs import CausalLMOutputWithPast12 13from .configuration_attn_ext import AttnExtConfig14 15 16def round_up(value: int, multiple: int) -> int:17 return multiple * math.ceil(value / multiple)18 19 20class RMSNorm(nn.Module):21 def __init__(self, dim: int, eps: float):22 super().__init__()23 self.weight = nn.Parameter(torch.ones(dim))24 self.eps = eps25 26 def forward(self, x):27 dtype = x.dtype28 xf = x.float()29 xf = xf * torch.rsqrt(30 xf.pow(2).mean(dim=-1, keepdim=True) + self.eps31 )32 return (xf * self.weight.float()).to(dtype)33 34 35def rotate_half(x):36 x1 = x[..., ::2]37 x2 = x[..., 1::2]38 return torch.stack((-x2, x1), dim=-1).flatten(-2)39 40 41class RotaryEmbedding(nn.Module):42 def __init__(self, dim, max_position, theta):43 super().__init__()44 45 inv_freq = 1.0 / (46 theta47 ** (48 torch.arange(0, dim, 2, dtype=torch.float32)49 / dim50 )51 )52 53 positions = torch.arange(54 max_position,55 dtype=torch.float32,56 )57 58 frequencies = torch.outer(positions, inv_freq)59 embedding = torch.repeat_interleave(60 frequencies,61 repeats=2,62 dim=-1,63 )64 65 self.register_buffer(66 "cos_cached",67 embedding.cos(),68 persistent=False,69 )70 self.register_buffer(71 "sin_cached",72 embedding.sin(),73 persistent=False,74 )75 76 def forward(self, q, k, position_ids=None):77 sequence_length = q.shape[-2]78 79 if position_ids is None:80 cos = self.cos_cached[:sequence_length][81 None, None, :, :82 ]83 sin = self.sin_cached[:sequence_length][84 None, None, :, :85 ]86 else:87 cos = self.cos_cached[position_ids][:, None, :, :]88 sin = self.sin_cached[position_ids][:, None, :, :]89 90 cos = cos.to(device=q.device, dtype=q.dtype)91 sin = sin.to(device=q.device, dtype=q.dtype)92 93 q = q * cos + rotate_half(q) * sin94 k = k * cos + rotate_half(k) * sin95 return q, k96 97 98class CausalSelfAttention(nn.Module):99 def __init__(self, config):100 super().__init__()101 102 self.d_model = config.d_model103 self.n_head = config.n_head104 self.head_dim = config.head_dim105 self.dropout_p = config.dropout106 107 self.q_proj = nn.Linear(108 config.d_model,109 config.d_model,110 bias=config.attention_bias,111 )112 self.k_proj = nn.Linear(113 config.d_model,114 config.d_model,115 bias=config.attention_bias,116 )117 self.v_proj = nn.Linear(118 config.d_model,119 config.d_model,120 bias=config.attention_bias,121 )122 self.o_proj = nn.Linear(123 config.d_model,124 config.d_model,125 bias=config.attention_bias,126 )127 128 self.rope = RotaryEmbedding(129 config.head_dim,130 config.block_size,131 config.rope_theta,132 )133 134 def forward(135 self,136 x,137 attention_mask=None,138 position_ids=None,139 ):140 batch_size, sequence_length, channels = x.shape141 142 q = self.q_proj(x).view(143 batch_size,144 sequence_length,145 self.n_head,146 self.head_dim,147 ).transpose(1, 2)148 149 k = self.k_proj(x).view(150 batch_size,151 sequence_length,152 self.n_head,153 self.head_dim,154 ).transpose(1, 2)155 156 v = self.v_proj(x).view(157 batch_size,158 sequence_length,159 self.n_head,160 self.head_dim,161 ).transpose(1, 2)162 163 q, k = self.rope(164 q,165 k,166 position_ids=position_ids,167 )168 169 dropout_p = self.dropout_p if self.training else 0.0170 171 if attention_mask is None or bool(attention_mask.all()):172 output = F.scaled_dot_product_attention(173 q,174 k,175 v,176 attn_mask=None,177 dropout_p=dropout_p,178 is_causal=True,179 )180 else:181 if attention_mask.shape != (182 batch_size,183 sequence_length,184 ):185 raise ValueError(186 "attention_mask must have shape "187 f"{(batch_size, sequence_length)}"188 )189 190 causal = torch.ones(191 sequence_length,192 sequence_length,193 device=x.device,194 dtype=torch.bool,195 ).tril()196 197 allowed = (198 causal[None, None, :, :]199 & attention_mask[:, None, None, :].bool()200 )201 202 output = F.scaled_dot_product_attention(203 q,204 k,205 v,206 attn_mask=allowed,207 dropout_p=dropout_p,208 is_causal=False,209 )210 211 output = output.transpose(1, 2).contiguous().view(212 batch_size,213 sequence_length,214 channels,215 )216 217 return self.o_proj(output)218 219 220class SwiGLU(nn.Module):221 def __init__(self, config):222 super().__init__()223 224 hidden_dim = round_up(225 int(config.ffn_multiplier * config.d_model),226 config.multiple_of,227 )228 229 self.gate_proj = nn.Linear(230 config.d_model,231 hidden_dim,232 bias=config.mlp_bias,233 )234 self.up_proj = nn.Linear(235 config.d_model,236 hidden_dim,237 bias=config.mlp_bias,238 )239 self.down_proj = nn.Linear(240 hidden_dim,241 config.d_model,242 bias=config.mlp_bias,243 )244 self.dropout = nn.Dropout(config.dropout)245 246 def forward(self, x):247 x = F.silu(self.gate_proj(x)) * self.up_proj(x)248 return self.dropout(self.down_proj(x))249 250 251class TransformerBlock(nn.Module):252 def __init__(self, config):253 super().__init__()254 255 self.input_norm = RMSNorm(256 config.d_model,257 config.rms_norm_eps,258 )259 self.post_attention_norm = RMSNorm(260 config.d_model,261 config.rms_norm_eps,262 )263 264 self.attention = CausalSelfAttention(config)265 self.mlp = SwiGLU(config)266 267 def forward(268 self,269 x,270 attention_mask=None,271 position_ids=None,272 ):273 x = x + self.attention(274 self.input_norm(x),275 attention_mask=attention_mask,276 position_ids=position_ids,277 )278 279 x = x + self.mlp(280 self.post_attention_norm(x)281 )282 283 return x284 285 286def canonical_binary_codebook(287 vocab_size,288 bits,289 encoding,290):291 token_ids = torch.arange(292 vocab_size,293 dtype=torch.int64,294 )295 shifts = torch.arange(296 bits,297 dtype=torch.int64,298 )299 300 codebook = (301 (token_ids[:, None] >> shifts[None, :]) & 1302 ).to(torch.float32)303 304 if encoding == "bipolar":305 codebook = codebook.mul(2.0).sub(1.0)306 307 return codebook.contiguous()308 309 310def gf2_rank(matrix):311 matrix = matrix.detach().cpu().to(312 torch.uint8313 ).clone()314 matrix &= 1315 316 rows, columns = matrix.shape317 rank = 0318 319 for column in range(columns):320 pivot = None321 322 for row in range(rank, rows):323 if int(matrix[row, column]) == 1:324 pivot = row325 break326 327 if pivot is None:328 continue329 330 if pivot != rank:331 temporary = matrix[rank].clone()332 matrix[rank] = matrix[pivot]333 matrix[pivot] = temporary334 335 for row in range(rows):336 if row != rank and int(337 matrix[row, column]338 ) == 1:339 matrix[row] ^= matrix[rank]340 341 rank += 1342 343 if rank == rows:344 break345 346 return rank347 348 349def make_invertible_gf2_matrix(350 bits,351 seed,352 min_row_weight,353 min_col_weight,354):355 generator = torch.Generator(device="cpu")356 generator.manual_seed(seed)357 358 for _ in range(1_000_000):359 matrix = torch.randint(360 0,361 2,362 (bits, bits),363 generator=generator,364 dtype=torch.uint8,365 )366 367 if bool(368 torch.any(369 matrix.sum(dim=1) < min_row_weight370 )371 ):372 continue373 374 if bool(375 torch.any(376 matrix.sum(dim=0) < min_col_weight377 )378 ):379 continue380 381 if gf2_rank(matrix) == bits:382 return matrix.contiguous()383 384 raise RuntimeError(385 "Could not construct an invertible GF(2) matrix"386 )387 388 389def gf2_binary_codebook(config):390 source = canonical_binary_codebook(391 config.vocab_size,392 config.binary_dim,393 "zero_one",394 ).to(torch.uint8)395 396 matrix = make_invertible_gf2_matrix(397 bits=config.binary_dim,398 seed=config.code_seed,399 min_row_weight=config.min_row_weight,400 min_col_weight=config.min_col_weight,401 )402 403 shift = torch.zeros(404 config.binary_dim,405 dtype=torch.uint8,406 )407 408 codebook = (409 source.to(torch.int16)410 @ matrix.to(torch.int16).T411 ).remainder(2).to(torch.uint8)412 413 codebook = codebook ^ shift414 415 if config.binary_encoding == "bipolar":416 codebook = (417 codebook.float().mul(2.0).sub(1.0)418 )419 else:420 codebook = codebook.float()421 422 return (423 codebook.contiguous(),424 matrix.contiguous(),425 shift.contiguous(),426 )427 428 429class FixedBinaryEmbedding(nn.Module):430 def __init__(self, config):431 super().__init__()432 433 if config.input_mode == "binary16":434 codebook = canonical_binary_codebook(435 config.vocab_size,436 config.binary_dim,437 config.binary_encoding,438 )439 matrix = None440 shift = None441 442 elif config.input_mode == "gf2":443 codebook, matrix, shift = (444 gf2_binary_codebook(config)445 )446 447 else:448 raise ValueError(449 "FixedBinaryEmbedding requires a "450 "frozen-code input mode"451 )452 453 self.register_buffer(454 "codebook",455 codebook,456 persistent=True,457 )458 459 if matrix is not None:460 self.register_buffer(461 "A_gf2",462 matrix,463 persistent=True,464 )465 self.register_buffer(466 "b_gf2",467 shift,468 persistent=True,469 )470 471 self.repeat = config.binary_repeat472 self.binary_scale = config.binary_scale473 474 @property475 def weight(self):476 return self.codebook477 478 def forward(self, input_ids):479 code = self.codebook[input_ids.long()]480 481 output = code.repeat(482 *([1] * (code.ndim - 1)),483 self.repeat,484 )485 486 if self.binary_scale != 1.0:487 output = output * self.binary_scale488 489 return output490 491 492class AttnExtPreTrainedModel(PreTrainedModel):493 config_class = AttnExtConfig494 base_model_prefix = "attn_ext"495 supports_gradient_checkpointing = True496 _supports_sdpa = True497 _no_split_modules = ["TransformerBlock"]498 499 def _init_weights(self, module):500 if isinstance(module, nn.Linear):501 nn.init.normal_(502 module.weight,503 mean=0.0,504 std=self.config.initializer_range,505 )506 if module.bias is not None:507 nn.init.zeros_(module.bias)508 509 elif isinstance(module, nn.Embedding):510 nn.init.normal_(511 module.weight,512 mean=0.0,513 std=self.config.initializer_range,514 )515 516 517class AttnExtForCausalLM(518 AttnExtPreTrainedModel,519 GenerationMixin,520):521 main_input_name = "input_ids"522 523 def __init__(self, config):524 super().__init__(config)525 526 if config.input_mode == "learned":527 self.token_embeddings = nn.Embedding(528 config.vocab_size,529 config.d_model,530 )531 else:532 self.token_embeddings = FixedBinaryEmbedding(533 config534 )535 536 self.layers = nn.ModuleList(537 [538 TransformerBlock(config)539 for _ in range(config.n_layer)540 ]541 )542 543 self.final_norm = RMSNorm(544 config.d_model,545 config.rms_norm_eps,546 )547 548 self.lm_head = nn.Linear(549 config.d_model,550 config.vocab_size,551 bias=False,552 )553 554 self.gradient_checkpointing = False555 self.post_init()556 557 residual_std = (558 config.initializer_range559 / math.sqrt(2 * config.n_layer)560 )561 562 for layer in self.layers:563 nn.init.normal_(564 layer.attention.o_proj.weight,565 mean=0.0,566 std=residual_std,567 )568 nn.init.normal_(569 layer.mlp.down_proj.weight,570 mean=0.0,571 std=residual_std,572 )573 574 def get_input_embeddings(self):575 return self.token_embeddings576 577 def set_input_embeddings(self, value):578 if self.config.input_mode != "learned":579 raise RuntimeError(580 "Frozen input codes cannot be replaced "581 "through set_input_embeddings"582 )583 self.token_embeddings = value584 585 def get_output_embeddings(self):586 return self.lm_head587 588 def set_output_embeddings(self, value):589 self.lm_head = value590 591 def prepare_inputs_for_generation(592 self,593 input_ids,594 attention_mask=None,595 **kwargs,596 ):597 if input_ids.shape[1] > self.config.block_size:598 input_ids = input_ids[599 :, -self.config.block_size:600 ]601 602 if attention_mask is not None:603 attention_mask = attention_mask[604 :, -self.config.block_size:605 ]606 607 position_ids = None608 609 if attention_mask is not None:610 position_ids = (611 attention_mask.long().cumsum(-1) - 1612 )613 position_ids.masked_fill_(614 attention_mask == 0,615 0,616 )617 618 return {619 "input_ids": input_ids,620 "attention_mask": attention_mask,621 "position_ids": position_ids,622 "use_cache": False,623 }624 625 def forward(626 self,627 input_ids=None,628 attention_mask=None,629 labels=None,630 position_ids=None,631 inputs_embeds=None,632 use_cache=None,633 return_dict=None,634 **kwargs,635 ):636 if input_ids is None and inputs_embeds is None:637 raise ValueError(638 "input_ids or inputs_embeds is required"639 )640 641 if inputs_embeds is not None:642 x = inputs_embeds643 batch_size, sequence_length, _ = x.shape644 else:645 batch_size, sequence_length = input_ids.shape646 x = self.token_embeddings(input_ids)647 648 if sequence_length > self.config.block_size:649 raise ValueError(650 f"Sequence length {sequence_length} exceeds "651 f"block_size={self.config.block_size}"652 )653 654 if attention_mask is not None:655 expected = (batch_size, sequence_length)656 if attention_mask.shape != expected:657 raise ValueError(658 f"attention_mask must have shape {expected}"659 )660 661 # HF_EXPORT_INPUT_DTYPE_FIX662 # Frozen floating-point buffers may remain FP32 after loading.663 # Match the residual stream to the backbone parameter dtype.664 x = x.to(dtype=self.layers[0].attention.q_proj.weight.dtype)665 666 for layer in self.layers:667 if self.gradient_checkpointing and self.training:668 def custom_forward(hidden_states, current_layer=layer):669 return current_layer(670 hidden_states,671 attention_mask=attention_mask,672 position_ids=position_ids,673 )674 675 x = torch.utils.checkpoint.checkpoint(676 custom_forward,677 x,678 use_reentrant=False,679 )680 else:681 x = layer(682 x,683 attention_mask=attention_mask,684 position_ids=position_ids,685 )686 687 x = self.final_norm(x)688 logits = self.lm_head(x)689 690 loss = None691 692 if labels is not None:693 if labels.shape != (694 batch_size,695 sequence_length,696 ):697 raise ValueError(698 "labels must have the same shape as input_ids"699 )700 701 shift_logits = logits[:, :-1, :].contiguous()702 shift_labels = labels[:, 1:].contiguous().clone()703 704 if attention_mask is not None:705 shift_labels.masked_fill_(706 attention_mask[:, 1:].eq(0),707 -100,708 )709 710 loss = F.cross_entropy(711 shift_logits.float().view(712 -1,713 self.config.vocab_size,714 ),715 shift_labels.view(-1),716 ignore_index=-100,717 )718 719 return_dict = (720 self.config.use_return_dict721 if return_dict is None722 else return_dict723 )724 725 if not return_dict:726 output = (logits,)727 return ((loss,) + output) if loss is not None else output728 729 return CausalLMOutputWithPast(730 loss=loss,731 logits=logits,732 past_key_values=None,733 )734 