Deep-Studio-Text / core /zero_shot.py
demeulemeesterxmaxime
Ajout des fichiers V1
636622a
Raw
History Blame Contribute Delete
1.8 kB
"""
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),
}