refactoring LLM adapters and service layer

This commit is contained in:
Adam Wilson
2025-06-12 20:25:18 -06:00
parent 0cdce03879
commit 3aca7df000
5 changed files with 43 additions and 67 deletions
@@ -1,18 +1,14 @@
"""
RAG implementation with local Phi-3-mini-4k-instruct-onnx and embeddings
"""
import logging
import sys
# LangChain imports
from langchain.prompts import PromptTemplate
from langchain_core.output_parsers import StrOutputParser
from langchain_core.runnables import RunnablePassthrough
from src.text_generation.adapters.llm.text_generation_model import TextGenerationFoundationModel
from src.text_generation.adapters.llm.abstract_language_model import AbstractLanguageModel
from src.text_generation.adapters.llm.text_generation_foundation_model import TextGenerationFoundationModel
class Phi3LanguageModel:
class LanguageModel(AbstractLanguageModel):
def __init__(self):
logger = logging.getLogger()
@@ -20,9 +16,14 @@ class Phi3LanguageModel:
handler = logging.StreamHandler(sys.stdout)
logger.addHandler(handler)
self.logger = logger
self.configure_model()
self._configure_model()
def configure_model(self):
def _extract_assistant_response(self, text):
if "<|assistant|>" in text:
return text.split("<|assistant|>")[-1].strip()
return text
def _configure_model(self):
# Create the LangChain LLM
llm = TextGenerationFoundationModel().build()
@@ -42,21 +43,15 @@ class Phi3LanguageModel:
| prompt
| llm
| StrOutputParser()
| self.extract_assistant_response
| self._extract_assistant_response
)
def extract_assistant_response(self, text):
if "<|assistant|>" in text:
return text.split("<|assistant|>")[-1].strip()
return text
def invoke(self, user_input: str) -> str:
def invoke(self, user_prompt: str) -> str:
try:
# Get response from the chain
response = self.chain.invoke(user_input)
response = self.chain.invoke(user_prompt)
return response
except Exception as e:
self.logger.error(f"Failed: {e}")
return e
raise e
@@ -1,11 +1,6 @@
"""
RAG implementation with local Phi-3-mini-4k-instruct-onnx and embeddings
"""
import logging
import sys
# LangChain imports
from langchain_huggingface import HuggingFaceEmbeddings
from langchain.text_splitter import RecursiveCharacterTextSplitter
from langchain_community.vectorstores import FAISS
@@ -17,10 +12,11 @@ from langchain_core.output_parsers import StrOutputParser
from langchain.chains import RetrievalQA
from langchain.prompts import PromptTemplate
from langchain.schema import Document
from src.text_generation.adapters.llm.text_generation_model import TextGenerationFoundationModel
from src.text_generation.adapters.llm.abstract_language_model import AbstractLanguageModel
from src.text_generation.adapters.llm.text_generation_foundation_model import TextGenerationFoundationModel
class Phi3LanguageModelWithRag:
class Phi3LanguageModelWithRag(AbstractLanguageModel):
def __init__(self):
logger = logging.getLogger()
@@ -131,9 +127,9 @@ class Phi3LanguageModelWithRag:
return raw_answer.strip()
def invoke(self, user_input: str) -> str:
def invoke(self, user_prompt: str) -> str:
context_docs = self.vectorstore.as_retriever(search_kwargs={"k": 3}).invoke(user_input)
context_docs = self.vectorstore.as_retriever(search_kwargs={"k": 3}).invoke(user_prompt)
context = self.format_docs(context_docs)
# PROMPT_TEMPLATE = """<|system|>
@@ -173,14 +169,10 @@ class Phi3LanguageModelWithRag:
chain = prompt | self.llm | StrOutputParser()
raw_answer = chain.invoke({
"context": context,
"question": user_input
"question": user_prompt
})
# Clean up the answer (remove any remaining template artifacts)
assistant_answer = self.parse_assistant_answer(raw_answer)
return {
# "question": user_input,
# "context": context,
"answer": assistant_answer
}
return assistant_answer
@@ -1,15 +1,8 @@
"""
RAG implementation with local Phi-3-mini-4k-instruct-onnx and embeddings
"""
import logging
import os
import sys
# LangChain imports
from langchain_huggingface import HuggingFacePipeline
# HuggingFace and ONNX imports
from optimum.onnxruntime import ORTModelForCausalLM
from transformers import AutoTokenizer, pipeline
@@ -35,7 +28,7 @@ class TextGenerationFoundationModel:
self.logger.debug(f'model_base_dir: {model_base_dir}')
self.logger.debug(f'model_cpu_dir: {model_cpu_dir}')
self.logger.debug(f"Loading Phi-3 model from: {model_path}")
self.logger.debug(f'Loading Phi-3 model from: {model_path}')
# Load the tokenizer and model
tokenizer = AutoTokenizer.from_pretrained(