snt-classifier / modeling_snt.py
lkvantaliani's picture
SNT classifier v0.5.1
8b26267 verified
Raw
History Blame Contribute Delete
4.12 kB
"""SNT news classifier — HF-native wrapper (uploaded to the HF repo as-is).
Usage:
from transformers import AutoModel, AutoTokenizer
model = AutoModel.from_pretrained("sweenk/snt-classifier", trust_remote_code=True)
tok = AutoTokenizer.from_pretrained("sweenk/snt-classifier")
enc = tok("Title. Body...", return_tensors="pt", truncation=True, max_length=512)
labels = model.predict_labels(**enc)
"""
from __future__ import annotations
import torch
import torch.nn as nn
from transformers import AutoConfig, AutoModel, PretrainedConfig, PreTrainedModel
class SNTConfig(PretrainedConfig):
model_type = "snt_classifier"
def __init__(
self,
encoder_name: str = "xlm-roberta-large",
l1_keys: list[str] | None = None,
l2_keys: list[str] | None = None,
l2_parent: dict[str, str] | None = None,
l1_thresholds: dict[str, float] | None = None,
l2_thresholds: dict[str, float] | None = None,
snt_version: str = "v0.5.1",
dropout: float = 0.1,
**kwargs,
):
self.encoder_name = encoder_name
self.l1_keys = l1_keys or []
self.l2_keys = l2_keys or []
self.l2_parent = l2_parent or {}
self.l1_thresholds = l1_thresholds or {}
self.l2_thresholds = l2_thresholds or {}
self.snt_version = snt_version
self.dropout = dropout
super().__init__(**kwargs)
@property
def n_l1(self) -> int:
return len(self.l1_keys)
@property
def n_l2(self) -> int:
return len(self.l2_keys)
class SNTForNewsClassification(PreTrainedModel):
config_class = SNTConfig
def __init__(self, config: SNTConfig):
super().__init__(config)
# Attribute names MUST match DualHeadModel so state_dicts load 1:1.
# from_config (not from_pretrained): weights come from this repo's
# safetensors; from_pretrained breaks under HF's meta-device loading.
self.encoder = AutoModel.from_config(AutoConfig.from_pretrained(config.encoder_name))
hidden = self.encoder.config.hidden_size
self.dropout = nn.Dropout(config.dropout)
self.head_top = nn.Linear(hidden, config.n_l1)
self.head_sub = nn.Linear(hidden, config.n_l2)
self.post_init()
def forward(self, input_ids, attention_mask, **kwargs):
out = self.encoder(input_ids=input_ids, attention_mask=attention_mask)
pooled = self.dropout(out.last_hidden_state[:, 0, :])
return {"l1_logits": self.head_top(pooled), "l2_logits": self.head_sub(pooled)}
@torch.no_grad()
def predict_labels(self, input_ids, attention_mask, **kwargs) -> list[dict]:
"""Thresholded multi-label prediction, one dict per batch row.
Logits are upcast to fp32 before sigmoid — bf16 sigmoid saturates to
exactly 1.0 above logit ~6.2, collapsing co-confident categories.
"""
out = self.forward(input_ids, attention_mask)
l1_probs = torch.sigmoid(out["l1_logits"].float())
l2_probs = torch.sigmoid(out["l2_logits"].float())
results = []
for row in range(l1_probs.shape[0]):
l1 = sorted(
(
{"key": k, "p": round(float(p), 4)}
for k, p in zip(self.config.l1_keys, l1_probs[row].tolist())
if p >= self.config.l1_thresholds.get(k, 0.5)
),
key=lambda hit: -hit["p"],
)
if not l1: # argmax fallback — never return unlabeled
idx = int(l1_probs[row].argmax())
l1 = [{"key": self.config.l1_keys[idx], "p": round(float(l1_probs[row][idx]), 4)}]
l2 = sorted(
(
{"key": k, "p": round(float(p), 4)}
for k, p in zip(self.config.l2_keys, l2_probs[row].tolist())
if p >= self.config.l2_thresholds.get(k, 0.5)
),
key=lambda hit: -hit["p"],
)
results.append({"l1": l1, "primary_l1": l1[0]["key"], "l2": l2})
return results