Download src/module.py from Deku21/RegFM: direct link, hf CLI and curl.
- Browser
- Download file 4.52 kB
-
https://huggingface.co/Deku21/RegFM/resolve/main/src/module.py
- Command line
-
hf download hf://Deku21/RegFM/src/module.py
-
curl -L -o module.py https://huggingface.co/Deku21/RegFM/resolve/main/src/module.py
4.52 kB
| # coding=utf-8 | |
| """TF encoder components used by RegFM.""" | |
| from copy import deepcopy | |
| import torch | |
| from torch import nn | |
| from torch.nn import CrossEntropyLoss | |
| from transformers.modeling_bert import * | |
| from transformers.modeling_bert import GenomicBertModelNew as TransContextModel | |
| class CrossAttention(nn.Module): | |
| def __init__(self, hidden_size, num_heads=4, dropout=0.1): | |
| super(CrossAttention, self).__init__() | |
| self.attention = nn.MultiheadAttention(embed_dim=hidden_size, num_heads=num_heads, dropout=dropout) | |
| self.linear1 = nn.Linear(hidden_size, hidden_size * 4) | |
| self.linear2 = nn.Linear(hidden_size * 4, hidden_size) | |
| self.norm1 = nn.LayerNorm(hidden_size) | |
| self.norm2 = nn.LayerNorm(hidden_size) | |
| self.dropout = nn.Dropout(dropout) | |
| self.activation = nn.ReLU() | |
| def forward(self, query, key, value, key_padding_mask=None): | |
| attn_output, attn_weights = self.attention(query, key, value, key_padding_mask=key_padding_mask) | |
| query = query + self.dropout(attn_output) | |
| query = self.norm1(query) | |
| ff_output = self.linear2(self.dropout(self.activation(self.linear1(query)))) | |
| output = query + self.dropout(ff_output) | |
| output = self.norm2(output) | |
| return output, attn_weights | |
| class TransContextForMaskedLM(BertPreTrainedModel): | |
| def __init__(self, config): | |
| super().__init__(config) | |
| tf_config = deepcopy(config) | |
| tf_config.vocab_size = 2108 | |
| tf_config.max_position_embeddings = 2112 | |
| self.bert = TransContextModel(tf_config, config) | |
| self.cls = BertOnlyMLMHead(config) | |
| self.vocab_size = config.vocab_size | |
| self.config = config | |
| self.tf_config = tf_config | |
| self.init_weights() | |
| def init_weights(self): | |
| for module_ in self.named_modules(): | |
| if isinstance(module_[1], (torch.nn.Linear, torch.nn.Embedding)): | |
| module_[1].weight.data.normal_(mean=0.0, std=self.config.initializer_range) | |
| elif isinstance(module_[1], torch.nn.LayerNorm): | |
| module_[1].bias.data.zero_() | |
| module_[1].weight.data.fill_(1.0) | |
| if isinstance(module_[1], torch.nn.Linear) and module_[1].bias is not None: | |
| module_[1].bias.data.zero_() | |
| def get_output_embeddings(self): | |
| return self.cls.predictions.decoder | |
| def forward( | |
| self, | |
| input_ids=None, | |
| trans_ids=None, | |
| attention_mask=None, | |
| token_type_ids=None, | |
| position_ids=None, | |
| head_mask=None, | |
| inputs_embeds=None, | |
| masked_lm_labels=None, | |
| encoder_hidden_states=None, | |
| encoder_attention_mask=None, | |
| lm_labels=None, | |
| l2_lambda=0.01, | |
| ): | |
| outputs = self.bert( | |
| input_ids, | |
| trans_ids, | |
| attention_mask=attention_mask, | |
| token_type_ids=token_type_ids, | |
| position_ids=position_ids, | |
| head_mask=head_mask, | |
| inputs_embeds=inputs_embeds, | |
| encoder_hidden_states=encoder_hidden_states, | |
| encoder_attention_mask=encoder_attention_mask, | |
| ) | |
| sequence_output = outputs[0] | |
| prediction_scores = self.cls(sequence_output) | |
| outputs = (prediction_scores, sequence_output) | |
| if masked_lm_labels is not None: | |
| class_counts = torch.bincount(masked_lm_labels[masked_lm_labels != -100], minlength=self.vocab_size) | |
| class_weights = 1.0 / (class_counts.float() + 1e-6) | |
| loss_fct = CrossEntropyLoss(weight=class_weights) | |
| masked_lm_loss = loss_fct(prediction_scores.view(-1, self.vocab_size), masked_lm_labels.view(-1)) | |
| _, predictions = torch.max(prediction_scores, dim=-1) | |
| masked_indices = masked_lm_labels != -100 | |
| masked_predictions = predictions[masked_indices] | |
| masked_labels = masked_lm_labels[masked_indices] | |
| accuracy = (masked_predictions == masked_labels).float().mean().item() | |
| outputs = (masked_lm_loss, accuracy) + outputs | |
| if lm_labels is not None: | |
| prediction_scores = prediction_scores[:, :-1, :].contiguous() | |
| lm_labels = lm_labels[:, 1:].contiguous() | |
| loss_fct = CrossEntropyLoss() | |
| ltr_lm_loss = loss_fct(prediction_scores.view(-1, self.vocab_size), lm_labels.view(-1)) | |
| outputs = (ltr_lm_loss,) + outputs | |
| return outputs | |