Spaces:
Runtime error
Runtime error
File size: 3,478 Bytes
68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 956aea5 68b1970 | 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 | import torch, faiss, json
import numpy as np
import gradio as gr
from transformers import (AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig,
AutoModelForSequenceClassification)
from sentence_transformers import SentenceTransformer
from huggingface_hub import hf_hub_download
import os
HF_REPO = os.environ.get("HF_REPO", "Premchan369/JurisGPT")
device = "cuda" if torch.cuda.is_available() else "cpu"
print("Loading models...")
bnb = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.float16)
tok = AutoTokenizer.from_pretrained(HF_REPO)
mdl = AutoModelForCausalLM.from_pretrained(HF_REPO, quantization_config=bnb, device_map="auto")
mdl.eval()
emb = SentenceTransformer("BAAI/bge-large-en")
idx = faiss.read_index(hf_hub_download(HF_REPO, "jurisgpt_faiss.index"))
with open(hf_hub_download(HF_REPO, "jurisgpt_metadata.json")) as f:
meta = json.load(f)
lb_tok = AutoTokenizer.from_pretrained("nlpaueb/legal-bert-base-uncased")
lb_mdl = AutoModelForSequenceClassification.from_pretrained(
"nlpaueb/legal-bert-base-uncased").to(device)
lb_mdl.eval()
print(f"Ready — {meta['total_docs']:,} docs loaded")
def retrieve(q, k=5):
qv = emb.encode([q], normalize_embeddings=True).astype("float32")
sc, ids = idx.search(qv, k)
return "\n\n".join(
f"[{meta['doc_sources'][i]}]\n{meta['all_docs'][i]}"
for s, i in zip(sc[0], ids[0]) if i < len(meta["all_docs"])
)
def predict(text):
inp = lb_tok(text, return_tensors="pt", truncation=True, max_length=512).to(device)
with torch.no_grad():
logits = lb_mdl(**inp).logits
probs = torch.softmax(logits, dim=-1).cpu().numpy()[0]
label = ["Unfavorable", "Favorable"][int(probs.argmax())]
return f"{label} ({float(probs.max()):.1%})"
def gen(prompt, max_t=700):
SYSTEM = "You are JurisGPT, expert Indian legal AI. Cite IPC sections and Constitutional articles. Be concise and practical."
msgs = [{"role":"system","content":SYSTEM}, {"role":"user","content":prompt}]
t = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
inp = tok(t, return_tensors="pt").to(mdl.device)
with torch.no_grad():
out = mdl.generate(**inp, max_new_tokens=max_t, temperature=0.3,
do_sample=True, repetition_penalty=1.1,
pad_token_id=tok.eos_token_id)
return tok.decode(out[0][inp["input_ids"].shape[1]:], skip_special_tokens=True).strip()
def chat(msg, hist):
ctx = retrieve(msg)
outcome = predict(msg)
return gen(f"Query: {msg}\nOutcome: {outcome}\nLaws:\n{ctx[:2000]}\n\nAnalysis:")
def quick(q):
return gen(f"Q: {q}\nContext:\n{retrieve(q,3)[:1000]}\n\nBrief answer:", 300)
with gr.Blocks(title="JurisGPT") as demo:
gr.Markdown(f"# JurisGPT — Indian Legal AI\n**{meta['total_docs']:,} docs** | Informational only.")
with gr.Tab("Legal Chat"):
gr.ChatInterface(fn=chat, examples=[
"Explain Article 21.", "IPC Section 302?",
"Landlord won\'t return deposit — what can I do?",
"Rights if arrested in India?",
"What is IPC Section 498A?",
])
with gr.Tab("Quick Q&A"):
q_in = gr.Textbox(label="Question", lines=2)
q_out= gr.Textbox(label="Answer", lines=8, show_copy_button=True)
gr.Button("Ask", variant="primary").click(quick, q_in, q_out)
demo.launch()
|