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

Zwiększyliśmy timeout dla TextIteratorStreamer do 60 sekund.

Browse files
Files changed (1) hide show
  1. app.py +16 -10
app.py CHANGED
@@ -4,7 +4,7 @@ import gradio as gr
4
  import torch
5
  from transformers import AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer
6
  from threading import Thread
7
- from queue import Queue
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"))
@@ -32,18 +32,25 @@ def generate_response(chat, kwargs):
32
  model = model_load_queue.get()
33
 
34
  inputs = tokenizer(chat, return_tensors="pt").to(model.device)
35
- streamer = TextIteratorStreamer(tokenizer, timeout=10., skip_prompt=True, skip_special_tokens=True)
 
 
 
 
36
 
37
  generation_kwargs = dict(inputs, streamer=streamer, **kwargs)
38
  thread = Thread(target=model.generate, kwargs=generation_kwargs)
39
  thread.start()
40
 
41
  output = ""
42
- for new_text in streamer:
43
- output += new_text
44
- if output.endswith("</s>"):
45
- output = output[:-4]
46
- break
 
 
 
47
 
48
  return output
49
 
@@ -57,8 +64,7 @@ def function(prompt, history=[]):
57
  do_sample=True,
58
  temperature=0.5,
59
  top_p=0.95,
60
- repetition_penalty=1.0,
61
- seed=1337
62
  )
63
 
64
  try:
@@ -66,7 +72,7 @@ def function(prompt, history=[]):
66
  return output
67
  except Exception as e:
68
  print(f"Error: {str(e)}")
69
- return ''
70
 
71
  interface = gr.ChatInterface(
72
  fn=function,
 
4
  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"))
 
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
 
 
64
  do_sample=True,
65
  temperature=0.5,
66
  top_p=0.95,
67
+ repetition_penalty=1.0
 
68
  )
69
 
70
  try:
 
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,