Download src/model.py from Deku21/RegFM: direct link, hf CLI and curl.
- Browser
- Download file 18.2 kB
-
https://huggingface.co/Deku21/RegFM/resolve/main/src/model.py
- Command line
-
hf download hf://Deku21/RegFM/src/model.py
-
curl -L -o model.py https://huggingface.co/Deku21/RegFM/resolve/main/src/model.py
18.2 kB
| import os | |
| import copy | |
| import math | |
| from copy import deepcopy | |
| import torch | |
| from torch import nn | |
| import torch.nn.functional as F | |
| from transformers.modeling_bert import * | |
| from huggingface_hub import hf_hub_download | |
| from custom_config import LongBERTConfig | |
| from module import CrossAttention, TransContextModel | |
| def clone(module, N): | |
| return nn.ModuleList([copy.deepcopy(module) for _ in range(N)]) | |
| class LongBERTOutput: | |
| def __init__(self): | |
| self.last_hidden_state = None | |
| self.pooled_output = None | |
| self.hidden_states = None | |
| self.last_attn_weights = None | |
| def dilated_attention(query, key, value, key_padding_mask=None, attn_mask=None, dropout=None, training=None): | |
| assert len(query.shape) == 5, "query must be (batch, n_head, n_seg, seg_len, dim)" | |
| assert query.shape == key.shape == value.shape, "q/k/v shape mismatch" | |
| batch_size, n_head, n_seg, seg_len, dim = query.shape | |
| scores = torch.matmul(query, key.transpose(-1, -2)) / math.sqrt(dim) | |
| if key_padding_mask is not None: | |
| key_padding_mask = key_padding_mask.unsqueeze(1).unsqueeze(3) | |
| if key_padding_mask.dtype != torch.bool: | |
| key_padding_mask = key_padding_mask.bool() | |
| scores = scores.masked_fill(key_padding_mask, float("-inf")) | |
| if attn_mask is not None: | |
| if len(attn_mask.shape) == 3: | |
| attn_mask = attn_mask.view(1, 1, n_seg, seg_len, seg_len) | |
| elif len(attn_mask.shape) == 4: | |
| attn_mask = attn_mask.view(batch_size, 1, n_seg, seg_len, seg_len) | |
| if attn_mask.dtype != torch.bool: | |
| attn_mask = attn_mask.bool() | |
| scores = scores.masked_fill(attn_mask, float("-inf")) | |
| attn_weights = F.softmax(scores, dim=-1) | |
| attn_weights = torch.nan_to_num(attn_weights) | |
| if dropout is not None and training: | |
| attn_weights = F.dropout(attn_weights, p=dropout, training=training) | |
| attn_output = torch.matmul(attn_weights, value) | |
| return attn_output, attn_weights, None | |
| class DilatedMultiheadAttention(nn.Module): | |
| def __init__(self, embedding_dim, n_head, segment_size, dilated_rate, dropout=0.1): | |
| super(DilatedMultiheadAttention, self).__init__() | |
| assert embedding_dim % n_head == 0, "The embedding dimension should be divisible by the number of heads" | |
| assert len(segment_size) == len(dilated_rate), "segment_size and dilated_rate should have the same length" | |
| self.d_proj = embedding_dim // n_head | |
| self.n_head = n_head | |
| self.segment_size = segment_size | |
| self.dilated_rate = dilated_rate | |
| self.dropout = dropout | |
| self.q_proj = nn.Linear(embedding_dim, embedding_dim, bias=False) | |
| self.k_proj = nn.Linear(embedding_dim, embedding_dim, bias=False) | |
| self.v_proj = nn.Linear(embedding_dim, embedding_dim, bias=False) | |
| def forward(self, query, key, value, key_padding_mask=None, attn_mask=None): | |
| batch_size, seq_len, embedding_dim = query.shape | |
| attn_output = torch.zeros_like(query) | |
| for seg_size, dil_rate in zip(self.segment_size, self.dilated_rate): | |
| pad_len = (seg_size - seq_len % seg_size) % seg_size | |
| _seq_len = seq_len + pad_len | |
| if pad_len > 0: | |
| pad = torch.zeros(batch_size, pad_len, embedding_dim, device=query.device, dtype=query.dtype) | |
| _query = torch.cat([query, pad], dim=1) | |
| _key = torch.cat([key, pad], dim=1) | |
| _value = torch.cat([value, pad], dim=1) | |
| if key_padding_mask is not None: | |
| pad_mask = torch.ones( | |
| batch_size, | |
| pad_len, | |
| device=key_padding_mask.device, | |
| dtype=key_padding_mask.dtype, | |
| ) | |
| _key_padding_mask = torch.cat([key_padding_mask, pad_mask], dim=1) | |
| else: | |
| _key_padding_mask = None | |
| else: | |
| _query, _key, _value = query, key, value | |
| _key_padding_mask = key_padding_mask | |
| n_segment = _seq_len // seg_size | |
| _query = _query.view(batch_size, n_segment, seg_size, embedding_dim) | |
| _key = _key.view(batch_size, n_segment, seg_size, embedding_dim) | |
| _value = _value.view(batch_size, n_segment, seg_size, embedding_dim) | |
| _query = _query[:, :, ::dil_rate, :] | |
| _key = _key[:, :, ::dil_rate, :] | |
| _value = _value[:, :, ::dil_rate, :] | |
| dil_seg_len = _query.shape[2] | |
| _query = self.q_proj(_query) | |
| _key = self.k_proj(_key) | |
| _value = self.v_proj(_value) | |
| _query = _query.reshape(batch_size, n_segment * dil_seg_len, self.n_head, self.d_proj) | |
| _key = _key.reshape(batch_size, n_segment * dil_seg_len, self.n_head, self.d_proj) | |
| _value = _value.reshape(batch_size, n_segment * dil_seg_len, self.n_head, self.d_proj) | |
| _query_flat = _query.permute(0, 2, 1, 3) | |
| _key_flat = _key.permute(0, 2, 1, 3) | |
| _value_flat = _value.permute(0, 2, 1, 3) | |
| cls_q = _query_flat[:, :, 0:1, :] | |
| cls_scores = torch.matmul(cls_q, _key_flat.transpose(-2, -1)) / (self.d_proj ** 0.5) | |
| if _key_padding_mask is not None: | |
| cls_key_padding_mask = _key_padding_mask.view(batch_size, n_segment, seg_size)[:, :, ::dil_rate] | |
| cls_key_padding_mask = cls_key_padding_mask.reshape(batch_size, 1, 1, n_segment * dil_seg_len) | |
| cls_scores = cls_scores.masked_fill(cls_key_padding_mask.bool(), float("-inf")) | |
| cls_attn = torch.softmax(cls_scores, dim=-1) | |
| cls_attn = torch.dropout(cls_attn, p=self.dropout, train=self.training) | |
| cls_global_out = torch.matmul(cls_attn, _value_flat) | |
| _query = _query_flat.view(batch_size, self.n_head, n_segment, dil_seg_len, self.d_proj) | |
| _key = _key_flat.view(batch_size, self.n_head, n_segment, dil_seg_len, self.d_proj) | |
| _value = _value_flat.view(batch_size, self.n_head, n_segment, dil_seg_len, self.d_proj) | |
| if _key_padding_mask is not None: | |
| _key_padding_mask = _key_padding_mask.view(batch_size, n_segment, seg_size)[:, :, ::dil_rate] | |
| _attn_out, _, _ = dilated_attention( | |
| _query, | |
| _key, | |
| _value, | |
| key_padding_mask=_key_padding_mask, | |
| attn_mask=attn_mask, | |
| dropout=self.dropout, | |
| training=self.training, | |
| ) | |
| attn_out_resized = torch.zeros( | |
| batch_size, | |
| n_segment, | |
| seg_size, | |
| self.n_head, | |
| self.d_proj, | |
| device=_attn_out.device, | |
| dtype=_attn_out.dtype, | |
| ) | |
| attn_out_resized[:, :, ::dil_rate, :, :] = _attn_out.permute(0, 2, 3, 1, 4) | |
| attn_out_resized[:, 0, 0, :, :] = attn_out_resized[:, 0, 0, :, :] + cls_global_out.squeeze(2) | |
| attn_out_flat = attn_out_resized.reshape(batch_size, n_segment, seg_size, embedding_dim) | |
| attn_out_seq = attn_out_flat.reshape(batch_size, _seq_len, embedding_dim) | |
| if pad_len > 0: | |
| attn_out_seq = attn_out_seq[:, :seq_len, :] | |
| attn_output += attn_out_seq / len(self.segment_size) | |
| return attn_output | |
| class LongBERTEmbeddings(nn.Module): | |
| def __init__(self, config): | |
| super(LongBERTEmbeddings, self).__init__() | |
| self.config = config | |
| self.word_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, padding_idx=3) | |
| self.position_embeddings = nn.Embedding(config.max_position_embeddings, config.hidden_size) | |
| self.token_type_embeddings = nn.Embedding(2, config.hidden_size) | |
| self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=1e-12, elementwise_affine=True) | |
| self.dropout = nn.Dropout(p=config.hidden_dropout_prob) | |
| def forward(self, input_ids, token_type_ids, position_ids): | |
| word_embeddings = self.word_embeddings(input_ids) | |
| position_embeddings = self.position_embeddings(position_ids) | |
| if token_type_ids is not None: | |
| token_type_embeddings = self.token_type_embeddings(token_type_ids) | |
| embeddings = word_embeddings + token_type_embeddings + position_embeddings | |
| else: | |
| embeddings = word_embeddings + position_embeddings | |
| embeddings = self.LayerNorm(embeddings) | |
| return self.dropout(embeddings) | |
| class LongBERTLayer(nn.Module): | |
| def __init__(self, config): | |
| super(LongBERTLayer, self).__init__() | |
| self.attention = DilatedMultiheadAttention( | |
| config.hidden_size, | |
| config.num_attention_heads, | |
| config.segment_size, | |
| config.dilated_rate, | |
| dropout=config.attention_probs_dropout_prob, | |
| ) | |
| self.linear1 = nn.Linear(config.hidden_size, config.hidden_size) | |
| self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=1e-12) | |
| self.dropout = nn.Dropout(p=config.hidden_dropout_prob) | |
| def forward(self, query, key, value, key_padding_mask=None, attn_mask=None): | |
| residual = query | |
| attn_output = self.attention(query, key, value, key_padding_mask=key_padding_mask, attn_mask=attn_mask) | |
| attn_output = residual + self.dropout(attn_output) | |
| attn_output = self.LayerNorm(attn_output) | |
| ffn_output = self.linear1(attn_output) | |
| ffn_output = residual + self.dropout(ffn_output) | |
| ffn_output = self.LayerNorm(ffn_output) | |
| return ffn_output | |
| class LongBERTPooler(nn.Module): | |
| def __init__(self, config): | |
| super(LongBERTPooler, self).__init__() | |
| self.dense = nn.Linear(config.hidden_size, config.hidden_size, bias=True) | |
| self.activation = nn.Tanh() | |
| def forward(self, hidden_state): | |
| return self.activation(self.dense(hidden_state[:, 0, :])) | |
| class LongBERTEncoder(nn.Module): | |
| def __init__(self, config): | |
| super(LongBERTEncoder, self).__init__() | |
| config.attention_probs_dropout_prob = 0.1 | |
| self.layer = clone(LongBERTLayer(config), config.num_hidden_layers) | |
| self.pooler = LongBERTPooler(config) | |
| self.longbert_output = LongBERTOutput() | |
| def forward(self, hidden_state, attention_mask=None, output_hidden_states=False): | |
| key_padding_mask = ~attention_mask.bool() if attention_mask is not None else None | |
| hidden_states = tuple() | |
| for layer in self.layer: | |
| hidden_state = layer(hidden_state, hidden_state, hidden_state, key_padding_mask=key_padding_mask) | |
| if output_hidden_states: | |
| hidden_states = hidden_states + (hidden_state,) | |
| self.longbert_output.pooled_output = self.pooler(hidden_state) | |
| self.longbert_output.last_hidden_state = hidden_state | |
| if output_hidden_states: | |
| self.longbert_output.hidden_states = hidden_states | |
| return self.longbert_output | |
| class LongBERTModel(nn.Module): | |
| def __init__(self, config=None): | |
| super(LongBERTModel, self).__init__() | |
| self.config = config | |
| self.embeddings = LongBERTEmbeddings(config) if config is not None else None | |
| self.encoder = LongBERTEncoder(config) if config is not None else None | |
| def from_config(cls, config): | |
| return cls(config=config) | |
| def from_pretrained(cls, ckpt, version="v2"): | |
| model_ckpt = hf_hub_download(repo_id=ckpt, filename=f"pytorch_model_{version}.bin") | |
| model_config = LongBERTConfig.from_pretrained(ckpt) | |
| model = cls(config=model_config) | |
| model.load_state_dict(torch.load(model_ckpt, map_location="cpu")) | |
| return model | |
| def save_pretrained(self, path): | |
| os.makedirs(path, exist_ok=True) | |
| torch.save(self.state_dict(), os.path.join(path, "pytorch_model.bin")) | |
| def forward(self, input_ids, attention_mask=None, token_type_ids=None, output_hidden_states=False): | |
| batch_size, seq_len = input_ids.size() | |
| position_ids = torch.arange(seq_len, dtype=torch.long, device=input_ids.device) | |
| position_ids = position_ids.unsqueeze(0).repeat(batch_size, 1) | |
| hidden_state = self.embeddings(input_ids, token_type_ids, position_ids) | |
| return self.encoder(hidden_state, attention_mask=attention_mask, output_hidden_states=output_hidden_states) | |
| class CisDNATrans(BertPreTrainedModel): | |
| def __init__(self, config): | |
| super().__init__(config) | |
| config.vocab_size = 150000 | |
| config.max_position_embeddings = 71680 | |
| config.intermediate_size = 3072 | |
| config.num_hidden_layers = 6 | |
| config.segment_size = [128, 512, 1024, 2048] | |
| config.dilated_rate = [16, 64, 256, 512] | |
| self.bert = LongBERTModel(config) | |
| self.cls = BertOnlyMLMHead(config) | |
| self.init_weights() | |
| def get_output_embeddings(self): | |
| return self.cls.predictions.decoder | |
| 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 forward( | |
| self, | |
| input_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, | |
| ): | |
| outputs = self.bert( | |
| input_ids, | |
| attention_mask=attention_mask, | |
| token_type_ids=token_type_ids, | |
| ) | |
| sequence_output = outputs.last_hidden_state | |
| if masked_lm_labels is not None: | |
| mask = masked_lm_labels != -100 | |
| selected_prediction_scores = self.cls(sequence_output[mask]) | |
| selected_labels = masked_lm_labels[mask] | |
| loss_fct = CrossEntropyLoss() | |
| masked_lm_loss = loss_fct(selected_prediction_scores, selected_labels) | |
| outputs = (masked_lm_loss,) | |
| return outputs | |
| class RegFM(BertPreTrainedModel): | |
| def __init__(self, epi_config): | |
| super().__init__(epi_config) | |
| config = deepcopy(epi_config) | |
| num_cross_attentions = 4 | |
| config.vocab_size = 2108 | |
| config.max_position_embeddings = 2112 | |
| dna_config = deepcopy(epi_config) | |
| dna_config.vocab_size = 150000 | |
| dna_config.max_position_embeddings = 71680 | |
| dna_config.intermediate_size = 3072 | |
| dna_config.num_hidden_layers = 6 | |
| dna_config.segment_size = [128, 512, 1024, 2048] | |
| dna_config.dilated_rate = [16, 64, 256, 512] | |
| config.attention_mode = "sparse" | |
| epi_config.attention_mode = "sparse" | |
| self.tf_bert = TransContextModel(config, epi_config) | |
| self.dna_bert = LongBERTModel(dna_config) | |
| self.cross_attentions = nn.ModuleList( | |
| [CrossAttention(dna_config.hidden_size, 2) for _ in range(num_cross_attentions)] | |
| ) | |
| self.dropout = nn.Dropout(config.hidden_dropout_prob) | |
| self.dropout2 = nn.Dropout(config.hidden_dropout_prob) | |
| self.predictor = nn.Linear(config.hidden_size, 1) | |
| self.relu = nn.LeakyReLU(negative_slope=0.01) | |
| 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 forward( | |
| self, | |
| input_ids=None, | |
| trans_ids=None, | |
| dna_ids=None, | |
| attention_mask=None, | |
| dna_attention_mask=None, | |
| token_type_ids=None, | |
| position_ids=None, | |
| head_mask=None, | |
| inputs_embeds=None, | |
| labels=None, | |
| ): | |
| outputs_tf = self.tf_bert( | |
| input_ids=input_ids, | |
| epi_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, | |
| ) | |
| outputs = self.dna_bert( | |
| dna_ids, | |
| attention_mask=dna_attention_mask, | |
| token_type_ids=token_type_ids, | |
| ) | |
| dna_token_output = self.dropout(outputs.last_hidden_state) | |
| tf_token_output = self.dropout2(outputs_tf[0]) | |
| out_attn = None | |
| for cross_attention in self.cross_attentions: | |
| query = dna_token_output.permute(1, 0, 2) | |
| key = tf_token_output.permute(1, 0, 2) | |
| value = tf_token_output.permute(1, 0, 2) | |
| mapped_output, attn_weights = cross_attention(query, key, value) | |
| if out_attn is None: | |
| out_attn = attn_weights | |
| else: | |
| out_attn += attn_weights | |
| dna_token_output = mapped_output.permute(1, 0, 2) | |
| cls_token_hidden_state = dna_token_output[:, 0, :] | |
| logits = self.relu(self.predictor(cls_token_hidden_state)) | |
| outputs = (logits, out_attn, outputs.last_hidden_state[:, :30, :], cls_token_hidden_state) | |
| if labels is not None: | |
| loss_fct = MSELoss() | |
| loss = loss_fct(logits.view(-1), labels.view(-1)) | |
| outputs = (loss,) + outputs | |
| return outputs | |