Eeppa commited on
Commit
199367f
·
verified ·
1 Parent(s): 47793c2

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +16 -5
app.py CHANGED
@@ -33,15 +33,21 @@ for name in MODELS:
33
 
34
 
35
  # ---- manual sampling fallback (works for ANY model, ignores custom generate) ----
36
- def manual_sample(model, input_ids, max_new_tokens, temperature, top_k):
37
  """Generic top-k sampler. Bypasses any custom .generate() method."""
38
  generated = input_ids.clone()
39
- for _ in range(int(max_new_tokens)):
 
 
40
  with torch.no_grad():
41
  outputs = model(generated)
42
  logits = outputs.logits if hasattr(outputs, "logits") else outputs[0]
43
  next_logits = logits[0, -1, :] / max(temperature, 1e-5)
44
 
 
 
 
 
45
  if top_k > 0:
46
  k = min(int(top_k), next_logits.size(-1))
47
  top_vals, top_idx = torch.topk(next_logits, k)
@@ -54,8 +60,9 @@ def manual_sample(model, input_ids, max_new_tokens, temperature, top_k):
54
 
55
  generated = torch.cat([generated, next_token.view(1, 1)], dim=-1)
56
 
57
- if hasattr(model.config, "eos_token_id") and model.config.eos_token_id is not None:
58
- if next_token.item() == model.config.eos_token_id:
 
59
  break
60
 
61
  return generated
@@ -98,6 +105,7 @@ def run_generate(model_name, prompt, max_new_tokens, temperature, top_k):
98
  out = model.generate(
99
  input_ids,
100
  max_new_tokens=int(max_new_tokens),
 
101
  temperature=float(temperature),
102
  top_k=int(top_k),
103
  do_sample=True,
@@ -105,7 +113,10 @@ def run_generate(model_name, prompt, max_new_tokens, temperature, top_k):
105
  return decode_output(tokenizer, out)
106
  except Exception as e:
107
  print(f"[fallback] {model_name} custom generate failed ({e}), using manual sampling")
108
- out = manual_sample(model, input_ids, max_new_tokens, temperature, top_k)
 
 
 
109
  return decode_output(tokenizer, out)
110
 
111
 
 
33
 
34
 
35
  # ---- manual sampling fallback (works for ANY model, ignores custom generate) ----
36
+ def manual_sample(model, input_ids, max_new_tokens, temperature, top_k, min_new_tokens=0):
37
  """Generic top-k sampler. Bypasses any custom .generate() method."""
38
  generated = input_ids.clone()
39
+ eos_id = getattr(model.config, "eos_token_id", None)
40
+
41
+ for i in range(int(max_new_tokens)):
42
  with torch.no_grad():
43
  outputs = model(generated)
44
  logits = outputs.logits if hasattr(outputs, "logits") else outputs[0]
45
  next_logits = logits[0, -1, :] / max(temperature, 1e-5)
46
 
47
+ # Suppress EOS for the first `min_new_tokens` steps
48
+ if i < min_new_tokens and eos_id is not None:
49
+ next_logits[eos_id] = float("-inf")
50
+
51
  if top_k > 0:
52
  k = min(int(top_k), next_logits.size(-1))
53
  top_vals, top_idx = torch.topk(next_logits, k)
 
60
 
61
  generated = torch.cat([generated, next_token.view(1, 1)], dim=-1)
62
 
63
+ # Only allow EOS to stop generation after min_new_tokens
64
+ if i >= min_new_tokens and eos_id is not None:
65
+ if next_token.item() == eos_id:
66
  break
67
 
68
  return generated
 
105
  out = model.generate(
106
  input_ids,
107
  max_new_tokens=int(max_new_tokens),
108
+ min_new_tokens=10, # <-- force at least 10 new tokens
109
  temperature=float(temperature),
110
  top_k=int(top_k),
111
  do_sample=True,
 
113
  return decode_output(tokenizer, out)
114
  except Exception as e:
115
  print(f"[fallback] {model_name} custom generate failed ({e}), using manual sampling")
116
+ out = manual_sample(
117
+ model, input_ids, max_new_tokens, temperature, top_k,
118
+ min_new_tokens=10, # <-- same here
119
+ )
120
  return decode_output(tokenizer, out)
121
 
122