Spaces:
Sleeping
Sleeping
File size: 1,801 Bytes
636622a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 | """
Core - Zero-Shot Topic Classification Module
Classification zero-shot avec BART-large-MNLI.
Catégorise un texte dans des labels arbitraires sans entraînement spécifique.
"""
from typing import Any
from transformers import pipeline
# ── Chargement lazy du modèle ──
_zs_pipeline = None
def _get_pipeline():
"""Charge le pipeline zero-shot en lazy loading."""
global _zs_pipeline
if _zs_pipeline is None:
_zs_pipeline = pipeline(
"zero-shot-classification",
model="facebook/bart-large-mnli",
)
return _zs_pipeline
def classify_zero_shot(text: str, categories: str) -> dict[str, Any]:
"""
Classifie un texte dans des catégories arbitraires.
Args:
text: Le texte à classifier.
categories: Catégories séparées par des virgules
(e.g., "Politics, Tech, Sports").
Returns:
Dict avec 'labels' et 'scores' (list), 'top_label' et 'top_score'.
"""
pipe = _get_pipeline()
# Parsing des catégories
candidate_labels = [c.strip() for c in categories.split(",") if c.strip()]
if not candidate_labels:
return {
"labels": [],
"scores": [],
"top_label": "N/A",
"top_score": 0.0,
}
# Classification
result = pipe(text, candidate_labels)
# Formater les scores
scores_formatted = [
{"label": label, "score": round(score * 100, 2)}
for label, score in zip(result["labels"], result["scores"])
]
return {
"labels": result["labels"],
"scores": [round(s * 100, 2) for s in result["scores"]],
"results": scores_formatted,
"top_label": result["labels"][0],
"top_score": round(result["scores"][0] * 100, 2),
}
|