rubert-tiny2-multitask-toxicity / modeling_multitask_toxicity.py
aurelianvolturi's picture
Upload folder using huggingface_hub
34b7f15 verified
Raw
History Blame Contribute Delete
1.9 kB
import torch
import torch.nn as nn
from transformers import AutoModel, PreTrainedModel, AutoConfig
from transformers.modeling_outputs import SequenceClassifierOutput
from .configuration_multitask_toxicity import MultiTaskToxicityConfig
class MultiTaskToxicityEncoder(PreTrainedModel):
config_class = MultiTaskToxicityConfig
base_model_prefix = "encoder"
def __init__(self, config):
super().__init__(config)
is_meta = False
import inspect
try:
for frame in inspect.stack():
if 'from_pretrained' in frame.function:
is_meta = True
break
except Exception:
pass
if is_meta:
enc_config = AutoConfig.from_pretrained(config.encoder_name, token=False)
self.encoder = AutoModel.from_config(enc_config)
else:
self.encoder = AutoModel.from_pretrained(config.encoder_name, token=False)
for param in self.encoder.parameters():
param.requires_grad = False
hidden_size = self.encoder.config.hidden_size
self.dropout = nn.Dropout(config.dropout)
self.heads = nn.ModuleDict({
label: nn.Linear(hidden_size, 1) for label in config.labels
})
self.post_init()
def forward(self, input_ids, attention_mask, labels=None, **kwargs):
cls_embedding = self.encoder(
input_ids=input_ids, attention_mask=attention_mask
).last_hidden_state[:, 0]
logits = torch.cat([
self.heads[label](self.dropout(cls_embedding))
for label in self.config.labels
], dim=1)
loss = None
if labels is not None:
loss = nn.functional.binary_cross_entropy_with_logits(logits, labels)
return SequenceClassifierOutput(loss=loss, logits=logits)