File size: 3,027 Bytes
bf2ef62 53b543e bf2ef62 faf85a4 373c574 bf2ef62 5923a7d bf2ef62 faf85a4 b2ed4f7 bf2ef62 b2ed4f7 bf2ef62 faf85a4 bf2ef62 faf85a4 bf2ef62 faf85a4 bf2ef62 faf85a4 bf2ef62 373c574 bf2ef62 373c574 bf2ef62 6a47450 bf2ef62 373c574 bf2ef62 373c574 bf2ef62 4ca0bde bf2ef62 53b543e 4ca0bde bf2ef62 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 | # 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()
|