mlandia commited on
Commit
e69ed13
·
verified ·
1 Parent(s): a41a548

Fix generation input handling for chat template

Browse files
Files changed (1) hide show
  1. app.py +14 -7
app.py CHANGED
@@ -41,6 +41,16 @@ EXAMPLES = [
41
  ]
42
 
43
 
 
 
 
 
 
 
 
 
 
 
44
  def _generate_one(
45
  key: str,
46
  prompt: str,
@@ -50,11 +60,8 @@ def _generate_one(
50
  ) -> tuple[str, float]:
51
  tokenizer = tokenizers[key]
52
  model = models[key]
53
- inputs = tokenizer.apply_chat_template(
54
- [{"role": "user", "content": prompt}],
55
- add_generation_prompt=True,
56
- return_tensors="pt",
57
- ).to(model.device)
58
 
59
  gen_kwargs: dict = {
60
  "max_new_tokens": max_new_tokens,
@@ -69,11 +76,11 @@ def _generate_one(
69
 
70
  started = time.perf_counter()
71
  with torch.inference_mode():
72
- output = model.generate(inputs, **gen_kwargs)
73
  elapsed = time.perf_counter() - started
74
 
75
  response = tokenizer.decode(
76
- output[0, inputs.shape[-1] :],
77
  skip_special_tokens=True,
78
  ).strip()
79
  return response, elapsed
 
41
  ]
42
 
43
 
44
+ def _build_inputs(tokenizer: AutoTokenizer, prompt: str, device: torch.device):
45
+ messages = [{"role": "user", "content": prompt}]
46
+ chat = tokenizer.apply_chat_template(
47
+ messages,
48
+ add_generation_prompt=True,
49
+ tokenize=False,
50
+ )
51
+ return tokenizer(chat, return_tensors="pt").to(device)
52
+
53
+
54
  def _generate_one(
55
  key: str,
56
  prompt: str,
 
60
  ) -> tuple[str, float]:
61
  tokenizer = tokenizers[key]
62
  model = models[key]
63
+ inputs = _build_inputs(tokenizer, prompt, model.device)
64
+ input_ids = inputs["input_ids"]
 
 
 
65
 
66
  gen_kwargs: dict = {
67
  "max_new_tokens": max_new_tokens,
 
76
 
77
  started = time.perf_counter()
78
  with torch.inference_mode():
79
+ output = model.generate(**inputs, **gen_kwargs)
80
  elapsed = time.perf_counter() - started
81
 
82
  response = tokenizer.decode(
83
+ output[0, input_ids.shape[-1] :],
84
  skip_special_tokens=True,
85
  ).strip()
86
  return response, elapsed