thehat-word-difficulty / modeling.py
nzinov's picture
Word difficulty (E) regressor: rubert-base + frequency feature
9359e92 verified
Raw
History Blame Contribute Delete
4.02 kB
# -*- coding: utf-8 -*-
"""Word-difficulty regressor for the game «Шляпа» (thehat.tech).
Predicts the Elo-style difficulty rating E of a Russian noun from the word
itself plus its corpus frequency. See README.md for metrics and background.
Usage:
from modeling import WordDifficultyPredictor
predictor = WordDifficultyPredictor.from_dir(".")
predictor.predict(["кровать", "соразмерность"]) # -> [~38.5, ~60.0]
"""
import json
import os
import numpy as np
import torch
from transformers import AutoModel, AutoTokenizer
MAX_LEN = 16
class BertFreqRegressor(torch.nn.Module):
def __init__(self, base_model="DeepPavlov/rubert-base-cased"):
super().__init__()
self.bert = AutoModel.from_pretrained(base_model)
hidden = self.bert.config.hidden_size
self.head = torch.nn.Sequential(
torch.nn.Linear(hidden + 2, 128),
torch.nn.GELU(),
torch.nn.Dropout(0.1),
torch.nn.Linear(128, 1),
)
def forward(self, ids, mask, extra):
out = self.bert(input_ids=ids, attention_mask=mask)
m = mask.unsqueeze(-1).float()
pooled = (out.last_hidden_state * m).sum(1) / m.sum(1).clamp(min=1)
return self.head(torch.cat([pooled, extra], dim=-1)).squeeze(-1)
class WordDifficultyPredictor:
def __init__(self, model, tokenizer, freq, mean, std, lf_mean, lf_std,
device=None):
self.model = model
self.tokenizer = tokenizer
self.freq = freq
self.mean, self.std = mean, std
self.lf_mean, self.lf_std = lf_mean, lf_std
self.device = device or torch.device(
"mps" if torch.backends.mps.is_available()
else "cuda" if torch.cuda.is_available() else "cpu")
self.model.to(self.device).eval()
@classmethod
def from_dir(cls, path):
"""Load from a local directory (e.g. a huggingface_hub snapshot)."""
ckpt = torch.load(os.path.join(path, "model.pt"), map_location="cpu",
weights_only=True)
model = BertFreqRegressor(ckpt.get("model_name",
"DeepPavlov/rubert-base-cased"))
model.load_state_dict(ckpt["state_dict"])
tokenizer = AutoTokenizer.from_pretrained(path)
with open(os.path.join(path, "word_frequency.json"),
encoding="utf-8") as fh:
freq = json.load(fh)
return cls(model, tokenizer, freq, ckpt["mean"], ckpt["std"],
ckpt["lf_mean"], ckpt["lf_std"])
@classmethod
def from_hub(cls, repo_id="nzinov/thehat-word-difficulty"):
from huggingface_hub import snapshot_download
return cls.from_dir(snapshot_download(repo_id))
@torch.no_grad()
def predict(self, words):
"""Difficulty ratings E for a list of Russian words (lowercase)."""
words = [w.strip().lower() for w in words]
enc = self.tokenizer(words, truncation=True, max_length=MAX_LEN,
padding="max_length", return_tensors="pt")
feats = []
for w in words:
f = self.freq.get(w)
has = 1.0 if f is not None else 0.0
lf = ((np.log1p(f) - self.lf_mean) / self.lf_std
if f is not None else 0.0)
feats.append([lf, has])
extra = torch.tensor(feats, dtype=torch.float32)
pred = self.model(enc["input_ids"].to(self.device),
enc["attention_mask"].to(self.device),
extra.to(self.device))
return (pred.cpu().numpy() * self.std + self.mean).tolist()
if __name__ == "__main__":
import sys
predictor = WordDifficultyPredictor.from_dir(
os.path.dirname(os.path.abspath(__file__)) or ".")
words = sys.argv[1:] or ["кровать", "луна", "синоним", "соразмерность"]
for word, e in zip(words, predictor.predict(words)):
print(f"{word:20s} E = {e:.1f}")