anirudh248 commited on
Commit
cda579b
·
verified ·
1 Parent(s): 01abe35

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +9 -16
app.py CHANGED
@@ -1,9 +1,9 @@
1
  import spaces
2
  import gradio as gr
3
  import torch
4
- from langchain_huggingface import HuggingFaceEmbeddings, HuggingFacePipeline
5
  from langchain_community.vectorstores import FAISS
6
- from transformers import pipeline, AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
 
7
  from langchain_classic.chains import create_retrieval_chain
8
  from langchain_classic.chains.combine_documents import create_stuff_documents_chain
9
  from langchain_core.prompts import PromptTemplate
@@ -14,7 +14,6 @@ from langchain_core.prompts import PromptTemplate
14
  print("Loading FAISS Vector Store...")
15
  embeddings = HuggingFaceEmbeddings(model_name='sentence-transformers/all-MiniLM-L6-v2')
16
 
17
- # allow_dangerous_deserialization is required for FAISS in newer LangChain versions
18
  vectorstore = FAISS.load_local(
19
  "faiss_upf_index",
20
  embeddings,
@@ -23,21 +22,16 @@ vectorstore = FAISS.load_local(
23
  retriever = vectorstore.as_retriever()
24
 
25
  # ==========================================
26
- # 2. Load Model & Tokenizer in 4-bit
27
  # ==========================================
28
  print("Loading Model...")
29
  model_id = "anirudh248/llama3-upf-generator"
30
 
31
- # Load in 4-bit to fit perfectly inside a GPU Space
32
- bnb_config = BitsAndBytesConfig(
33
- load_in_4bit=True,
34
- bnb_4bit_compute_dtype=torch.float16
35
- )
36
-
37
  tokenizer = AutoTokenizer.from_pretrained(model_id)
 
38
  model = AutoModelForCausalLM.from_pretrained(
39
  model_id,
40
- quantization_config=bnb_config,
41
  device_map="auto"
42
  )
43
  model.generation_config.pad_token_id = tokenizer.eos_token_id
@@ -74,16 +68,15 @@ prompt = PromptTemplate.from_template(prompt_template)
74
  document_chain = create_stuff_documents_chain(llm, prompt)
75
  rag_chain = create_retrieval_chain(retriever, document_chain)
76
 
 
 
 
77
 
78
  # ==========================================
79
- # 4. Gradio UI & Inference
80
  # ==========================================
81
 
82
  @spaces.GPU(duration=120)
83
- def generate_upf_code(power_intent_description):
84
- result = rag_chain.invoke({"input": power_intent_description})
85
- return result['answer'].strip()
86
-
87
  def user_interaction(user_message, history):
88
  history = history or []
89
  response = generate_upf_code(user_message)
 
1
  import spaces
2
  import gradio as gr
3
  import torch
 
4
  from langchain_community.vectorstores import FAISS
5
+ from langchain_huggingface import HuggingFaceEmbeddings, HuggingFacePipeline
6
+ from transformers import pipeline, AutoModelForCausalLM, AutoTokenizer
7
  from langchain_classic.chains import create_retrieval_chain
8
  from langchain_classic.chains.combine_documents import create_stuff_documents_chain
9
  from langchain_core.prompts import PromptTemplate
 
14
  print("Loading FAISS Vector Store...")
15
  embeddings = HuggingFaceEmbeddings(model_name='sentence-transformers/all-MiniLM-L6-v2')
16
 
 
17
  vectorstore = FAISS.load_local(
18
  "faiss_upf_index",
19
  embeddings,
 
22
  retriever = vectorstore.as_retriever()
23
 
24
  # ==========================================
25
+ # 2. Load Model & Tokenizer in 16-bit (ZeroGPU Fix)
26
  # ==========================================
27
  print("Loading Model...")
28
  model_id = "anirudh248/llama3-upf-generator"
29
 
 
 
 
 
 
 
30
  tokenizer = AutoTokenizer.from_pretrained(model_id)
31
+
32
  model = AutoModelForCausalLM.from_pretrained(
33
  model_id,
34
+ torch_dtype=torch.bfloat16,
35
  device_map="auto"
36
  )
37
  model.generation_config.pad_token_id = tokenizer.eos_token_id
 
68
  document_chain = create_stuff_documents_chain(llm, prompt)
69
  rag_chain = create_retrieval_chain(retriever, document_chain)
70
 
71
+ def generate_upf_code(power_intent_description):
72
+ result = rag_chain.invoke({"input": power_intent_description})
73
+ return result['answer'].strip()
74
 
75
  # ==========================================
76
+ # 4. Gradio UI
77
  # ==========================================
78
 
79
  @spaces.GPU(duration=120)
 
 
 
 
80
  def user_interaction(user_message, history):
81
  history = history or []
82
  response = generate_upf_code(user_message)