PascalNotin/Tranception_design
4
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 