wiklif commited on
Commit
e234eec
·
1 Parent(s): e3fd506

lepsze logowanie błędów, timeout zwiększony do 120s

Browse files
Files changed (1) hide show
  1. app.py +52 -37
app.py CHANGED
@@ -5,6 +5,10 @@ import torch
5
  from transformers import AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer
6
  from threading import Thread
7
  from queue import Queue, Empty
 
 
 
 
8
 
9
  model_id = "meta-llama/Meta-Llama-3.1-8B"
10
  tokenizer = AutoTokenizer.from_pretrained(model_id, token=os.environ.get("MY_API_LLAMA_3_1"))
@@ -14,45 +18,61 @@ model_load_queue = Queue()
14
 
15
  def load_model():
16
  global model
17
- if model is None:
18
- model = AutoModelForCausalLM.from_pretrained(
19
- model_id,
20
- token=os.environ.get("MY_API_LLAMA_3_1"),
21
- torch_dtype=torch.bfloat16,
22
- device_map="auto",
23
- low_cpu_mem_usage=True
24
- )
25
- model_load_queue.put(model)
 
 
 
 
 
 
26
 
27
- @spaces.GPU(duration=60)
28
  def generate_response(chat, kwargs):
29
  global model
30
- if model is None:
31
- Thread(target=load_model).start()
32
- model = model_load_queue.get()
 
 
 
 
33
 
34
- inputs = tokenizer(chat, return_tensors="pt").to(model.device)
35
- streamer = TextIteratorStreamer(tokenizer, timeout=60., skip_prompt=True, skip_special_tokens=True)
 
36
 
37
- # Usuwamy 'seed' z kwargs, ponieważ nie jest obsługiwany przez model
38
- if 'seed' in kwargs:
39
- del kwargs['seed']
40
 
41
- generation_kwargs = dict(inputs, streamer=streamer, **kwargs)
42
- thread = Thread(target=model.generate, kwargs=generation_kwargs)
43
- thread.start()
44
 
45
- output = ""
46
- try:
47
- for new_text in streamer:
48
- output += new_text
49
- if output.endswith("</s>"):
50
- output = output[:-4]
51
- break
52
- except Empty:
53
- print("Timeout occurred during generation")
54
 
55
- return output
 
 
 
 
 
 
 
 
 
 
 
 
 
 
56
 
57
  def function(prompt, history=[]):
58
  chat = "<s>"
@@ -67,12 +87,7 @@ def function(prompt, history=[]):
67
  repetition_penalty=1.0
68
  )
69
 
70
- try:
71
- output = generate_response(chat, kwargs)
72
- return output
73
- except Exception as e:
74
- print(f"Error: {str(e)}")
75
- return 'Wystąpił błąd podczas generowania odpowiedzi.'
76
 
77
  interface = gr.ChatInterface(
78
  fn=function,
 
5
  from transformers import AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer
6
  from threading import Thread
7
  from queue import Queue, Empty
8
+ import logging
9
+
10
+ logging.basicConfig(level=logging.INFO)
11
+ logger = logging.getLogger(__name__)
12
 
13
  model_id = "meta-llama/Meta-Llama-3.1-8B"
14
  tokenizer = AutoTokenizer.from_pretrained(model_id, token=os.environ.get("MY_API_LLAMA_3_1"))
 
18
 
19
  def load_model():
20
  global model
21
+ try:
22
+ if model is None:
23
+ logger.info("Loading model...")
24
+ model = AutoModelForCausalLM.from_pretrained(
25
+ model_id,
26
+ token=os.environ.get("MY_API_LLAMA_3_1"),
27
+ torch_dtype=torch.bfloat16,
28
+ device_map="auto",
29
+ low_cpu_mem_usage=True
30
+ )
31
+ logger.info("Model loaded successfully")
32
+ model_load_queue.put(model)
33
+ except Exception as e:
34
+ logger.error(f"Error loading model: {str(e)}")
35
+ model_load_queue.put(None)
36
 
37
+ @spaces.GPU(duration=120)
38
  def generate_response(chat, kwargs):
39
  global model
40
+ try:
41
+ if model is None:
42
+ logger.info("Starting model loading thread")
43
+ Thread(target=load_model).start()
44
+ model = model_load_queue.get(timeout=120)
45
+ if model is None:
46
+ return "Nie udało się załadować modelu. Proszę spróbować ponownie później."
47
 
48
+ logger.info("Preparing input for generation")
49
+ inputs = tokenizer(chat, return_tensors="pt").to(model.device)
50
+ streamer = TextIteratorStreamer(tokenizer, timeout=120., skip_prompt=True, skip_special_tokens=True)
51
 
52
+ if 'seed' in kwargs:
53
+ del kwargs['seed']
 
54
 
55
+ generation_kwargs = dict(inputs, streamer=streamer, **kwargs)
 
 
56
 
57
+ logger.info("Starting generation thread")
58
+ thread = Thread(target=model.generate, kwargs=generation_kwargs)
59
+ thread.start()
 
 
 
 
 
 
60
 
61
+ output = ""
62
+ try:
63
+ for new_text in streamer:
64
+ output += new_text
65
+ if output.endswith("</s>"):
66
+ output = output[:-4]
67
+ break
68
+ except Empty:
69
+ logger.warning("Timeout occurred during generation")
70
+
71
+ logger.info("Generation completed")
72
+ return output
73
+ except Exception as e:
74
+ logger.error(f"Error in generate_response: {str(e)}")
75
+ return f"Wystąpił błąd: {str(e)}"
76
 
77
  def function(prompt, history=[]):
78
  chat = "<s>"
 
87
  repetition_penalty=1.0
88
  )
89
 
90
+ return generate_response(chat, kwargs)
 
 
 
 
 
91
 
92
  interface = gr.ChatInterface(
93
  fn=function,