SimpliTax / app.py
Wiefdw's picture
Update app.py
6a47450 verified
Raw
History Blame Contribute Delete
3.03 kB
# app.py
import gradio as gr
from huggingface_hub import snapshot_download
from transformers import AutoModelForCausalLM, AutoTokenizer, pipeline, BitsAndBytesConfig
from langchain_community.vectorstores import FAISS
from langchain_community.embeddings import HuggingFaceEmbeddings
import torch
# --- Config ---
MODEL_REPO_ID = "Wiefdw/merged-tax-raft-mistral-7b"
VECTOR_REPO_ID = "Wiefdw/tax-indonesia-vectordb"
EMBEDDING_MODEL_NAME = "sentence-transformers/all-MiniLM-L6-v2"
LOCAL_DB_PATH = "./vector_db_pajak"
# --- Load Vector Database ---
print("πŸ“¦ Downloading vector database...")
snapshot_download(repo_id=VECTOR_REPO_ID, repo_type="dataset", local_dir=LOCAL_DB_PATH)
print("πŸ”Ž Loading embeddings and vectorstore...")
embeddings = HuggingFaceEmbeddings(model_name=EMBEDDING_MODEL_NAME)
vectorstore = FAISS.load_local(LOCAL_DB_PATH, embeddings, allow_dangerous_deserialization=True)
retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
# --- Load LLM ---
print("πŸš€ Loading LLM model...")
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
)
model = AutoModelForCausalLM.from_pretrained(
MODEL_REPO_ID,
quantization_config=bnb_config,
device_map="auto",
torch_dtype=torch.bfloat16,
trust_remote_code=True
)
tokenizer = AutoTokenizer.from_pretrained(MODEL_REPO_ID, trust_remote_code=True)
tokenizer.pad_token = tokenizer.eos_token
llm_pipeline = pipeline(
"text-generation",
model=model,
tokenizer=tokenizer,
max_new_tokens=1024,
temperature=0.6,
top_p=0.9,
do_sample=True,
)
# --- Function to answer ---
def chatbot_fn(message, history):
try:
# Retrieve context
docs = retriever.invoke(message)
context = "\n\n".join([doc.page_content for doc in docs])
# Build prompt
prompt = f"""<s>[INST]
Jawab pertanyaan-pertanyaan berikut HANYA dan SELALU dalam Bahasa Indonesia, berdasarkan konteks yang diberikan.
🎯 Fokuskan jawaban hanya pada POIN-POIN UTAMA yang relevan.
🚫 Jangan bertele-tele.
βœ… Gunakan gaya jawab yang padat, jelas, dan langsung ke inti.
πŸ›‘ Pastikan jawaban selesai sebelum limit token 1024 habis.
Pertanyaan: {message}
Konteks:
{context}
[/INST]"""
# Generate
response = llm_pipeline(prompt)
answer = response[0]["generated_text"].split("[/INST]")[-1].strip().replace("</s>", "")
return answer
except Exception as e:
return f"❌ Error: {str(e)}"
# --- Gradio Chat UI ---
with gr.Blocks(theme=gr.themes.Soft(primary_hue="violet")) as demo:
gr.Markdown("""
# πŸ’¬ **Chatbot Pajak Indonesia**
Diskusi soal perpajakan Indonesia dengan LLM + RAG.
Versi demo di Hugging Face Spaces.
""")
chatbot = gr.ChatInterface(
fn=chatbot_fn,
title="Chatbot Pajak (RAG)",
description="Tanyakan apa saja tentang perpajakan Indonesia, SPT, NPWP, dll.",
theme="soft",
)
if __name__ == "__main__":
demo.launch()