import json import logging import os import pickle from collections.abc import Callable, Generator from string import Template from threading import Thread from typing import Any, Tuple from gradio import MessageDict from accelerate import Accelerator from sentence_transformers import SentenceTransformer from transformers import TextIteratorStreamer, pipeline import torch from config import EMBEDDING_MODEL_ID, LLM_ID, KNOWLEDGE_CACHE_FILE, \ DOCUMENT_FOLLOWUP_THRESHOLD, DOCUMENT_SELECTION_THRESHOLD, DOMAIN_SELECTION_THRESHOLD, \ RESPONSE_MAX_NEW_TOKENS, RESPONSE_SAMPLING_TEMPERATURE, RESPONSE_SAMPLING_TOP_P from helpers import download_cv_data logger = logging.getLogger(__name__) accelerator = Accelerator() device = accelerator.device class DomainKnowledge: def __init__( self, desc: str, intro_message: str, documents: list[str], document_embeddings: torch.Tensor, ): self.description = desc self.intro_message = intro_message self.documents = documents self.document_embeddings = document_embeddings class DomainKnowledgeConfig: def __init__(self, description, file_name: str, parse_file_fn: Callable[[Any], Tuple[str, list[str]]]): self.desc = description self.file_name = file_name self.parse_file = parse_file_fn def initialize(self, embedding_model: SentenceTransformer) -> DomainKnowledge: with open(self.file_name) as f: document_file = json.load(f) intro_message, documents = self.parse_file(document_file) document_embeddings: torch.Tensor = embedding_model.encode_document( documents, convert_to_tensor=True, device=device, ) # type: ignore return DomainKnowledge(self.desc, intro_message, documents, document_embeddings) class Knowledge: def __init__( self, welcome_message: str, no_match_response: str, reset_query: str, domain_knowledges: list[DomainKnowledge], domain_knowledge_description_embeddings: torch.Tensor, ): self.welcome_message = welcome_message self.no_match_response = no_match_response self.reset_query = reset_query self.domain_knowledges = domain_knowledges self.dk_desc_embeddings = domain_knowledge_description_embeddings class KnowledgeConfig: def __init__( self, welcome_message_template: Template, no_match_response: str, reset_query: str, domain_knowledge_configs: list[DomainKnowledgeConfig], ): domain_descriptions = "\n".join([f"➢ {dk_cfg.desc}" for dk_cfg in domain_knowledge_configs]) self.welcome_message = welcome_message_template.substitute(domains=domain_descriptions) self.no_match_response = no_match_response self.reset_query = reset_query self.domain_knowledge_configs = domain_knowledge_configs def initialize( self, embedding_model: SentenceTransformer, generate_cache: bool, download_data: bool, ) -> Knowledge: if generate_cache or not os.path.exists(KNOWLEDGE_CACHE_FILE): logger.info("⚠️ Initializing cv knowledge...") if download_data: logger.info("⚠️ Downloading cv data...") download_cv_data() domain_knowledges = [] domain_knowledge_descriptions = [] for dk_cfg in self.domain_knowledge_configs: domain_knowledge = dk_cfg.initialize(embedding_model) domain_knowledges.append(domain_knowledge) domain_knowledge_descriptions.append(domain_knowledge.description) domain_knowledge_description_embeddings: torch.Tensor = embedding_model.encode_document( domain_knowledge_descriptions, convert_to_tensor=True, device=device, ) # type: ignore knowledge = Knowledge( self.welcome_message, self.no_match_response, self.reset_query, domain_knowledges, domain_knowledge_description_embeddings, ) with open(KNOWLEDGE_CACHE_FILE, 'wb') as f: pickle.dump(knowledge, f) logger.info(f"✅ Cached cv knowledge to '{KNOWLEDGE_CACHE_FILE}'.") else: logger.info(f"✅ Found cache for cv knowledge. Loading data from '{KNOWLEDGE_CACHE_FILE}'...") with open(KNOWLEDGE_CACHE_FILE, 'rb') as f: knowledge = pickle.load(f) return knowledge class ChatContext: """ Holds the conversational state that uses knowledge base and LLM to answer questions. Context tree: 1. Learn about jobs a. position a, company x b. position b, company x ... 2. Learn about tech stack and skills a. languages b. domain of expertise c. tech stacks 3. Learn about papers a. paper 1 b. paper 2 ... 4. Learn about education """ def __init__(self): self.domain_index = -1 self.document_index = -1 class KnowledgeBase: def __init__( self, knowledge_config: KnowledgeConfig, generate_cache: bool, download_data: bool, ): self.embedding_model = SentenceTransformer(EMBEDDING_MODEL_ID, device='cuda') self.llm_pipeline = pipeline(task="text-generation", model=LLM_ID, device='cuda') self.knowledge = knowledge_config.initialize(self.embedding_model, generate_cache, download_data) def semantic_search_domains(self, query: str) -> int: """symmetric semantic search on domain descriptions""" query_embedding: torch.Tensor = self.embedding_model.encode_query(query, convert_to_tensor=True) # type: ignore similarity_scores = self.embedding_model.similarity(query_embedding, self.knowledge.dk_desc_embeddings) if similarity_scores.numel() == 0: logger.debug("Calculating similarities returned an empty tensor.") return -1 best_index = similarity_scores.argmax() best_similarity = similarity_scores[0, best_index] if best_similarity >= DOMAIN_SELECTION_THRESHOLD: logger.debug(f"domain description index {best_index} with similarity {best_similarity} exceeds threshold {DOMAIN_SELECTION_THRESHOLD}.") return int(best_index) logger.debug(f"best domain description at {best_index} with similarity {best_similarity} is below threshold {DOMAIN_SELECTION_THRESHOLD}.") return -1 def semantic_search_documents(self, query: str, chat_context: ChatContext) -> int: """asymmetric semantic search on documents""" query_embedding: torch.Tensor = self.embedding_model.encode_query(query, convert_to_tensor=True) # type: ignore similarity_threshold = DOCUMENT_SELECTION_THRESHOLD if chat_context.document_index == -1 else DOCUMENT_FOLLOWUP_THRESHOLD document_embeddings = self.knowledge.domain_knowledges[chat_context.domain_index].document_embeddings similarity_scores = self.embedding_model.similarity(query_embedding, document_embeddings) if similarity_scores.numel() == 0: logger.debug("Calculating similarities returned an empty tensor.") return -1 best_index = similarity_scores.argmax() best_similarity = similarity_scores[0, best_index] if best_similarity >= similarity_threshold: logger.debug(f"in domain {chat_context.domain_index}, document at {best_index} with similarity {best_similarity} exceeds threshold {similarity_threshold}.") return int(best_index) logger.debug(f"in domain {chat_context.domain_index}, document at {best_index} with best similarity {best_similarity} is below threshold {similarity_threshold}.") return -1 def get_system_prompt(self, context_str: str) -> str: return f""" \rAnswer the following QUESTION based only on the CONTEXT provided. If the answer cannot be found in the CONTEXT, write \"{self.knowledge.no_match_response}\" \r--- \rCONTEXT: \r{context_str}""" def get_llm_prompt(self, chat_context: ChatContext, user_input: str, history: list[MessageDict]) -> list[MessageDict]: """Returns the prompt to ask LLM to generate response.""" document_content = self.knowledge.domain_knowledges[chat_context.domain_index].documents[chat_context.document_index] llm_prompt = list() llm_prompt.append(MessageDict(role="system", content=self.get_system_prompt(document_content))) llm_prompt.extend(history[1:]) # skip welcome message llm_prompt.append(MessageDict(role="user", content=f"QUESTION:\n{user_input}")) logger.debug("generated llm prompt:") logger.debug(f"system: {llm_prompt[0]["content"]}") logger.debug(f"user: {llm_prompt[-1]["content"]}") return llm_prompt def intro_response(self, chat_context: ChatContext) -> Generator[Tuple[str, ChatContext]]: yield self.knowledge.welcome_message, chat_context def no_match_response(self, chat_context: ChatContext) -> Generator[Tuple[str, ChatContext]]: yield self.knowledge.no_match_response, chat_context def answer_query( self, user_message: str, history: list[MessageDict], chat_context: ChatContext, ) -> Generator[Tuple[str, ChatContext]]: """Generates a response based on the user message and chat context.""" logger.debug(f"user input: {user_message}") if user_message == self.knowledge.reset_query: # reset chat logger.debug("resetting chat") chat_context.domain_index = -1 chat_context.document_index = -1 yield from self.intro_response(chat_context) return elif len(history) == 0: # new convo logger.debug("starting a new chat") yield from self.intro_response(chat_context) return elif chat_context.domain_index == -1: # domain selection chat domain_index = self.semantic_search_domains(user_message) if domain_index >= 0: chat_context.domain_index = domain_index yield self.knowledge.domain_knowledges[domain_index].intro_message, chat_context return else: logger.debug("failed to find matching domain") yield from self.no_match_response(chat_context) return else: # domain chat document_index = self.semantic_search_documents(user_message, chat_context) if document_index >= 0: # found content similar to user's input. use that content. chat_context.document_index = document_index logger.debug(f"using document at index {chat_context.document_index} in domain {chat_context.domain_index} that matches user's input.") elif chat_context.document_index >= 0: # no specific content matching user's input, so use previous document in chat_context. logger.debug(f"no content matches user's input. using previous document at index {chat_context.document_index} in domain {chat_context.domain_index}.") else: # no content found and no previous content. yield no match. logger.debug("no content found and no previous content.") yield from self.no_match_response(chat_context) return # Generate response using LLM llm_prompt = self.get_llm_prompt(chat_context, user_message, history) llm_response_streamer = TextIteratorStreamer(self.llm_pipeline.tokenizer, skip_prompt=True) # type: ignore generation_args = dict( text_inputs=llm_prompt, streamer=llm_response_streamer, skip_special_tokens=True, max_new_tokens=RESPONSE_MAX_NEW_TOKENS, do_sample=True, top_p=RESPONSE_SAMPLING_TOP_P, temperature=RESPONSE_SAMPLING_TEMPERATURE, ) thread = Thread( target=self.llm_pipeline, kwargs=generation_args ) thread.start() response = "" for new_text in llm_response_streamer: response += new_text yield response, chat_context logger.debug(f"[assistant] {response}")