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()