File size: 3,665 Bytes
bc9904d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
70
71
72
73
74
75
76
77
78
79
80
81
82
import json
import numpy as np
import faiss
from sentence_transformers import SentenceTransformer
import logging

class AgentClassifier:
    """
    Clase encargada de clasificar prompts para seleccionar el 
    agente adecuado usando similitud cosenoidal.
    """
    def __init__(self, descriptions_file: str, model_name: str = "all-MiniLM-L6-v2", similarity_threshold: float = 0.5):
        """
        Inicializa el clasificador con descripciones de agentes y un modelo de embeddings.

        Args:
            descriptions_file (str): archivo de descripciones representativas de cada agente
            model_name (str): Nombre del modelo a utilizar
            similarity_threshold (float): Umbral de similitud entre las palabras
        """
        self.logger = logging.getLogger(__name__)
        self.model = SentenceTransformer(model_name)
        self.similarity_threshold = similarity_threshold
        self.agent_descriptions = {}
        self.agent_embeddings = None
        self.index = None
        self.agent_ids = []
        # Cargar descripciones
        try:
            with open(descriptions_file, 'r', encoding='utf-8') as f:
                self.agent_descriptions = json.load(f)
        except Exception as e:
            self.logger.error(f"Error al cargar descripciones: {e}")
            raise
        # Generar y almacenar embeddings
        self._initialize_embeddings()

    def _initialize_embeddings(self):
        """
        Genera embeddings de las descripciones y los almacena en un indice FAISS.

        Args: 
            None
        """
        self.agent_ids = list(self.agent_descriptions.keys())
        descriptions = list(self.agent_descriptions.values())
        self.agent_embeddings = self.model.encode(descriptions, show_progress_bar=False) # Generar embeddings
        # Se crean los indices faiss
        dimension = self.agent_embeddings.shape[1]
        self.index = faiss.IndexFlatIP(dimension)  # Usar Inner Product para similitud cosenoidal
        faiss.normalize_L2(self.agent_embeddings)  # Normalizar para similitud cosenoidal
        self.index.add(self.agent_embeddings)
        self.logger.info("Embeddings de agentes inicializados y almacenados en FAISS.")

    def classify(self, prompt: str) -> str:
        """
        Metodo que clasifica un prompt y devuelve el ID del agente adecuado.

        Args:
            prompt (str): Texto del prompt para su posterior analisis y definir el id del agente adecuado 

        Returns:
            (str): ID del agente a usar despues del analisis
        """
        try:
            prompt_embedding = self.model.encode([prompt], show_progress_bar=False) # Generar embedding del prompt
            faiss.normalize_L2(prompt_embedding)
            scores, indices = self.index.search(prompt_embedding, 1) # Buscar el agente mas similar
            max_score = scores[0][0]
            agent_index = indices[0][0]
            if max_score < self.similarity_threshold: # Verificar umbral de similitud
                self.logger.warning(f"Similitud maxima {max_score} por debajo del umbral {self.similarity_threshold}")
                return None  
            selected_agent = self.agent_ids[agent_index]
            self.logger.info(f"Prompt clasificado: agente seleccionado = {selected_agent}, score = {max_score}")
            return selected_agent
        except Exception as e:
            self.logger.error(f"Error al clasificar prompt: {e}")
            return None

def create_agent(descriptions_file: str, model_name: str = "all-MiniLM-L6-v2", similarity_threshold: float = 0.5):
    return AgentClassifier(descriptions_file=descriptions_file)