Chat-Bot / src /rag_chain.py
SandyBot's picture
Upload 2 files
2df8973 verified
Raw
History Blame
3.03 kB
import os
from youtube_transcript_api import YouTubeTranscriptApi, TranscriptsDisabled
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_community.vectorstores import FAISS
from langchain_core.prompts import PromptTemplate
from langchain_huggingface import HuggingFaceEmbeddings
from langchain_groq import ChatGroq
from langchain_core.output_parsers import StrOutputParser
from langchain_core.runnables import RunnableParallel, RunnablePassthrough, RunnableLambda
VECTORSTORE_DIR = "vectorstores"
EMBEDDING_MODEL_NAME = os.environ["EMBEDDING_MODEL"]
_embedding_model = HuggingFaceEmbeddings(model_name=EMBEDDING_MODEL_NAME)
def load_transcript(video_id: str) -> str:
"""Fetch and flatten a YouTube video's transcript."""
try:
yt_api = YouTubeTranscriptApi()
transcript_list = yt_api.fetch(video_id, languages=["en"])
transcript = " ".join(chunk.text for chunk in transcript_list)
return transcript
except TranscriptsDisabled:
raise ValueError("No captions available for this video.")
def split_transcript(transcript: str):
splitter = RecursiveCharacterTextSplitter(
chunk_size=1000,
chunk_overlap=200,
)
return splitter.create_documents([transcript])
def get_or_build_vectorstore(video_id: str):
"""Load a cached FAISS index for this video, or build + save one if missing."""
path = os.path.join(VECTORSTORE_DIR, video_id)
if os.path.exists(path):
return FAISS.load_local(
path,
_embedding_model,
allow_dangerous_deserialization=True,
)
transcript = load_transcript(video_id)
chunks = split_transcript(transcript)
vectorstore = FAISS.from_documents(chunks, _embedding_model)
os.makedirs(VECTORSTORE_DIR, exist_ok=True)
vectorstore.save_local(path)
return vectorstore
def format_docs(retrieved_docs):
return "\n\n".join(doc.page_content for doc in retrieved_docs)
def get_rag_chain(video_id: str):
"""Builds and returns the full RAG chain for a given YouTube video."""
vectorstore = get_or_build_vectorstore(video_id)
retriever = vectorstore.as_retriever(search_kwargs={"k": 2})
llm = ChatGroq(
model="llama-3.1-8b-instant",
temperature=0,
api_key=os.environ["GROQ_API_KEY"],
)
prompt = PromptTemplate(
template="""
You are a helpful assistant.
Answer ONLY from the provided transcript context.
If the context is insufficient, just say you don't know.
{context}
Question: {question}
""",
input_variables=["context", "question"],
)
parser = StrOutputParser()
parallel_chain = RunnableParallel({
"context": retriever | RunnableLambda(format_docs),
"question": RunnablePassthrough(),
})
main_chain = parallel_chain | prompt | llm | parser
return main_chain