Taykhoom/CodonBERT
0133
1import math2from typing import Optional, Tuple, Union3 4import torch5import torch.nn as nn6import torch.nn.functional as F7import torch.utils.checkpoint8from transformers.activations import ACT2FN9from transformers import PreTrainedModel, PretrainedConfig10from transformers.modeling_outputs import BaseModelOutputWithPooling, MaskedLMOutput11 12from .configuration_bert_updated import BertUpdatedConfig13 14 15class BertSelfAttention(nn.Module):16 17 def __init__(self, config):18 super().__init__()19 self.num_attention_heads = config.num_attention_heads20 self.attention_head_size = config.hidden_size // config.num_attention_heads21 self.all_head_size = self.num_attention_heads * self.attention_head_size22 23 self.query = nn.Linear(config.hidden_size, self.all_head_size)24 self.key = nn.Linear(config.hidden_size, self.all_head_size)25 self.value = nn.Linear(config.hidden_size, self.all_head_size)26 self.dropout = nn.Dropout(config.attention_probs_dropout_prob)27 28 def _split_heads(self, x: torch.Tensor) -> torch.Tensor:29 B, T, _ = x.shape30 return x.view(B, T, self.num_attention_heads, self.attention_head_size).permute(0, 2, 1, 3)31 32 def forward(33 self,34 hidden_states: torch.Tensor,35 key_padding_mask: Optional[torch.Tensor] = None,36 output_attentions: bool = False,37 head_mask: Optional[torch.Tensor] = None,38 ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:39 q = self._split_heads(self.query(hidden_states))40 k = self._split_heads(self.key(hidden_states))41 v = self._split_heads(self.value(hidden_states))42 43 scale = math.sqrt(self.attention_head_size)44 scores = torch.matmul(q, k.transpose(-1, -2)) / scale45 if key_padding_mask is not None:46 scores = scores.masked_fill(key_padding_mask[:, None, None, :], float("-inf"))47 probs = torch.nan_to_num(48 F.softmax(scores.float(), dim=-1),49 nan=0.0,50 )51 context_probs = self.dropout(probs)52 if head_mask is not None:53 context_probs = context_probs * head_mask54 context = torch.matmul(context_probs.to(v.dtype), v)55 56 B, _, T, _ = context.shape57 context = context.permute(0, 2, 1, 3).contiguous().view(B, T, self.all_head_size)58 59 if output_attentions:60 return context, probs61 return context, None62 63 64class BertSdpaSelfAttention(BertSelfAttention):65 66 def forward(67 self,68 hidden_states: torch.Tensor,69 key_padding_mask: Optional[torch.Tensor] = None,70 output_attentions: bool = False,71 head_mask: Optional[torch.Tensor] = None,72 ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:73 if output_attentions or head_mask is not None:74 return super().forward(75 hidden_states,76 key_padding_mask,77 output_attentions=output_attentions,78 head_mask=head_mask,79 )80 81 B, T, _ = hidden_states.shape82 q = self._split_heads(self.query(hidden_states))83 k = self._split_heads(self.key(hidden_states))84 v = self._split_heads(self.value(hidden_states))85 86 attn_mask = None87 if key_padding_mask is not None:88 attn_mask = torch.zeros(B, 1, 1, T, dtype=q.dtype, device=q.device)89 attn_mask = attn_mask.masked_fill(key_padding_mask[:, None, None, :], float("-inf"))90 91 context = F.scaled_dot_product_attention(92 q,93 k,94 v,95 attn_mask=attn_mask,96 dropout_p=self.dropout.p if self.training else 0.0,97 )98 context = context.permute(0, 2, 1, 3).contiguous().view(B, T, self.all_head_size)99 return context, None100 101 102class BertFlashSelfAttention(BertSelfAttention):103 104 def forward(105 self,106 hidden_states: torch.Tensor,107 key_padding_mask: Optional[torch.Tensor] = None,108 output_attentions: bool = False,109 head_mask: Optional[torch.Tensor] = None,110 ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:111 if (112 output_attentions113 or head_mask is not None114 or (115 key_padding_mask is not None116 and (~key_padding_mask).sum(dim=-1).eq(0).any()117 )118 ):119 return super().forward(120 hidden_states,121 key_padding_mask,122 output_attentions=output_attentions,123 head_mask=head_mask,124 )125 126 try:127 from flash_attn import flash_attn_func, flash_attn_varlen_func128 from flash_attn.bert_padding import pad_input, unpad_input129 except ImportError as e:130 raise ImportError(131 "flash_attn is required for attn_implementation='flash_attention_2'. "132 "Install with: pip install flash-attn --no-build-isolation"133 ) from e134 135 B, T, _ = hidden_states.shape136 q = self._split_heads(self.query(hidden_states)).permute(0, 2, 1, 3)137 k = self._split_heads(self.key(hidden_states)).permute(0, 2, 1, 3)138 v = self._split_heads(self.value(hidden_states)).permute(0, 2, 1, 3)139 140 if q.dtype not in (torch.float16, torch.bfloat16):141 raise ValueError(142 "flash_attention_2 requires float16 or bfloat16 model weights. "143 f"Received {q.dtype}."144 )145 146 if key_padding_mask is not None and key_padding_mask.any():147 attend = ~key_padding_mask148 q_u, indices, cu_seqlens, max_seqlen, _ = unpad_input(q, attend)149 k_u, _, _, _, _ = unpad_input(k, attend)150 v_u, _, _, _, _ = unpad_input(v, attend)151 out_u = flash_attn_varlen_func(152 q_u, k_u, v_u,153 cu_seqlens_q=cu_seqlens, cu_seqlens_k=cu_seqlens,154 max_seqlen_q=max_seqlen, max_seqlen_k=max_seqlen,155 dropout_p=self.dropout.p if self.training else 0.0,156 causal=False,157 )158 out = pad_input(out_u, indices, B, T)159 else:160 out = flash_attn_func(161 q,162 k,163 v,164 dropout_p=self.dropout.p if self.training else 0.0,165 causal=False,166 )167 168 out = out.reshape(B, T, self.all_head_size)169 return out, None170 171 172BERT_SELF_ATTENTION_CLASSES = {173 "eager": BertSelfAttention,174 "sdpa": BertSdpaSelfAttention,175 "flash_attention_2": BertFlashSelfAttention,176}177 178 179class BertSelfOutput(nn.Module):180 def __init__(self, config):181 super().__init__()182 self.dense = nn.Linear(config.hidden_size, config.hidden_size)183 self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)184 self.dropout = nn.Dropout(config.hidden_dropout_prob)185 186 def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor:187 hidden_states = self.dropout(self.dense(hidden_states))188 return self.LayerNorm(hidden_states + input_tensor)189 190 191class BertAttention(nn.Module):192 def __init__(self, config):193 super().__init__()194 attn_cls = BERT_SELF_ATTENTION_CLASSES[getattr(config, "_attn_implementation", "eager")]195 self.self = attn_cls(config)196 self.output = BertSelfOutput(config)197 198 def forward(199 self,200 hidden_states: torch.Tensor,201 key_padding_mask: Optional[torch.Tensor],202 output_attentions: bool = False,203 head_mask: Optional[torch.Tensor] = None,204 ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:205 self_out, attn_weights = self.self(206 hidden_states,207 key_padding_mask,208 output_attentions,209 head_mask,210 )211 return self.output(self_out, hidden_states), attn_weights212 213 214class BertIntermediate(nn.Module):215 def __init__(self, config):216 super().__init__()217 self.dense = nn.Linear(config.hidden_size, config.intermediate_size)218 if config.hidden_act in (None, "identity", "linear"):219 self.intermediate_act_fn = lambda value: value220 elif isinstance(config.hidden_act, str):221 self.intermediate_act_fn = ACT2FN[config.hidden_act]222 else:223 self.intermediate_act_fn = config.hidden_act224 225 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:226 return self.intermediate_act_fn(self.dense(hidden_states))227 228 229class BertOutput(nn.Module):230 def __init__(self, config):231 super().__init__()232 self.dense = nn.Linear(config.intermediate_size, config.hidden_size)233 self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)234 self.dropout = nn.Dropout(config.hidden_dropout_prob)235 236 def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor:237 hidden_states = self.dropout(self.dense(hidden_states))238 return self.LayerNorm(hidden_states + input_tensor)239 240 241class BertLayer(nn.Module):242 def __init__(self, config):243 super().__init__()244 self.attention = BertAttention(config)245 self.intermediate = BertIntermediate(config)246 self.output = BertOutput(config)247 248 def forward(249 self,250 hidden_states: torch.Tensor,251 key_padding_mask: Optional[torch.Tensor],252 output_attentions: bool = False,253 head_mask: Optional[torch.Tensor] = None,254 ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:255 attn_out, attn_weights = self.attention(256 hidden_states,257 key_padding_mask,258 output_attentions,259 head_mask,260 )261 return self.output(self.intermediate(attn_out), attn_out), attn_weights262 263 264class BertEncoder(nn.Module):265 def __init__(self, config):266 super().__init__()267 self.layer = nn.ModuleList([BertLayer(config) for _ in range(config.num_hidden_layers)])268 self.gradient_checkpointing = False269 270 def forward(271 self,272 hidden_states: torch.Tensor,273 key_padding_mask: Optional[torch.Tensor],274 output_hidden_states: bool = False,275 output_attentions: bool = False,276 head_mask: Optional[torch.Tensor] = None,277 ) -> Tuple:278 all_hidden_states = (hidden_states,) if output_hidden_states else None279 all_attentions = () if output_attentions else None280 281 for layer_index, layer in enumerate(self.layer):282 layer_head_mask = (283 head_mask[layer_index] if head_mask is not None else None284 )285 if self.gradient_checkpointing and self.training:286 hidden_states, attn_weights = torch.utils.checkpoint.checkpoint(287 layer.__call__,288 hidden_states,289 key_padding_mask,290 output_attentions,291 layer_head_mask,292 use_reentrant=False,293 )294 else:295 hidden_states, attn_weights = layer(296 hidden_states,297 key_padding_mask,298 output_attentions,299 layer_head_mask,300 )301 if output_hidden_states:302 all_hidden_states = all_hidden_states + (hidden_states,)303 if output_attentions:304 all_attentions = all_attentions + (attn_weights,)305 306 return hidden_states, all_hidden_states, all_attentions307 308 309class BertEmbeddings(nn.Module):310 def __init__(self, config):311 super().__init__()312 self.word_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id)313 self.position_embeddings = nn.Embedding(config.max_position_embeddings, config.hidden_size)314 self.token_type_embeddings = nn.Embedding(config.type_vocab_size, config.hidden_size)315 self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)316 self.dropout = nn.Dropout(config.hidden_dropout_prob)317 self.register_buffer("position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)), persistent=False)318 319 def forward(320 self,321 input_ids: Optional[torch.LongTensor] = None,322 token_type_ids: Optional[torch.LongTensor] = None,323 position_ids: Optional[torch.LongTensor] = None,324 inputs_embeds: Optional[torch.FloatTensor] = None,325 ) -> torch.Tensor:326 if (input_ids is None) == (inputs_embeds is None):327 raise ValueError("Specify exactly one of input_ids or inputs_embeds.")328 B, T = (329 input_ids.shape330 if input_ids is not None331 else inputs_embeds.shape[:2]332 )333 if token_type_ids is None:334 token_type_ids = torch.zeros(335 (B, T),336 dtype=torch.long,337 device=(338 input_ids.device339 if input_ids is not None340 else inputs_embeds.device341 ),342 )343 if position_ids is None:344 position_ids = self.position_ids[:, :T]345 x = (346 self.word_embeddings(input_ids)347 if inputs_embeds is None348 else inputs_embeds349 )350 x = x + self.token_type_embeddings(token_type_ids)351 x = x + self.position_embeddings(position_ids)352 return self.dropout(self.LayerNorm(x))353 354 355class BertPooler(nn.Module):356 def __init__(self, config):357 super().__init__()358 self.dense = nn.Linear(config.hidden_size, config.hidden_size)359 self.activation = nn.Tanh()360 361 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:362 return self.activation(self.dense(hidden_states[:, 0]))363 364 365class BertPredictionHeadTransform(nn.Module):366 def __init__(self, config):367 super().__init__()368 self.dense = nn.Linear(config.hidden_size, config.hidden_size)369 self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)370 hidden_act = (371 config.mlm_hidden_act372 if getattr(config, "mlm_hidden_act", None) is not None373 else config.hidden_act374 )375 if hidden_act in (None, "identity", "linear"):376 self.transform_act_fn = lambda value: value377 elif isinstance(hidden_act, str):378 self.transform_act_fn = ACT2FN[hidden_act]379 else:380 self.transform_act_fn = hidden_act381 382 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:383 return self.LayerNorm(self.transform_act_fn(self.dense(hidden_states)))384 385 386class BertModel(PreTrainedModel):387 config_class = BertUpdatedConfig388 base_model_prefix = "bert"389 _supports_sdpa = True390 _supports_flash_attn_2 = True391 supports_gradient_checkpointing = True392 _keys_to_ignore_on_load_missing = [r"pooler\."]393 394 def __init__(self, config, add_pooling_layer=True):395 super().__init__(config)396 self.embeddings = BertEmbeddings(config)397 self.encoder = BertEncoder(config)398 self.pooler = BertPooler(config) if add_pooling_layer else None399 self.post_init()400 401 def get_input_embeddings(self):402 return self.embeddings.word_embeddings403 404 def set_input_embeddings(self, value):405 self.embeddings.word_embeddings = value406 407 def forward(408 self,409 input_ids: Optional[torch.LongTensor] = None,410 attention_mask: Optional[torch.Tensor] = None,411 token_type_ids: Optional[torch.LongTensor] = None,412 position_ids: Optional[torch.LongTensor] = None,413 head_mask: Optional[torch.Tensor] = None,414 inputs_embeds: Optional[torch.FloatTensor] = None,415 output_hidden_states: Optional[bool] = None,416 output_attentions: Optional[bool] = None,417 return_dict: Optional[bool] = None,418 ) -> Union[Tuple, BaseModelOutputWithPooling]:419 output_hidden_states = output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states420 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions421 return_dict = return_dict if return_dict is not None else self.config.use_return_dict422 423 if (input_ids is None) == (inputs_embeds is None):424 raise ValueError("Specify exactly one of input_ids or inputs_embeds.")425 input_shape = (426 input_ids.shape427 if input_ids is not None428 else inputs_embeds.shape[:2]429 )430 input_device = (431 input_ids.device432 if input_ids is not None433 else inputs_embeds.device434 )435 if attention_mask is None:436 attention_mask = torch.ones(input_shape, device=input_device)437 key_padding_mask = attention_mask.eq(0)438 if not key_padding_mask.any():439 key_padding_mask = None440 head_mask = self.get_head_mask(441 head_mask,442 self.config.num_hidden_layers,443 )444 if all(mask is None for mask in head_mask):445 head_mask = None446 447 x = self.embeddings(448 input_ids=input_ids,449 token_type_ids=token_type_ids,450 position_ids=position_ids,451 inputs_embeds=inputs_embeds,452 )453 last_hidden_state, all_hidden_states, all_attentions = self.encoder(454 x, key_padding_mask,455 output_hidden_states=output_hidden_states,456 output_attentions=output_attentions,457 head_mask=head_mask,458 )459 pooled = self.pooler(last_hidden_state) if self.pooler is not None else None460 461 if not return_dict:462 output = (last_hidden_state, pooled)463 if output_hidden_states:464 output += (all_hidden_states,)465 if output_attentions:466 output += (all_attentions,)467 return output468 469 return BaseModelOutputWithPooling(470 last_hidden_state=last_hidden_state,471 pooler_output=pooled,472 hidden_states=all_hidden_states,473 attentions=all_attentions,474 )475 476 477class BertForMaskedLM(PreTrainedModel):478 config_class = BertUpdatedConfig479 base_model_prefix = "bert"480 _supports_sdpa = True481 _supports_flash_attn_2 = True482 supports_gradient_checkpointing = True483 _keys_to_ignore_on_load_unexpected = [r"bert\.pooler\."]484 485 def __init__(self, config):486 super().__init__(config)487 self.bert = BertModel(config, add_pooling_layer=False)488 self.transform = BertPredictionHeadTransform(config)489 self.cls = nn.Linear(config.hidden_size, config.vocab_size)490 self.post_init()491 492 def get_input_embeddings(self):493 return self.bert.embeddings.word_embeddings494 495 def set_input_embeddings(self, value):496 self.bert.embeddings.word_embeddings = value497 498 def get_output_embeddings(self):499 return self.cls500 501 def set_output_embeddings(self, value):502 self.cls = value503 504 def forward(505 self,506 input_ids: Optional[torch.LongTensor] = None,507 attention_mask: Optional[torch.Tensor] = None,508 token_type_ids: Optional[torch.LongTensor] = None,509 position_ids: Optional[torch.LongTensor] = None,510 head_mask: Optional[torch.Tensor] = None,511 inputs_embeds: Optional[torch.FloatTensor] = None,512 labels: Optional[torch.LongTensor] = None,513 output_hidden_states: Optional[bool] = None,514 output_attentions: Optional[bool] = None,515 return_dict: Optional[bool] = None,516 ) -> Union[Tuple, MaskedLMOutput]:517 return_dict = return_dict if return_dict is not None else self.config.use_return_dict518 519 outputs = self.bert(520 input_ids=input_ids,521 attention_mask=attention_mask,522 token_type_ids=token_type_ids,523 position_ids=position_ids,524 head_mask=head_mask,525 inputs_embeds=inputs_embeds,526 output_hidden_states=output_hidden_states,527 output_attentions=output_attentions,528 return_dict=True,529 )530 logits = self.cls(self.transform(outputs.last_hidden_state))531 532 loss = None533 if labels is not None:534 loss = F.cross_entropy(logits.view(-1, self.config.vocab_size), labels.view(-1), ignore_index=-100)535 536 if not return_dict:537 output = (logits,)538 if output_hidden_states:539 output += (outputs.hidden_states,)540 if output_attentions:541 output += (outputs.attentions,)542 return (loss,) + output if loss is not None else output543 544 return MaskedLMOutput(545 loss=loss, logits=logits,546 hidden_states=outputs.hidden_states, attentions=outputs.attentions,547 )548 