Team Ai
Modelpublic

E6E831728/ab_ext_binary16

sourceHugging Faceupdated 15d agoView on Hugging Face
0likes469downloads
modeling_attn_ext.py734 linesDownload Raw Back to root
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