Non-SHADovcy/synthetic-cpp-code-detection
011
1import torch2import torch.nn as nn3from transformers import PreTrainedModel, AutoModel4from .model_config import CustomConfig5 6class LogRegClassifier(nn.Module):7 def __init__(self, transformer_output_dim):8 super(LogRegClassifier, self).__init__()9 self.linear = nn.Linear(transformer_output_dim, 1)10 11 def forward(self, x):12 return torch.sigmoid(self.linear(x))13 14class CombinedModel(PreTrainedModel):15 config_class = CustomConfig16 17 def __init__(self, config):18 super().__init__(config)19 self.transformer = AutoModel.from_pretrained(config.transformer_type)20 self.classifier = LogRegClassifier(config.transformer_output_dim)21 22 def forward(self, input_ids, attention_mask):23 outputs = self.transformer(input_ids=input_ids, attention_mask=attention_mask)24 pooled_output = outputs.last_hidden_state[:, 0, :]25 return self.classifier(pooled_output)26 