Spaces:
Running on Zero
Running on Zero
Download knowledge.py from will123654/cvchatbot: direct link, hf CLI and curl.
- Browser
- Download file 12.7 kB
-
https://huggingface.co/spaces/will123654/cvchatbot/resolve/main/knowledge.py
- Command line
-
hf download hf://spaces/will123654/cvchatbot/knowledge.py
-
curl -L -o knowledge.py https://huggingface.co/spaces/will123654/cvchatbot/resolve/main/knowledge.py
12.7 kB
| 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}") | |