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 @classmethod def from_config(cls, config): return cls(config=config) @classmethod 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_() @add_start_docstrings_to_callable(BERT_INPUTS_DOCSTRING) 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_() @add_start_docstrings_to_callable(BERT_INPUTS_DOCSTRING) 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