Spaces:
Sleeping
Sleeping
| 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) |