authorship-api / app.py
MSK1202's picture
Create app.py
2ede3a4 verified
Raw History Blame Contribute Delete
1.66 kB
import torch
from fastapi import FastAPI
from pydantic import BaseModel
from huggingface_hub import list_repo_files
MODEL_ID = "peterkirby/modernbert-large-pan2020-authorship-verification"
# Определяем тип модели: bi-encoder (sentence-transformers) или pair classifier
IS_BIENCODER = "modules.json" in list_repo_files(MODEL_ID)
if IS_BIENCODER:
from sentence_transformers import SentenceTransformer, util
model = SentenceTransformer(MODEL_ID)
else:
from transformers import AutoTokenizer, AutoModelForSequenceClassification
tok = AutoTokenizer.from_pretrained(MODEL_ID)
model = AutoModelForSequenceClassification.from_pretrained(MODEL_ID).eval()
class Req(BaseModel):
text: str
text_pair: str
app = FastAPI()
@app.get("/")
def health():
return {"ok": True, "mode": "bi-encoder" if IS_BIENCODER else "pair-classifier"}
@app.post("/verify")
def verify(r: Req):
if IS_BIENCODER:
e = model.encode([r.text, r.text_pair], convert_to_tensor=True,
normalize_embeddings=True)
p = (float(util.cos_sim(e[0], e[1])) + 1) / 2
else:
enc = tok(r.text, r.text_pair, return_tensors="pt",
truncation=True, max_length=512)
with torch.no_grad():
logits = model(**enc).logits[0]
p = float(torch.sigmoid(logits[0])) if logits.numel() == 1 \
else float(torch.softmax(logits, -1)[1])
# тот же формат, что возвращал HF API, чтобы остальной код не менять
return [[{"label": "LABEL_1", "score": p},
{"label": "LABEL_0", "score": 1 - p}]]