Madushani-Weerasekara commited on
Commit
bb62c99
·
verified ·
1 Parent(s): 79869bb

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +10 -13
app.py CHANGED
@@ -2,47 +2,43 @@ import os
2
  import streamlit as st
3
  from transformers import pipeline, AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
4
 
5
- import os
6
- import streamlit as st
7
- from transformers import pipeline, AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
8
-
9
  def load_model():
10
  model_id = "TheBloke/Mistral-7B-Instruct-v0.1-GPTQ"
11
-
12
  access_token = os.getenv("hf_mistral_token")
13
 
14
  tokenizer = AutoTokenizer.from_pretrained(model_id, use_fast=True, token=access_token)
15
 
16
- # Proper way to load GPTQ models (auto-gptq must be installed)
 
 
 
 
 
 
17
  model = AutoModelForCausalLM.from_pretrained(
18
  model_id,
 
19
  device_map="auto",
20
- trust_remote_code=True,
21
  token=access_token
22
  )
23
 
24
  pipe = pipeline("text-generation", model=model, tokenizer=tokenizer)
25
  return pipe
26
 
27
-
28
  def main():
29
  st.title("ChatGPT-Clone")
30
 
31
- # Load the generator model only once
32
  if "generator" not in st.session_state:
33
  with st.spinner("Loading model..."):
34
  st.session_state.generator = load_model()
35
 
36
- # Message history
37
  if "messages" not in st.session_state:
38
  st.session_state.messages = []
39
 
40
- # Display past messages
41
  for msg in st.session_state.messages:
42
  with st.chat_message(msg["role"]):
43
  st.markdown(msg["content"])
44
 
45
- # Chat input
46
  if prompt := st.chat_input("Ask anything..."):
47
  st.session_state.messages.append({"role": "user", "content": prompt})
48
 
@@ -55,8 +51,9 @@ def main():
55
  prompt,
56
  max_new_tokens=512,
57
  temperature=0.7,
58
- do_sample=True
59
  )[0]["generated_text"]
 
60
  st.markdown(result)
61
 
62
  st.session_state.messages.append({"role": "assistant", "content": result})
 
2
  import streamlit as st
3
  from transformers import pipeline, AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
4
 
 
 
 
 
5
  def load_model():
6
  model_id = "TheBloke/Mistral-7B-Instruct-v0.1-GPTQ"
 
7
  access_token = os.getenv("hf_mistral_token")
8
 
9
  tokenizer = AutoTokenizer.from_pretrained(model_id, use_fast=True, token=access_token)
10
 
11
+ quant_config = BitsAndBytesConfig(
12
+ load_in_4bit=True,
13
+ bnb_4bit_use_double_quant=True,
14
+ bnb_4bit_quant_type="nf4",
15
+ bnb_4bit_compute_dtype="float16"
16
+ )
17
+
18
  model = AutoModelForCausalLM.from_pretrained(
19
  model_id,
20
+ quantization_config=quant_config,
21
  device_map="auto",
 
22
  token=access_token
23
  )
24
 
25
  pipe = pipeline("text-generation", model=model, tokenizer=tokenizer)
26
  return pipe
27
 
 
28
  def main():
29
  st.title("ChatGPT-Clone")
30
 
 
31
  if "generator" not in st.session_state:
32
  with st.spinner("Loading model..."):
33
  st.session_state.generator = load_model()
34
 
 
35
  if "messages" not in st.session_state:
36
  st.session_state.messages = []
37
 
 
38
  for msg in st.session_state.messages:
39
  with st.chat_message(msg["role"]):
40
  st.markdown(msg["content"])
41
 
 
42
  if prompt := st.chat_input("Ask anything..."):
43
  st.session_state.messages.append({"role": "user", "content": prompt})
44
 
 
51
  prompt,
52
  max_new_tokens=512,
53
  temperature=0.7,
54
+ do_sample=True,
55
  )[0]["generated_text"]
56
+
57
  st.markdown(result)
58
 
59
  st.session_state.messages.append({"role": "assistant", "content": result})