multitask-toxicity-comments / modeling_multitask_toxicity.py
pogram1st's picture
Upload folder using huggingface_hub
78b8ecc verified
Raw
History Blame Contribute Delete
1.35 kB
import torch
import torch.nn as nn
from transformers import PreTrainedModel, AutoModel, AutoConfig
from .configuration_multitask_toxicity import MultiTaskToxicityConfig
class MultiTaskToxicityEncoder(PreTrainedModel):
config_class = MultiTaskToxicityConfig
def __init__(self, config):
super().__init__(config)
self.config = config
base_config = AutoConfig.from_pretrained(config.base_model_name)
self.encoder = AutoModel.from_config(base_config)
hidden_size = base_config.hidden_size
self.dropout = nn.Dropout(config.dropout_prob)
self.profanity_head = nn.Linear(hidden_size, 1)
self.threat_head = nn.Linear(hidden_size, 1)
self.illegal_head = nn.Linear(hidden_size, 1)
self.post_init()
def forward(self, input_ids, attention_mask=None, **kwargs):
outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask)
cls_embedding = outputs.last_hidden_state[:, 0, :]
cls_embedding = self.dropout(cls_embedding)
profanity_logit = self.profanity_head(cls_embedding)
threat_logit = self.threat_head(cls_embedding)
illegal_logit = self.illegal_head(cls_embedding)
return profanity_logit, threat_logit, illegal_logit