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)