Eeppa commited on
Commit
0fd172b
·
verified ·
1 Parent(s): 6125f9e

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +32 -10
app.py CHANGED
@@ -17,8 +17,15 @@ def load_model(name):
17
  return
18
  repo = MODELS[name]
19
  try:
20
- tokenizer = AutoTokenizer.from_pretrained(repo, trust_remote_code=True)
21
- model = AutoModelForCausalLM.from_pretrained(repo, trust_remote_code=True)
 
 
 
 
 
 
 
22
  model.eval()
23
  loaded[name] = (model, tokenizer)
24
  print(f"[loaded] {name}")
@@ -29,6 +36,7 @@ def load_model(name):
29
  for name in MODELS:
30
  load_model(name)
31
 
 
32
  # ---- manual sampling fallback (works for ANY model, ignores custom generate) ----
33
  def manual_sample(model, input_ids, max_new_tokens, temperature, top_k):
34
  """Generic top-k sampler. Bypasses any custom .generate() method."""
@@ -39,7 +47,6 @@ def manual_sample(model, input_ids, max_new_tokens, temperature, top_k):
39
  logits = outputs.logits if hasattr(outputs, "logits") else outputs[0]
40
  next_logits = logits[0, -1, :] / max(temperature, 1e-5)
41
 
42
- # top-k filtering
43
  if top_k > 0:
44
  k = min(int(top_k), next_logits.size(-1))
45
  top_vals, top_idx = torch.topk(next_logits, k)
@@ -52,34 +59,47 @@ def manual_sample(model, input_ids, max_new_tokens, temperature, top_k):
52
 
53
  generated = torch.cat([generated, next_token.view(1, 1)], dim=-1)
54
 
55
- # stop if we hit an EOS
56
  if hasattr(model.config, "eos_token_id") and model.config.eos_token_id is not None:
57
  if next_token.item() == model.config.eos_token_id:
58
  break
59
 
60
  return generated
61
 
 
62
  def encode_prompt(tokenizer, prompt):
 
63
  if hasattr(tokenizer, "encode"):
64
  ids = tokenizer.encode(prompt)
65
  if hasattr(ids, "ids"):
66
  ids = ids.ids
67
  else:
68
  ids = tokenizer(prompt)["input_ids"]
 
 
 
 
 
 
 
69
  return torch.tensor([ids], dtype=torch.long)
70
 
 
71
  def decode_output(tokenizer, ids):
72
- try:
73
- return tokenizer.decode(ids[0].tolist() if torch.is_tensor(ids[0]) else ids[0], skip_special_tokens=True)
74
- except TypeError:
75
- return tokenizer.decode(ids[0], skip_special_tokens=True)
 
 
 
76
 
77
  @spaces.GPU(duration=30)
78
  def run_generate(model_name, prompt, max_new_tokens, temperature, top_k):
79
  model, tokenizer = loaded[model_name]
80
  input_ids = encode_prompt(tokenizer, prompt)
81
 
82
- # try the model's own generate first, fall back to manual sampling
 
83
  try:
84
  with torch.no_grad():
85
  out = model.generate(
@@ -90,11 +110,12 @@ def run_generate(model_name, prompt, max_new_tokens, temperature, top_k):
90
  do_sample=True,
91
  )
92
  return decode_output(tokenizer, out)
93
- except (TypeError, AttributeError, ValueError) as e:
94
  print(f"[fallback] {model_name} custom generate failed ({e}), using manual sampling")
95
  out = manual_sample(model, input_ids, max_new_tokens, temperature, top_k)
96
  return decode_output(tokenizer, out)
97
 
 
98
  def compare_all(prompt, max_new_tokens, temperature, top_k):
99
  results = []
100
  for name in MODELS:
@@ -111,6 +132,7 @@ def compare_all(prompt, max_new_tokens, temperature, top_k):
111
  results.append(text)
112
  return results
113
 
 
114
  with gr.Blocks(title="TinyBuddy Trilogy") as demo:
115
  gr.Markdown(
116
  """
 
17
  return
18
  repo = MODELS[name]
19
  try:
20
+ # FIX 1: use_fast=False forces the slow GPT-2 tokenizer path,
21
+ # which reads vocab.json + merges.txt directly and never touches
22
+ # the broken tokenizer.json -> avoids the "missing field trim_offsets" error.
23
+ tokenizer = AutoTokenizer.from_pretrained(
24
+ repo, trust_remote_code=True, use_fast=False
25
+ )
26
+ model = AutoModelForCausalLM.from_pretrained(
27
+ repo, trust_remote_code=True
28
+ )
29
  model.eval()
30
  loaded[name] = (model, tokenizer)
31
  print(f"[loaded] {name}")
 
36
  for name in MODELS:
37
  load_model(name)
38
 
39
+
40
  # ---- manual sampling fallback (works for ANY model, ignores custom generate) ----
41
  def manual_sample(model, input_ids, max_new_tokens, temperature, top_k):
42
  """Generic top-k sampler. Bypasses any custom .generate() method."""
 
47
  logits = outputs.logits if hasattr(outputs, "logits") else outputs[0]
48
  next_logits = logits[0, -1, :] / max(temperature, 1e-5)
49
 
 
50
  if top_k > 0:
51
  k = min(int(top_k), next_logits.size(-1))
52
  top_vals, top_idx = torch.topk(next_logits, k)
 
59
 
60
  generated = torch.cat([generated, next_token.view(1, 1)], dim=-1)
61
 
 
62
  if hasattr(model.config, "eos_token_id") and model.config.eos_token_id is not None:
63
  if next_token.item() == model.config.eos_token_id:
64
  break
65
 
66
  return generated
67
 
68
+
69
  def encode_prompt(tokenizer, prompt):
70
+ # FIX 2: more robust handling for both fast and slow tokenizers.
71
  if hasattr(tokenizer, "encode"):
72
  ids = tokenizer.encode(prompt)
73
  if hasattr(ids, "ids"):
74
  ids = ids.ids
75
  else:
76
  ids = tokenizer(prompt)["input_ids"]
77
+
78
+ # Ensure we always have a flat list of ints
79
+ if torch.is_tensor(ids):
80
+ ids = ids.tolist()
81
+ if isinstance(ids, list) and len(ids) > 0 and isinstance(ids[0], list):
82
+ ids = ids[0]
83
+
84
  return torch.tensor([ids], dtype=torch.long)
85
 
86
+
87
  def decode_output(tokenizer, ids):
88
+ # FIX 3: normalize to a flat python list of ints before decoding.
89
+ if torch.is_tensor(ids):
90
+ ids = ids.tolist()
91
+ if isinstance(ids, list) and len(ids) > 0 and isinstance(ids[0], list):
92
+ ids = ids[0]
93
+ return tokenizer.decode(ids, skip_special_tokens=True)
94
+
95
 
96
  @spaces.GPU(duration=30)
97
  def run_generate(model_name, prompt, max_new_tokens, temperature, top_k):
98
  model, tokenizer = loaded[model_name]
99
  input_ids = encode_prompt(tokenizer, prompt)
100
 
101
+ # FIX 4: some tiny custom models don't implement .generate() correctly,
102
+ # so we try it first and fall back to manual sampling on ANY exception.
103
  try:
104
  with torch.no_grad():
105
  out = model.generate(
 
110
  do_sample=True,
111
  )
112
  return decode_output(tokenizer, out)
113
+ except Exception as e:
114
  print(f"[fallback] {model_name} custom generate failed ({e}), using manual sampling")
115
  out = manual_sample(model, input_ids, max_new_tokens, temperature, top_k)
116
  return decode_output(tokenizer, out)
117
 
118
+
119
  def compare_all(prompt, max_new_tokens, temperature, top_k):
120
  results = []
121
  for name in MODELS:
 
132
  results.append(text)
133
  return results
134
 
135
+
136
  with gr.Blocks(title="TinyBuddy Trilogy") as demo:
137
  gr.Markdown(
138
  """