Eeppa commited on
Commit
47793c2
·
verified ·
1 Parent(s): 0fd172b

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +7 -14
app.py CHANGED
@@ -12,20 +12,14 @@ MODELS = {
12
  loaded = {}
13
  load_errors = {}
14
 
 
15
  def load_model(name):
16
  if name in loaded or name in load_errors:
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}")
@@ -33,6 +27,7 @@ def load_model(name):
33
  load_errors[name] = f"{type(e).__name__}: {e}"
34
  print(f"[load failed] {name}: {e}")
35
 
 
36
  for name in MODELS:
37
  load_model(name)
38
 
@@ -67,7 +62,7 @@ def manual_sample(model, input_ids, max_new_tokens, temperature, top_k):
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"):
@@ -75,7 +70,6 @@ def encode_prompt(tokenizer, prompt):
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):
@@ -85,7 +79,7 @@ def encode_prompt(tokenizer, prompt):
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):
@@ -98,8 +92,7 @@ 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(
 
12
  loaded = {}
13
  load_errors = {}
14
 
15
+
16
  def load_model(name):
17
  if name in loaded or name in load_errors:
18
  return
19
  repo = MODELS[name]
20
  try:
21
+ tokenizer = AutoTokenizer.from_pretrained(repo, trust_remote_code=True)
22
+ model = AutoModelForCausalLM.from_pretrained(repo, trust_remote_code=True)
 
 
 
 
 
 
 
23
  model.eval()
24
  loaded[name] = (model, tokenizer)
25
  print(f"[loaded] {name}")
 
27
  load_errors[name] = f"{type(e).__name__}: {e}"
28
  print(f"[load failed] {name}: {e}")
29
 
30
+
31
  for name in MODELS:
32
  load_model(name)
33
 
 
62
 
63
 
64
  def encode_prompt(tokenizer, prompt):
65
+ """Handle both fast and slow tokenizers, and flatten nested outputs."""
66
  if hasattr(tokenizer, "encode"):
67
  ids = tokenizer.encode(prompt)
68
  if hasattr(ids, "ids"):
 
70
  else:
71
  ids = tokenizer(prompt)["input_ids"]
72
 
 
73
  if torch.is_tensor(ids):
74
  ids = ids.tolist()
75
  if isinstance(ids, list) and len(ids) > 0 and isinstance(ids[0], list):
 
79
 
80
 
81
  def decode_output(tokenizer, ids):
82
+ """Normalize to a flat python list of ints before decoding."""
83
  if torch.is_tensor(ids):
84
  ids = ids.tolist()
85
  if isinstance(ids, list) and len(ids) > 0 and isinstance(ids[0], list):
 
92
  model, tokenizer = loaded[model_name]
93
  input_ids = encode_prompt(tokenizer, prompt)
94
 
95
+ # Try the model's own generate first, fall back to manual sampling on failure
 
96
  try:
97
  with torch.no_grad():
98
  out = model.generate(