cvchatbot / knowledge.py
will123654's picture
init
bc300ce
Raw History Blame Contribute Delete
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}")