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),
    }