PeptideBERT / hemo /network.py
Cat-591's picture
Upload 22 files
532bce5 verified
Raw
History Blame Contribute Delete
1.89 kB
import torch
from transformers import BertModel, BertConfig, logging
logging.set_verbosity_error()
class PeptideBERT(torch.nn.Module):
def __init__(self, bert_config):
super(PeptideBERT, self).__init__()
self.protbert = BertModel.from_pretrained(
'Rostlab/prot_bert_bfd',
config=bert_config,
ignore_mismatched_sizes=True
)
self.head = torch.nn.Sequential(
torch.nn.Linear(bert_config.hidden_size, 1),
torch.nn.Sigmoid()
)
def forward(self, inputs, attention_mask):
output = self.protbert(inputs, attention_mask=attention_mask)
return self.head(output.pooler_output)
def create_model(config):
bert_config = BertConfig(
vocab_size=config['vocab_size'],
hidden_size=config['network']['hidden_size'],
num_hidden_layers=config['network']['hidden_layers'],
num_attention_heads=config['network']['attn_heads'],
hidden_dropout_prob=config['network']['dropout']
)
model = PeptideBERT(bert_config).to(config['device'])
return model
def cri_opt_sch(config, model):
criterion = torch.nn.BCELoss()
optimizer = torch.optim.AdamW(model.parameters(), lr=config['optim']['lr'])
if config['sch']['name'] == 'onecycle':
scheduler = torch.optim.lr_scheduler.OneCycleLR(
optimizer,
max_lr=config['optim']['lr'],
epochs=config['epochs'],
steps_per_epoch=config['sch']['steps']
)
elif config['sch']['name'] == 'lronplateau':
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
optimizer,
mode='max',
factor=config['sch']['factor'],
patience=config['sch']['patience']
)
return criterion, optimizer, scheduler