Team Ai
Apppublic

PascalNotin/Tranception_design

sourceHugging Facemitupdated 4y agoView on Hugging Face
4likes
model_pytorch.py931 linesDownload Raw Back to tranception
1from dataclasses import dataclass2from typing import Optional, Tuple3import math4import os5import pandas as pd6 7import torch8from torch import nn9from torch.nn import CrossEntropyLoss, NLLLoss10import torch.nn.functional as F11from transformers import GPT2PreTrainedModel12 13from transformers.modeling_utils import (14    Conv1D,15    PreTrainedModel,16    SequenceSummary,17    find_pruneable_heads_and_indices,18    prune_conv1d_layer,19)20from transformers.file_utils import (21    ModelOutput,22    add_code_sample_docstrings,23    add_start_docstrings,24    add_start_docstrings_to_model_forward,25    replace_return_docstrings26)27from transformers.modeling_outputs import (28    BaseModelOutputWithPastAndCrossAttentions,29    CausalLMOutputWithCrossAttentions,30    SequenceClassifierOutputWithPast,31    TokenClassifierOutput32)33from transformers.utils.model_parallel_utils import assert_device_map, get_device_map34 35from tranception.activations import tranception_ACT2FN36from tranception.config import TranceptionConfig37from tranception.outputs import (38    TranceptionCausalLMOutputWithCrossAttentions,39)40from tranception.utils import msa_utils41from tranception.utils import scoring_utils42 43def nanmean(v, *args, inplace=False, **kwargs):44    if not inplace:45        v = v.clone()46    is_nan = torch.isnan(v)47    v[is_nan] = 048    return v.sum(*args, **kwargs) / (~is_nan).float().sum(*args, **kwargs)49 50def get_slopes(n, mode="standard_alibi", verbose=False):51    """52    Function to compute the m constant for each attention head. Code has been adapted from the official ALiBi codebase at:53    https://github.com/ofirpress/attention_with_linear_biases/blob/master/fairseq/models/transformer.py54    """55    def get_slopes_power_of_2(n):56        start = (2**(-2**-(math.log2(n)-3)))57        ratio = start58        return [start*ratio**i for i in range(n)]59    if mode=="grouped_alibi":60        n = n // 461    if math.log2(n).is_integer():62        result = get_slopes_power_of_2(n)                   63    else:64        #Workaround when the number of heads is not a power of 265        closest_power_of_2 = 2**math.floor(math.log2(n))  66        result = get_slopes_power_of_2(closest_power_of_2) + get_slopes(2*closest_power_of_2)[0::2][:n-closest_power_of_2]67    if mode=="grouped_alibi":68        result = result * 469        if verbose:70            print("ALiBi slopes: {}".format(result))71    return result72 73class SpatialDepthWiseConvolution(nn.Module):74    def __init__(self, head_dim: int, kernel_size: int = 3):75        super().__init__()76        self.kernel_size = kernel_size77        self.conv = nn.Conv1d(in_channels=head_dim, out_channels=head_dim, kernel_size=(kernel_size,), padding=(kernel_size - 1,), groups=head_dim)78    79    def forward(self, x: torch.Tensor):80        batch_size, heads, seq_len, head_dim = x.shape81        x = x.permute(0, 1, 3, 2).contiguous()82        x = x.view(batch_size * heads, head_dim, seq_len)83        x = self.conv(x)84        if self.kernel_size>1:85            x = x[:, :, :-(self.kernel_size - 1)]86        x = x.view(batch_size, heads, head_dim, seq_len)87        x = x.permute(0, 1, 3, 2)88        return x89 90class TranceptionBlockAttention(nn.Module):91    def __init__(self, config, is_cross_attention=False, SDWC_kernel_size=None):92        super().__init__()93 94        max_positions = config.max_position_embeddings95        self.register_buffer(96            "bias",97            torch.tril(torch.ones((max_positions, max_positions), dtype=torch.uint8)).view(98                1, 1, max_positions, max_positions99            ),100        )101        self.register_buffer("masked_bias", torch.tensor(-1e4))102 103        self.embed_dim = config.hidden_size104        self.num_heads = config.num_attention_heads105        self.head_dim = self.embed_dim // self.num_heads106        self.split_size = self.embed_dim107        if self.head_dim * self.num_heads != self.embed_dim:108            raise ValueError(109                f"`embed_dim` must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`: {self.num_heads})."110            )111 112        self.scale_attn_weights = config.scale_attn_weights113        self.is_cross_attention = is_cross_attention114 115        if self.is_cross_attention:116            self.c_attn = Conv1D(2 * self.embed_dim, self.embed_dim)117            self.q_attn = Conv1D(self.embed_dim, self.embed_dim)118        else:119            self.c_attn = Conv1D(3 * self.embed_dim, self.embed_dim)120        self.c_proj = Conv1D(self.embed_dim, self.embed_dim)121 122        self.attn_dropout = nn.Dropout(config.attn_pdrop)123        self.resid_dropout = nn.Dropout(config.resid_pdrop)124 125        self.pruned_heads = set()126 127        self.attention_mode=config.attention_mode128        129        if self.attention_mode=="tranception":130            assert self.num_heads%4==0, "Invalid number of heads. Tranception requires the number of heads to be a multiple of 4."131            self.num_heads_per_kernel_size = self.num_heads // 4132            self.query_depthwiseconv = nn.ModuleDict()133            self.key_depthwiseconv = nn.ModuleDict()134            self.value_depthwiseconv = nn.ModuleDict()135            for kernel_idx, kernel in enumerate([3,5,7]):136                self.query_depthwiseconv[str(kernel_idx)] = SpatialDepthWiseConvolution(self.head_dim,kernel)137                self.key_depthwiseconv[str(kernel_idx)]   = SpatialDepthWiseConvolution(self.head_dim,kernel)138                self.value_depthwiseconv[str(kernel_idx)] = SpatialDepthWiseConvolution(self.head_dim,kernel)139 140    def prune_heads(self, heads):141        if len(heads) == 0:142            return143        heads, index = find_pruneable_heads_and_indices(heads, self.num_heads, self.head_dim, self.pruned_heads)144        index_attn = torch.cat([index, index + self.split_size, index + (2 * self.split_size)])145 146        # Prune conv1d layers147        self.c_attn = prune_conv1d_layer(self.c_attn, index_attn, dim=1)148        self.c_proj = prune_conv1d_layer(self.c_proj, index, dim=0)149 150        # Update hyper params151        self.split_size = (self.split_size // self.num_heads) * (self.num_heads - len(heads))152        self.num_heads = self.num_heads - len(heads)153        self.pruned_heads = self.pruned_heads.union(heads)154 155    def _attn(self, query, key, value, attention_mask=None, head_mask=None, alibi_bias=None):156        attn_weights = torch.matmul(query, key.transpose(-1, -2))157 158        if self.scale_attn_weights:159            attn_weights = attn_weights / (float(value.size(-1)) ** 0.5)160 161        if not self.is_cross_attention:162            # if only "normal" attention layer implements causal mask163            query_length, key_length = query.size(-2), key.size(-2)164            causal_mask = self.bias[:, :, key_length - query_length : key_length, :key_length].bool()165            attn_weights = torch.where(causal_mask, attn_weights, self.masked_bias.to(attn_weights.dtype))166 167        if alibi_bias is not None:168            attn_weights = attn_weights + alibi_bias[:,:,:attn_weights.size(-1)]169 170        if attention_mask is not None:171            # Apply the attention mask172            attn_weights = attn_weights + attention_mask173 174        attn_weights = nn.Softmax(dim=-1)(attn_weights)175        attn_weights = self.attn_dropout(attn_weights)176 177        # Mask heads if we want to178        if head_mask is not None:179            attn_weights = attn_weights * head_mask180 181        attn_output = torch.matmul(attn_weights, value)182 183        return attn_output, attn_weights184 185    def _split_heads(self, tensor, num_heads, attn_head_size):186        """187        Splits hidden_size dim into attn_head_size and num_heads188        """189        new_shape = tensor.size()[:-1] + (num_heads, attn_head_size)190        tensor = tensor.view(*new_shape)191        return tensor.permute(0, 2, 1, 3)  # (batch, head, seq_length, head_features)192 193    def _merge_heads(self, tensor, num_heads, attn_head_size):194        """195        Merges attn_head_size dim and num_attn_heads dim into hidden_size196        """197        tensor = tensor.permute(0, 2, 1, 3).contiguous()198        new_shape = tensor.size()[:-2] + (num_heads * attn_head_size,)199        return tensor.view(new_shape)200 201    def forward(202        self,203        hidden_states,204        layer_past=None,205        attention_mask=None,206        head_mask=None,207        encoder_hidden_states=None,208        encoder_attention_mask=None,209        use_cache=False,210        output_attentions=False,211        alibi_bias=None,212    ):213        if encoder_hidden_states is not None:214            if not hasattr(self, "q_attn"):215                raise ValueError(216                    "If class is used as cross attention, the weights `q_attn` have to be defined. "217                    "Please make sure to instantiate class with `GPT2Attention(..., is_cross_attention=True)`."218                )219 220            query = self.q_attn(hidden_states)221            key, value = self.c_attn(encoder_hidden_states).split(self.split_size, dim=2)222            attention_mask = encoder_attention_mask223        else:224            query, key, value = self.c_attn(hidden_states).split(self.split_size, dim=2)225 226        query = self._split_heads(query, self.num_heads, self.head_dim)227        key = self._split_heads(key, self.num_heads, self.head_dim)228        value = self._split_heads(value, self.num_heads, self.head_dim)229 230        if layer_past is not None:231            past_key, past_value = layer_past232            key = torch.cat((past_key, key), dim=-2)233            value = torch.cat((past_value, value), dim=-2)234 235        if use_cache is True:236            present = (key, value)237        else:238            present = None239        240        if self.attention_mode=="tranception":241            # We do not do anything on the first self.num_heads_per_kernel_size heads (kernel =1)242            query_list=[query[:,:self.num_heads_per_kernel_size,:,:]]243            key_list=[key[:,:self.num_heads_per_kernel_size,:,:]]244            value_list=[value[:,:self.num_heads_per_kernel_size,:,:]]245            for kernel_idx in range(3):246                query_list.append(self.query_depthwiseconv[str(kernel_idx)](query[:,(kernel_idx+1)*self.num_heads_per_kernel_size:(kernel_idx+2)*self.num_heads_per_kernel_size,:,:]))247                key_list.append(self.key_depthwiseconv[str(kernel_idx)](key[:,(kernel_idx+1)*self.num_heads_per_kernel_size:(kernel_idx+2)*self.num_heads_per_kernel_size,:,:]))248                value_list.append(self.value_depthwiseconv[str(kernel_idx)](value[:,(kernel_idx+1)*self.num_heads_per_kernel_size:(kernel_idx+2)*self.num_heads_per_kernel_size,:,:]))249            query=torch.cat(query_list, dim=1)250            key=torch.cat(key_list, dim=1)251            value=torch.cat(value_list, dim=1)252        253        attn_output, attn_weights = self._attn(query, key, value, attention_mask, head_mask, alibi_bias=alibi_bias)254 255        attn_output = self._merge_heads(attn_output, self.num_heads, self.head_dim)256        attn_output = self.c_proj(attn_output)257        attn_output = self.resid_dropout(attn_output)258 259        outputs = (attn_output, present)260        if output_attentions:261            outputs += (attn_weights,)262 263        return outputs  # a, present, (attentions)264 265class TranceptionBlockMLP(nn.Module):266    def __init__(self, intermediate_size, config):267        super().__init__()268        embed_dim = config.hidden_size269        self.c_fc = Conv1D(intermediate_size, embed_dim)270        self.c_proj = Conv1D(embed_dim, intermediate_size)271        self.act = tranception_ACT2FN[config.activation_function]272        self.dropout = nn.Dropout(config.resid_pdrop)273    274    def forward(self, hidden_states):275        hidden_states = self.c_fc(hidden_states)276        hidden_states = self.act(hidden_states)277        hidden_states = self.c_proj(hidden_states)278        hidden_states = self.dropout(hidden_states)279        return hidden_states280 281class TranceptionBlock(nn.Module):282    def __init__(self, config, SDWC_kernel_size=None):283        super().__init__()284        hidden_size = config.hidden_size285        inner_dim = config.n_inner if config.n_inner is not None else 4 * hidden_size286 287        self.ln_1 = nn.LayerNorm(hidden_size, eps=config.layer_norm_epsilon)288        self.attn = TranceptionBlockAttention(config, SDWC_kernel_size=SDWC_kernel_size)289        self.ln_2 = nn.LayerNorm(hidden_size, eps=config.layer_norm_epsilon)290 291        if config.add_cross_attention:292            self.crossattention = TranceptionBlockAttention(config, is_cross_attention=True, SDWC_kernel_size=SDWC_kernel_size)293            self.ln_cross_attn = nn.LayerNorm(hidden_size, eps=config.layer_norm_epsilon)294 295        self.mlp = TranceptionBlockMLP(inner_dim, config)296    297    def forward(298        self,299        hidden_states,300        layer_past=None,301        attention_mask=None,302        head_mask=None,303        encoder_hidden_states=None,304        encoder_attention_mask=None,305        use_cache=False,306        output_attentions=False,307        alibi_bias=None,308    ):309        residual = hidden_states310        hidden_states = self.ln_1(hidden_states)311        attn_outputs = self.attn(312            hidden_states,313            layer_past=layer_past,314            attention_mask=attention_mask,315            head_mask=head_mask,316            use_cache=use_cache,317            output_attentions=output_attentions,318            alibi_bias=alibi_bias,319        )320        attn_output = attn_outputs[0]  # output_attn: a, present, (attentions)321        outputs = attn_outputs[1:]322        # residual connection323        hidden_states = attn_output + residual324 325        if encoder_hidden_states is not None:326            # add one self-attention block for cross-attention327            if not hasattr(self, "crossattention"):328                raise ValueError(329                    f"If `encoder_hidden_states` are passed, {self} has to be instantiated with "330                    "cross-attention layers by setting `config.add_cross_attention=True`"331                )332            residual = hidden_states333            hidden_states = self.ln_cross_attn(hidden_states)334            cross_attn_outputs = self.crossattention(335                hidden_states,336                attention_mask=attention_mask,337                head_mask=head_mask,338                encoder_hidden_states=encoder_hidden_states,339                encoder_attention_mask=encoder_attention_mask,340                output_attentions=output_attentions,341            )342            attn_output = cross_attn_outputs[0]343            # residual connection344            hidden_states = residual + attn_output345            outputs = outputs + cross_attn_outputs[2:]  # add cross attentions if we output attention weights346 347        residual = hidden_states348        hidden_states = self.ln_2(hidden_states)349 350        feed_forward_hidden_states = self.mlp(hidden_states)351        352        # residual connection353        hidden_states = residual + feed_forward_hidden_states354 355        if use_cache:356            outputs = (hidden_states,) + outputs357        else:358            outputs = (hidden_states,) + outputs[1:]359 360        return outputs  # hidden_states, present, (attentions, cross_attentions)361 362class TranceptionModel(GPT2PreTrainedModel):363    _keys_to_ignore_on_load_missing = ["attn.masked_bias"]364    def __init__(self, config):365        super().__init__(config)366 367        self.embed_dim = config.hidden_size368        self.wte = nn.Embedding(config.vocab_size, self.embed_dim)369        self.position_embedding = config.position_embedding if hasattr(config, "position_embedding") else "learned"370        if self.position_embedding=="learned":371            self.wpe = nn.Embedding(config.max_position_embeddings, self.embed_dim)372            self.alibi = None373        elif self.position_embedding=="grouped_alibi":374            maxpos = config.n_positions375            attn_heads = config.n_head376            self.slopes = torch.Tensor(get_slopes(attn_heads, mode=self.position_embedding))377            #The softmax operation is invariant to translation, and bias functions used are always linear. 378            alibi = self.slopes.unsqueeze(1).unsqueeze(1) * torch.arange(maxpos).unsqueeze(0).unsqueeze(0).expand(attn_heads, -1, -1)379            alibi = alibi.view(attn_heads, 1, maxpos)380            self.register_buffer('alibi',alibi)381 382        self.drop = nn.Dropout(config.embd_pdrop)383        self.h = nn.ModuleList([TranceptionBlock(config) for _ in range(config.num_hidden_layers)])384        self.ln_f = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_epsilon)385 386        self.init_weights()387 388        # Model parallel389        self.model_parallel = False390        self.device_map = None391        self.gradient_checkpointing = False392    393    def parallelize(self, device_map=None, num_cores=None):394        self.device_map = (395                get_device_map(len(self.h), range(torch.cuda.device_count())) if device_map is None else device_map396            )397        device_prefix="cuda:"398        assert_device_map(self.device_map, len(self.h))399        self.model_parallel = True400        self.first_device = "cpu" if "cpu" in self.device_map.keys() else device_prefix + str(min(self.device_map.keys()))401        self.last_device = device_prefix + str(max(self.device_map.keys()))402        self.wte = self.wte.to(self.first_device)403        if self.position_embedding=="learned":404            self.wpe = self.wpe.to(self.first_device)405        for k, v in self.device_map.items():406            print("k,v :"+str(k)+","+str(v))407            for block in v:408                cuda_device = device_prefix + str(k)409                self.h[block] = self.h[block].to(cuda_device)410        self.ln_f = self.ln_f.to(self.last_device)411    412    def deparallelize(self):413        self.model_parallel = False414        self.device_map = None415        self.first_device = "cpu"416        self.last_device = "cpu"417        self.wte = self.wte.to("cpu")418        if self.position_embedding=="learned":419            self.wpe = self.wpe.to("cpu")420        for index in range(len(self.h)):421            self.h[index] = self.h[index].to("cpu")422        self.ln_f = self.ln_f.to("cpu")423        torch.cuda.empty_cache()424 425    def get_input_embeddings(self):426        return self.wte427 428    def set_input_embeddings(self, new_embeddings):429        self.wte = new_embeddings430 431    def _prune_heads(self, heads_to_prune):432        """433        Prunes heads of the model. heads_to_prune: dict of {layer_num: list of heads to prune in this layer}434        """435        for layer, heads in heads_to_prune.items():436            self.h[layer].attn.prune_heads(heads)437 438    def forward(439        self,440        input_ids=None,441        past_key_values=None,442        attention_mask=None,443        token_type_ids=None,444        position_ids=None,445        head_mask=None,446        inputs_embeds=None,447        encoder_hidden_states=None,448        encoder_attention_mask=None,449        use_cache=None,450        output_attentions=None,451        output_hidden_states=None,452        return_dict=None,453    ):454        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions455        output_hidden_states = (456            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states457        )458        use_cache = use_cache if use_cache is not None else self.config.use_cache459        return_dict = return_dict if return_dict is not None else self.config.use_return_dict460 461        if input_ids is not None and inputs_embeds is not None:462            raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")463        elif input_ids is not None:464            input_shape = input_ids.size()465            input_ids = input_ids.view(-1, input_shape[-1])466            batch_size = input_ids.shape[0]467        elif inputs_embeds is not None:468            input_shape = inputs_embeds.size()[:-1]469            batch_size = inputs_embeds.shape[0]470        else:471            raise ValueError("You have to specify either input_ids or inputs_embeds")472 473        device = input_ids.device if input_ids is not None else inputs_embeds.device474 475        if token_type_ids is not None:476            token_type_ids = token_type_ids.view(-1, input_shape[-1])477        if position_ids is not None:478            position_ids = position_ids.view(-1, input_shape[-1])479 480        if past_key_values is None:481            past_length = 0482            past_key_values = tuple([None] * len(self.h))483        else:484            past_length = past_key_values[0][0].size(-2)485        if position_ids is None:486            position_ids = torch.arange(past_length, input_shape[-1] + past_length, dtype=torch.long, device=device)487            position_ids = position_ids.unsqueeze(0).view(-1, input_shape[-1])488 489        # GPT2Attention mask.490        if attention_mask is not None:491            if batch_size <= 0:492                raise ValueError("batch_size has to be defined and > 0")493            attention_mask = attention_mask.view(batch_size, -1)494            # We create a 3D attention mask from a 2D tensor mask.495            # Sizes are [batch_size, 1, 1, to_seq_length]496            # So we can broadcast to [batch_size, num_heads, from_seq_length, to_seq_length]497            # this attention mask is more simple than the triangular masking of causal attention498            # used in OpenAI GPT, we just need to prepare the broadcast dimension here.499            attention_mask = attention_mask[:, None, None, :]500 501            # Since attention_mask is 1.0 for positions we want to attend and 0.0 for502            # masked positions, this operation will create a tensor which is 0.0 for503            # positions we want to attend and -10000.0 for masked positions.504            # Since we are adding it to the raw scores before the softmax, this is505            # effectively the same as removing these entirely.506            attention_mask = attention_mask.to(dtype=self.dtype)  # fp16 compatibility507            attention_mask = (1.0 - attention_mask) * -10000.0508 509        # If a 2D ou 3D attention mask is provided for the cross-attention510        # we need to make broadcastable to [batch_size, num_heads, seq_length, seq_length]511        if self.config.add_cross_attention and encoder_hidden_states is not None:512            encoder_batch_size, encoder_sequence_length, _ = encoder_hidden_states.size()513            encoder_hidden_shape = (encoder_batch_size, encoder_sequence_length)514            if encoder_attention_mask is None:515                encoder_attention_mask = torch.ones(encoder_hidden_shape, device=device)516            encoder_attention_mask = self.invert_attention_mask(encoder_attention_mask)517        else:518            encoder_attention_mask = None519 520        # Prepare head mask if needed521        # 1.0 in head_mask indicate we keep the head522        # attention_probs has shape bsz x n_heads x N x N523        # head_mask has shape n_layer x batch x n_heads x N x N524        head_mask = self.get_head_mask(head_mask, self.config.n_layer)525 526        if inputs_embeds is None:527            inputs_embeds = self.wte(input_ids)528        if self.position_embedding=="learned":529            position_embeds = self.wpe(position_ids)530            hidden_states = inputs_embeds + position_embeds531        else:532            hidden_states = inputs_embeds533 534        if token_type_ids is not None:535            token_type_embeds = self.wte(token_type_ids)536            hidden_states = hidden_states + token_type_embeds537 538        hidden_states = self.drop(hidden_states)539 540        output_shape = input_shape + (hidden_states.size(-1),)541 542        presents = () if use_cache else None543        all_self_attentions = () if output_attentions else None544        all_cross_attentions = () if output_attentions and self.config.add_cross_attention else None545        all_hidden_states = () if output_hidden_states else None546        547        for i, (block, layer_past) in enumerate(zip(self.h, past_key_values)):548            # Model parallel549            if self.model_parallel:550                torch.cuda.set_device(hidden_states.device)551                # Ensure layer_past is on same device as hidden_states (might not be correct)552                if layer_past is not None:553                    layer_past = tuple(past_state.to(hidden_states.device) for past_state in layer_past)554                # Ensure that attention_mask is always on the same device as hidden_states555                if attention_mask is not None:556                    attention_mask = attention_mask.to(hidden_states.device)557                if isinstance(head_mask, torch.Tensor):558                    head_mask = head_mask.to(hidden_states.device)559            if output_hidden_states:560                all_hidden_states = all_hidden_states + (hidden_states,)561 562            if self.gradient_checkpointing and self.training:563                if use_cache:564                    logger.warning(565                        "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."566                    )567                    use_cache = False568 569                def create_custom_forward(module):570                    def custom_forward(*inputs):571                        # None for past_key_value572                        return module(*inputs, use_cache, output_attentions)573 574                    return custom_forward575 576                outputs = torch.utils.checkpoint.checkpoint(577                    create_custom_forward(block),578                    hidden_states,579                    None,580                    attention_mask,581                    head_mask[i],582                    encoder_hidden_states,583                    encoder_attention_mask,584                )585            else:586                outputs = block(587                    hidden_states,588                    layer_past=layer_past,589                    attention_mask=attention_mask,590                    head_mask=head_mask[i],591                    encoder_hidden_states=encoder_hidden_states,592                    encoder_attention_mask=encoder_attention_mask,593                    use_cache=use_cache,594                    output_attentions=output_attentions,595                    alibi_bias=self.alibi if hasattr(self, "alibi") else None596                )597 598            hidden_states = outputs[0]599            600            if use_cache is True:601                presents = presents + (outputs[1],)602 603            if output_attentions:604                all_self_attentions = all_self_attentions + (outputs[2 if use_cache else 1],)605                if self.config.add_cross_attention:606                    all_cross_attentions = all_cross_attentions + (outputs[3 if use_cache else 2],)607 608            if self.model_parallel:609                device_prefix="cuda:"610                for k, v in self.device_map.items():611                    if i == v[-1] and device_prefix + str(k) != self.last_device:612                        hidden_states = hidden_states.to(device_prefix + str(k + 1))613 614        hidden_states = self.ln_f(hidden_states)615 616        hidden_states = hidden_states.view(*output_shape)617        # Add last hidden state618        if output_hidden_states:619            all_hidden_states = all_hidden_states + (hidden_states,)620 621        if not return_dict:622            return tuple(623                v624                for v in [hidden_states, presents, all_hidden_states, all_self_attentions, all_cross_attentions, moe_loss]625                if v is not None626            )627        628        return BaseModelOutputWithPastAndCrossAttentions(629                last_hidden_state=hidden_states,630                past_key_values=presents,631                hidden_states=all_hidden_states,632                attentions=all_self_attentions,633                cross_attentions=all_cross_attentions,634            )635 636class TranceptionLMHeadModel(GPT2PreTrainedModel):637    _keys_to_ignore_on_load_missing = [r"attn.masked_bias", r"attn.bias", r"lm_head.weight"]638    def __init__(self, config):639        super().__init__(config)640        self.transformer = TranceptionModel(config)641        self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)642        self.config = config643 644        self.init_weights()645        646        self.default_model_device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")647        # Model parallel648        self.model_parallel = False649        self.device_map = None650        651        self.retrieval_aggregation_mode = config.retrieval_aggregation_mode if hasattr(config, "retrieval_aggregation_mode") else None652        if self.retrieval_aggregation_mode is not None:653            print("Model leverages both autoregressive and retrieval inference")654            self.MSA_filename = config.MSA_filename if hasattr(config, "MSA_filename") else False655            self.MSA_folder = '/'.join(self.MSA_filename.split(os.sep)[:-1])656            self.MSA_name = self.MSA_filename.split(os.sep)[-1]657            self.retrieval_inference_weight_LR = config.retrieval_inference_weight if hasattr(config, "retrieval_inference_weight") else 0.6658            self.retrieval_inference_weight_RL = config.retrieval_inference_weight if hasattr(config, "retrieval_inference_weight") else 0.6659            self.MSA_start=config.MSA_start660            self.MSA_end=config.MSA_end661            self.full_protein_length = config.full_protein_length if hasattr(config, "full_protein_length") else -1662            663            self.MSA_log_prior = torch.log(torch.tensor(664                                                        msa_utils.get_msa_prior(665                                                            MSA_data_file=self.MSA_filename, 666                                                            MSA_weight_file_name=config.MSA_weight_file_name, 667                                                            retrieval_aggregation_mode=self.retrieval_aggregation_mode,668                                                            MSA_start=self.MSA_start,669                                                            MSA_end=self.MSA_end,670                                                            len_target_seq=self.full_protein_length, 671                                                            vocab=config.tokenizer.get_vocab(), 672                                                            verbose=False673                                                        )674                                            ).float().to(self.default_model_device))675        else:676            print("Model only uses autoregressive inference")677 678    def parallelize(self, device_map=None, num_cores=None, num_pipelines=1):679        self.num_pipelines=num_pipelines680        self.device_map = (681                get_device_map(len(self.transformer.h), range(torch.cuda.device_count()))682                if device_map is None683                else device_map684            )685        assert_device_map(self.device_map, len(self.transformer.h))686        self.transformer.parallelize(self.device_map, num_cores=num_cores)687        self.lm_head = self.lm_head.to(self.transformer.first_device)688        self.model_parallel = True689 690    def deparallelize(self):691        self.transformer.deparallelize()692        self.transformer = self.transformer.to("cpu")693        self.lm_head = self.lm_head.to("cpu")694        self.model_parallel = False695        torch.cuda.empty_cache()696 697    def get_output_embeddings(self):698        return self.lm_head699 700    def set_output_embeddings(self, new_embeddings):701        self.lm_head = new_embeddings702 703    def prepare_inputs_for_generation(self, input_ids, past=None, **kwargs):704        token_type_ids = kwargs.get("token_type_ids", None)705        # only last token for inputs_ids if past is defined in kwargs706        if past:707            input_ids = input_ids[:, -1].unsqueeze(-1)708            if token_type_ids is not None:709                token_type_ids = token_type_ids[:, -1].unsqueeze(-1)710 711        attention_mask = kwargs.get("attention_mask", None)712        position_ids = kwargs.get("position_ids", None)713 714        if attention_mask is not None and position_ids is None:715            # create position_ids on the fly for batch generation716            position_ids = attention_mask.long().cumsum(-1) - 1717            position_ids.masked_fill_(attention_mask == 0, 1)718            if past:719                position_ids = position_ids[:, -1].unsqueeze(-1)720        else:721            position_ids = None722        723        return {724                "input_ids": input_ids,725                "past_key_values": past,726                "use_cache": kwargs.get("use_cache"),727                "position_ids": position_ids,728                "attention_mask": attention_mask,729                "token_type_ids": token_type_ids,730                "flip": kwargs.get("flip", None),731            }732 733    def forward(734        self,735        input_ids=None,736        past_key_values=None,737        attention_mask=None,738        token_type_ids=None,739        position_ids=None,740        head_mask=None,741        inputs_embeds=None,742        encoder_hidden_states=None,743        encoder_attention_mask=None,744        labels=None,745        use_cache=None,746        output_attentions=None,747        output_hidden_states=None,748        return_dict=None,749        flip=None,750        start_slice=None,751        end_slice=None,752        mutated_sequence=None,753    ):754        r"""755        labels (:obj:`torch.LongTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):756            Labels for language modeling. Note that the labels **are shifted** inside the model, i.e. you can set757            ``labels = input_ids`` Indices are selected in ``[-100, 0, ..., config.vocab_size]`` All labels set to758            ``-100`` are ignored (masked), the loss is only computed for labels in ``[0, ..., config.vocab_size]``759        """760        return_dict = return_dict if return_dict is not None else self.config.use_return_dict761        762        transformer_outputs = self.transformer(763            input_ids,764            past_key_values=past_key_values,765            attention_mask=attention_mask,766            token_type_ids=token_type_ids,767            position_ids=position_ids,768            head_mask=head_mask,769            inputs_embeds=inputs_embeds,770            encoder_hidden_states=encoder_hidden_states,771            encoder_attention_mask=encoder_attention_mask,772            use_cache=use_cache,773            output_attentions=output_attentions,774            output_hidden_states=output_hidden_states,775            return_dict=return_dict776        )777        hidden_states = transformer_outputs[0]778            779        # Set device for model parallelism780        if self.model_parallel:781            torch.cuda.set_device(self.transformer.first_device)782            hidden_states = hidden_states.to(self.lm_head.weight.device)783            self.MSA_log_prior = self.MSA_log_prior.to(self.lm_head.weight.device)784 785        lm_logits = self.lm_head(hidden_states)786 787        loss = None788        if labels is not None:789            # Shift so that tokens < n predict n790            shift_logits = lm_logits[..., :-1, :].contiguous()791            shift_labels = labels[..., 1:].contiguous()792            793            if self.retrieval_aggregation_mode is not None:794                batch_size = input_ids.size(0)795                796                if self.retrieval_aggregation_mode=="aggregate_indel":797                    assert batch_size==1, "Aggregate indel is only supported for batch size of 1"798                    truncated_sequence_text = mutated_sequence[0][start_slice[0]:end_slice[0]]799                    if len(truncated_sequence_text)!=shift_logits.shape[1]-1: # shift_logits only has one extra token compared to truncated_sequence_text (the BOS token)800                        print("Tokenization error -- seq length: {} and shift_logits length - 1 : {}".format(len(mutated_sequence),shift_logits.shape[1]-1))801                    MSA_log_prior, MSA_start, MSA_end = msa_utils.update_retrieved_MSA_log_prior_indel(self, self.MSA_log_prior, self.MSA_start, self.MSA_end, mutated_sequence[0])  802                803                elif self.retrieval_aggregation_mode=="aggregate_substitution":804                    MSA_log_prior=self.MSA_log_prior805                    MSA_start=self.MSA_start806                    MSA_end=self.MSA_end807                808                shift_log_probas = torch.log_softmax(shift_logits,dim=-1)809                fused_shift_log_probas = shift_log_probas.clone()810                if flip is None:811                    flip = torch.zeros(batch_size).to(fused_shift_log_probas.device)812                flip = flip > 0813                814                for seq_index in range(batch_size):815                    min_prior_slice = max(start_slice[seq_index], MSA_start) 816                    max_prior_slice = min(end_slice[seq_index], MSA_end)817                    818                    if max_prior_slice <= min_prior_slice:819                        print("Non overlapping region detected: min_prior_slice {} and max_prior_slice {}".format(min_prior_slice,max_prior_slice))820                        continue821                    822                    slice_prior = MSA_log_prior[min_prior_slice:max_prior_slice,:].to(fused_shift_log_probas.device) 823                    if flip[seq_index]:824                        slice_prior = torch.flip(slice_prior,dims=(0,))825                        min_logits_slice = max(0,end_slice[seq_index]-MSA_end) 826                        max_logits_slice = min_logits_slice + (max_prior_slice-min_prior_slice)827                        fused_shift_log_probas[seq_index,min_logits_slice:max_logits_slice,:] = (1-self.retrieval_inference_weight_RL)*shift_log_probas[seq_index,min_logits_slice:max_logits_slice,:] + self.retrieval_inference_weight_RL*slice_prior828                    else:829                        min_logits_slice = max(0, MSA_start-start_slice[seq_index]) 830                        max_logits_slice = min_logits_slice + (max_prior_slice-min_prior_slice)831                        fused_shift_log_probas[seq_index,min_logits_slice:max_logits_slice,:] = (1-self.retrieval_inference_weight_LR)*shift_log_probas[seq_index,min_logits_slice:max_logits_slice,:] + self.retrieval_inference_weight_LR*slice_prior832                833                if self.retrieval_aggregation_mode=="aggregate_indel":834                    try:835                        # If a given residue colume is an added zero-column, then we overwrite prior fusion and only predict based on the autoregressive transformer inference mode.836                        inserted_retrieval_positions = [True if slice_prior[i].sum()==0 else False for i in range(len(slice_prior))]+[True] #Last True is for the end of sentence token837                        fused_shift_log_probas[:,inserted_retrieval_positions,:]=shift_log_probas[:,inserted_retrieval_positions,:]838                    except:839                        print("Error when adding zero column(s) to account for insertion mutations.")840                841                loss_fct = NLLLoss(reduction='none')842                loss = loss_fct(input=fused_shift_log_probas.view(-1, fused_shift_log_probas.size(-1)), target=shift_labels.view(-1)).view(fused_shift_log_probas.shape[0],fused_shift_log_probas.shape[1])843                mask = attention_mask[..., 1:].float()844                mask[mask==0]=float('nan')845                loss *= mask846                loss = nanmean(loss, dim=1).mean()847            else:848                loss_fct = CrossEntropyLoss()849                loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))850                fused_shift_log_probas = None851 852        if not return_dict:853            output = (lm_logits,) + transformer_outputs[1:] 854            return ((loss,) + output) if loss is not None else output855        856        return TranceptionCausalLMOutputWithCrossAttentions(857            loss=loss,858            logits=lm_logits,859            past_key_values=transformer_outputs.past_key_values,860            hidden_states=transformer_outputs.hidden_states,861            attentions=transformer_outputs.attentions,862            cross_attentions=transformer_outputs.cross_attentions,863            fused_shift_log_probas=fused_shift_log_probas864        )865 866 867    @staticmethod868    def _reorder_cache(past: Tuple[Tuple[torch.Tensor]], beam_idx: torch.Tensor) -> Tuple[Tuple[torch.Tensor]]:869        """870        This function is used to re-order the :obj:`past_key_values` cache if871        :meth:`~transformers.PreTrainedModel.beam_search` or :meth:`~transformers.PreTrainedModel.beam_sample` is872        called. This is required to match :obj:`past_key_values` with the correct beam_idx at every generation step.873        """874        return tuple(875            tuple(past_state.index_select(0, beam_idx.to(past_state.device)) for past_state in layer_past)876            for layer_past in past877        )878    879    def score_mutants(self, DMS_data, target_seq=None, scoring_mirror=True, batch_size_inference=10, num_workers=10, indel_mode=False):880        """881        Method to score mutants in an input DMS file.882        DMS_data: (dataframe) Dataframe containing the list of mutated sequences for scoring.883        target_seq: (string) Full reference sequence (wild type) that is mutated in the DMS assay. If not None, returned scores are delta log likelihood wrt that sequence.884        scoring_mirror: (bool) Whether to score mutated sequences from both directions (Left->Right and Right->Left).885        batch_size_inference: (int) Batch size for scoring.886        num_workers: (int) Number of workers to be used in the data loader.887        indel_mode: (bool) Flag to be used when scoring insertions and deletions. Otherwise assumes substitutions.888        """889        df = DMS_data.copy()890        if ('mutated_sequence' not in df) and (not indel_mode): df['mutated_sequence'] = df['mutant'].apply(lambda x: scoring_utils.get_mutated_sequence(target_seq, x))891        assert ('mutated_sequence' in df), "DMS file to score does not have mutated_sequence column"892        #if 'mutant' not in df: df['mutant'] = df['mutated_sequence'] #if mutant not in DMS file we default to mutated_sequence893        if 'DMS_score' in df: del df['DMS_score'] 894        if 'DMS_score_bin' in df: del df['DMS_score_bin'] 895        if target_seq is not None:896            df_left_to_right_slices = scoring_utils.get_sequence_slices(df, target_seq=target_seq, model_context_len = self.config.n_ctx - 2, indel_mode=indel_mode, scoring_window=self.config.scoring_window)897        else:898            df_left_to_right_slices = scoring_utils.get_sequence_slices(df, target_seq=list(df['mutated_sequence'])[0], model_context_len = self.config.n_ctx - 2, indel_mode=indel_mode, scoring_window='sliding')899        print("Scoring sequences from left to right")900        scores_L_to_R = scoring_utils.get_tranception_scores_mutated_sequences(model=self, mutated_sequence_df=df_left_to_right_slices, batch_size_inference=batch_size_inference, score_var_name='avg_score_L_to_R', target_seq=target_seq, num_workers=num_workers, indel_mode=indel_mode)901        if scoring_mirror: 902            print("Scoring sequences from right to left")903            df_right_to_left_slices = df_left_to_right_slices.copy()904            df_right_to_left_slices['sliced_mutated_sequence'] = df_right_to_left_slices['sliced_mutated_sequence'].apply(lambda x: x[::-1])905            scores_R_to_L = scoring_utils.get_tranception_scores_mutated_sequences(model=self, mutated_sequence_df=df_right_to_left_slices, batch_size_inference=batch_size_inference, score_var_name='avg_score_R_to_L', target_seq=target_seq, num_workers=num_workers, reverse=True, indel_mode=indel_mode)906            all_scores = pd.merge(scores_L_to_R, scores_R_to_L, on='mutated_sequence', how='left', suffixes=('','_R_to_L'))907            all_scores['avg_score'] = (all_scores['avg_score_L_to_R'] + all_scores['avg_score_R_to_L']) / 2.0908        else:909            all_scores = scores_L_to_R910            all_scores['avg_score'] = all_scores['avg_score_L_to_R']911        #By design "get_tranception_scores_mutated_sequences" drops the WT from the output. We add it back if that was one of the sequences to score in the DMS (score=0 by definition)912        if target_seq in DMS_data.mutated_sequence.values:913            print("LEMON")914            if scoring_mirror:915                wt_row = pd.DataFrame([[target_seq,0,0,0]], columns=['mutated_sequence','avg_score_L_to_R','avg_score_R_to_L','avg_score'])916            else:917                wt_row = pd.DataFrame([[target_seq,0,0]], columns=['mutated_sequence','avg_score_L_to_R','avg_score'])918            all_scores = pd.concat([all_scores,wt_row], ignore_index=True)919        return all_scores920 921    def encode_batch(self, protein_sequence, sequence_name="sliced_mutated_sequence"):922        """923        Method to process an input AA sequence batch (protein_sequence) and return a tokenized sequence (via the tokenizer associated to the model).924        """925        protein_sequence[sequence_name] = scoring_utils.sequence_replace(sequences=protein_sequence[sequence_name], char_to_replace='X', char_replacements='ACDEFGHIKLMNPQRSTVWY')926        protein_sequence[sequence_name] = scoring_utils.sequence_replace(sequences=protein_sequence[sequence_name], char_to_replace='B', char_replacements='DN')927        protein_sequence[sequence_name] = scoring_utils.sequence_replace(sequences=protein_sequence[sequence_name], char_to_replace='J', char_replacements='IL')928        protein_sequence[sequence_name] = scoring_utils.sequence_replace(sequences=protein_sequence[sequence_name], char_to_replace='Z', char_replacements='EQ')929        return self.config.tokenizer(list(protein_sequence[sequence_name]), add_special_tokens=True, truncation=True, padding=True, max_length=self.config.n_ctx)930 931