daisy-rrhh-backend / src /agents /classifier_agent.py
octaviofr8hub's picture
add rrhh agent
01b6dbe
Raw
History Blame
3.67 kB
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)