import torch.nn as nn import torch.nn.functional as F class SimpleClassifier(nn.Module): def __init__(self, embedding_dim=768): super().__init__() self.layernorm = nn.LayerNorm(embedding_dim) self.linear = nn.Linear(embedding_dim, 1, bias=True) def forward(self, x): return F.sigmoid(self.linear(self.layernorm(x)))