Gemini commited on
Commit
073b1d0
·
1 Parent(s): f600c0f

Switch to CPU-compatible model

Browse files
Files changed (2) hide show
  1. app.py +6 -14
  2. requirements.txt +1 -2
app.py CHANGED
@@ -1,25 +1,17 @@
1
  from fastapi import FastAPI
2
  from pydantic import BaseModel
3
- from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
4
  import torch
5
  from typing import List
6
 
7
  app = FastAPI()
8
 
9
- model_id = "unsloth/gemma-2-2b-it-bnb-4bit"
10
-
11
- bnb_config = BitsAndBytesConfig(
12
- load_in_4bit=True,
13
- bnb_4bit_compute_dtype=torch.bfloat16,
14
- bnb_4bit_use_double_quant=True,
15
- bnb_4bit_quant_type="nf4"
16
- )
17
 
18
  tokenizer = AutoTokenizer.from_pretrained(model_id)
19
  model = AutoModelForCausalLM.from_pretrained(
20
  model_id,
21
- quantization_config=bnb_config,
22
- device_map="auto"
23
  )
24
 
25
  class InferenceRequest(BaseModel):
@@ -38,7 +30,7 @@ def read_root():
38
 
39
  @app.post("/infer")
40
  def infer(request: InferenceRequest):
41
- inputs = tokenizer(request.text, return_tensors="pt").to("cuda")
42
  outputs = model.generate(**inputs, max_new_tokens=50)
43
  decoded_output = tokenizer.decode(outputs[0], skip_special_tokens=True)
44
  return {"generated_text": decoded_output}
@@ -47,9 +39,9 @@ def infer(request: InferenceRequest):
47
  def chat(request: ChatRequest):
48
  chat_history = [message.dict() for message in request.messages]
49
  prompt = tokenizer.apply_chat_template(chat_history, tokenize=False, add_generation_prompt=True)
50
- inputs = tokenizer(prompt, return_tensors="pt").to("cuda")
51
  outputs = model.generate(**inputs, max_new_tokens=150)
52
  decoded_output = tokenizer.decode(outputs[0], skip_special_tokens=True)
53
  # The output from the model includes the prompt, so we need to remove it.
54
  response = decoded_output.split("<end_of_turn>\n")[-1]
55
- return {"generated_text": response}
 
1
  from fastapi import FastAPI
2
  from pydantic import BaseModel
3
+ from transformers import AutoModelForCausalLM, AutoTokenizer
4
  import torch
5
  from typing import List
6
 
7
  app = FastAPI()
8
 
9
+ model_id = "google/gemma-2b-it"
 
 
 
 
 
 
 
10
 
11
  tokenizer = AutoTokenizer.from_pretrained(model_id)
12
  model = AutoModelForCausalLM.from_pretrained(
13
  model_id,
14
+ device_map="cpu"
 
15
  )
16
 
17
  class InferenceRequest(BaseModel):
 
30
 
31
  @app.post("/infer")
32
  def infer(request: InferenceRequest):
33
+ inputs = tokenizer(request.text, return_tensors="pt")
34
  outputs = model.generate(**inputs, max_new_tokens=50)
35
  decoded_output = tokenizer.decode(outputs[0], skip_special_tokens=True)
36
  return {"generated_text": decoded_output}
 
39
  def chat(request: ChatRequest):
40
  chat_history = [message.dict() for message in request.messages]
41
  prompt = tokenizer.apply_chat_template(chat_history, tokenize=False, add_generation_prompt=True)
42
+ inputs = tokenizer(prompt, return_tensors="pt")
43
  outputs = model.generate(**inputs, max_new_tokens=150)
44
  decoded_output = tokenizer.decode(outputs[0], skip_special_tokens=True)
45
  # The output from the model includes the prompt, so we need to remove it.
46
  response = decoded_output.split("<end_of_turn>\n")[-1]
47
+ return {"generated_text": response}
requirements.txt CHANGED
@@ -2,8 +2,7 @@ fastapi
2
  uvicorn
3
  torch
4
  transformers
5
- bitsandbytes
6
  accelerate
7
  sentencepiece
8
  python-dotenv
9
- pydantic
 
2
  uvicorn
3
  torch
4
  transformers
 
5
  accelerate
6
  sentencepiece
7
  python-dotenv
8
+ pydantic