frb-chatbot / app.py
jaydatech's picture
Update app.py
adc3697 verified
Raw
History Blame Contribute Delete
8.25 kB
import os
import shutil
import json
# Set custom cache directories to avoid permission issues. New
os.environ["HF_HOME"] = "/tmp/huggingface"
os.makedirs("/tmp/huggingface", exist_ok=True)
os.environ["XDG_CACHE_HOME"] = "/tmp/.cache"
os.makedirs("/tmp/.cache", exist_ok=True)
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from typing import List, Optional
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, AutoConfig
import re
# Hugging Face model config
REPO_NAME = "jaydatech/phi3-finetuned-project"
BASE_MODEL = "microsoft/Phi-3-mini-4k-instruct"
HF_TOKEN = os.getenv("HF_TOKEN") # Load from Render environment variable
# # Optional: Cleanup if corrupted config is detected
# def check_and_cleanup_corrupt_cache(repo_name: str):
# cache_dir = os.environ["HF_HOME"]
# model_dir = os.path.join(cache_dir, f"models--{repo_name.replace('/', '--')}")
# if os.path.exists(model_dir):
# for root, dirs, files in os.walk(model_dir):
# for file in files:
# if file == "config.json":
# path = os.path.join(root, file)
# try:
# with open(path, "r") as f:
# json.load(f)
# except json.JSONDecodeError:
# print(f"Corrupted config file detected at {path}, cleaning up...")
# shutil.rmtree(model_dir, ignore_errors=True)
# return
# check_and_cleanup_corrupt_cache(REPO_NAME)
app = FastAPI()
# Enable CORS for frontend access
app.add_middleware(
CORSMiddleware,
allow_origins=["*"], # Or specify your frontend domain
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# Hugging Face model config
# REPO_NAME = "jaydatech/phi3-finetuned-project"
# BASE_MODEL = "microsoft/Phi-3-mini-4k-instruct"
# HF_TOKEN = os.getenv("HF_TOKEN") # Load from Render environment variable
config = AutoConfig.from_pretrained(REPO_NAME, token=HF_TOKEN)
device = "cuda" if torch.cuda.is_available() else "cpu"
tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(
REPO_NAME,
config = config,
token=HF_TOKEN,
torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,
device_map="auto"
)
# Message and request models
class ChatMessage(BaseModel):
role: str
text: str
class ChatRequest(BaseModel):
message: str
history: Optional[List[ChatMessage]] = []
def is_farewell(message: str) -> bool:
farewells = ["bye", "goodbye", "see you", "farewell", "exit", "quit", "end"]
message_lower = message.lower().strip()
return any(re.search(rf"\b{re.escape(f)}\b", message_lower) for f in farewells)
# ============================= Recognizes "attend" but with "end" in it... Fix is above ^ =============================
# farewells = ["bye", "goodbye", "see you", "farewell", "exit", "quit", "end"]
# message_lower = message.lower().strip()
# return any(farewell in message_lower for farewell in farewells)
def clean_input_text(text: str) -> str:
# Remove any "Instruction N: ..." or similar phrases
return re.sub(r"Instruction\s*\d+\s*\(.*?\):", "", text, flags=re.IGNORECASE)
@app.post("/chat")
async def chat(request: ChatRequest):
try:
history = request.history
user_message = request.message
if is_farewell(user_message):
return {
"response": "Goodbye! Feel free to chat again if you have more questions.",
"terminate": True
}
conversation = (
"<|system|>\nYou are an AI assistant for the Federal Reserve Bank of St. Louis. "
"Answer questions based ONLY on your knowledge of the Federal Reserve Bank of St. Louis. "
"If the answer is NOT in the training data, respond with: 'I don't think this information is available. Maybe rephrase for me!'. "
"Answer ONLY what the user asks. Do not volunteer information unless specifically requested. "
"Do NOT ask: 'How can I assist you today?' or 'What can I do for you today?' after every response you give. "
"Provide concise answers to the exact question asked and nothing more.\n"
)
seen_user_messages = set()
seen_model_messages = set()
for msg in history:
if msg.role == "user":
cleaned_user_text = clean_input_text(msg.text.strip())
if cleaned_user_text and cleaned_user_text not in seen_user_messages:
conversation += f"<|user|>\n{cleaned_user_text}\n"
seen_user_messages.add(cleaned_user_text)
elif msg.role == "model":
cleaned_model_text = clean_input_text(msg.text.strip())
if cleaned_model_text and cleaned_model_text not in seen_model_messages:
conversation += f"<|assistant|>\n{cleaned_model_text}\n"
seen_model_messages.add(cleaned_model_text)
# seen_messages = set()
# for msg in history:
# if msg.role == "user" and msg.text.strip() not in seen_messages:
# cleaned_user_text = clean_input_text(msg.text.strip())
# conversation += f"<|user|>\n{cleaned_user_text}\n"
# seen_messages.add(cleaned_user_text)
# elif msg.role == "model":
# cleaned_model_text = clean_input_text(msg.text.strip())
# conversation += f"<|assistant|>\n{cleaned_model_text}\n"
#conversation += f"<|user|>\n{user_message.strip()}\n<|assistant|>"
conversation += f"<|user|>\n{clean_input_text(user_message.strip())}\n<|assistant|>"
inputs = tokenizer(conversation, return_tensors="pt", padding=True, truncation=True, max_length=4096).to(device)
with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=130,
do_sample=True,
temperature=0.1,
top_k=5,
pad_token_id=tokenizer.eos_token_id
)
full_response = tokenizer.decode(outputs[0], skip_special_tokens=False)
assistant_response = ""
if "<|assistant|>" in full_response:
assistant_sections = full_response.split("<|assistant|>")
for section in reversed(assistant_sections):
cleaned = section.strip()
if cleaned:
cleaned = re.split(
r"(<\|user\|>|<\|system\|>|<\|assistant\|>|\nuser[:\s]|<\|endoftext\|>)",
cleaned
)[0]
cleaned = re.sub(r"\n?(User|Assistant)\s*[::\-–]\s*.*", "", cleaned, flags=re.IGNORECASE).strip()
assistant_response = cleaned
break
if not assistant_response:
assistant_response = "⚠️ Sorry, I couldn't generate a response."
assistant_response = re.sub(r'\*\* Instruction \*\*:.*?(?=\n\n|\n$|$)', '', assistant_response, flags=re.DOTALL)
assistant_response = re.sub(r'\*\* Instruction \*\*.*?(?=\n\n|\n$|$)', '', assistant_response, flags=re.DOTALL)
assistant_response = re.sub(r'\n{3,}', '\n\n', assistant_response).strip()
assistant_response = re.sub(r"(How can I assist you today\?|What else can I help you with\?|How can I help you today\?)", "", assistant_response, flags=re.IGNORECASE).strip()
def remove_repeated_sentences(response):
sentences = response.split(". ")
seen = set()
cleaned = []
for sentence in sentences:
if sentence not in seen:
cleaned.append(sentence)
seen.add(sentence)
return ". ".join(cleaned)
assistant_response = remove_repeated_sentences(assistant_response)
return {"response": assistant_response}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))