Andhs commited on
Commit
9ebf833
·
verified ·
1 Parent(s): 69dfb7d

Upload app.py

Browse files
Files changed (1) hide show
  1. app.py +9 -8
app.py CHANGED
@@ -107,7 +107,7 @@ def generate(model, tokenizer, prompt, steps=128, max_new_tokens=128, block_size
107
  if isinstance(prompt, torch.Tensor):
108
  x = prompt.to(device).long()
109
  else:
110
- if isinstance(prompt, (list, tuple)):
111
  max_len = max(len(p) for p in prompt)
112
  x = torch.full((len(prompt), max_len), pad_id, device=device, dtype=torch.long)
113
  for i, p in enumerate(prompt):
@@ -165,11 +165,12 @@ def generate(model, tokenizer, prompt, steps=128, max_new_tokens=128, block_size
165
  if capture_interval > 0 and total_step % capture_interval == 0:
166
  intermediates.append(x.clone())
167
  total_step += 1
168
- if tokenizer.eos_token_id is not None:
169
- finished |= (x_blk_new == tokenizer.eos_token_id).any(dim=1)
 
170
  if finished.all():
171
  break
172
- generated += cur_len
173
  if capture_interval > 0:
174
  return x, intermediates
175
  return x
@@ -184,7 +185,7 @@ def generate_stream(model, tokenizer, prompt, steps=128, max_new_tokens=128, blo
184
  if isinstance(prompt, torch.Tensor):
185
  x = prompt.to(device).long()
186
  else:
187
- if isinstance(prompt, (list, tuple)):
188
  max_len = max(len(p) for p in prompt)
189
  x = torch.full((len(prompt), max_len), pad_id, device=device, dtype=torch.long)
190
  for i, p in enumerate(prompt):
@@ -261,7 +262,7 @@ def generate_stream(model, tokenizer, prompt, steps=128, max_new_tokens=128, blo
261
  finished |= (x_blk_new == tokenizer.eos_token_id).any(dim=1)
262
  if finished.all():
263
  break
264
- generated += cur_len
265
 
266
  if finished.all():
267
  break
@@ -370,7 +371,7 @@ def generate_text_stream():
370
  return jsonify({"error": "Model workspace offline"}), 503
371
 
372
  data = request.get_json() or {}
373
- if 'prompt' not in data:
374
  return jsonify({"error": "Missing 'prompt' operational field"}), 400
375
 
376
  prompt = data['prompt']
@@ -422,7 +423,7 @@ def generate_text_sse():
422
  return jsonify({"error": "Model workspace offline"}), 503
423
 
424
  data = request.get_json() or {}
425
- if 'prompt' not in data:
426
  return jsonify({"error": "Missing 'prompt' operational field"}), 400
427
 
428
  prompt = data['prompt']
 
107
  if isinstance(prompt, torch.Tensor):
108
  x = prompt.to(device).long()
109
  else:
110
+ if isinstance(prompt[0], (list, tuple)):
111
  max_len = max(len(p) for p in prompt)
112
  x = torch.full((len(prompt), max_len), pad_id, device=device, dtype=torch.long)
113
  for i, p in enumerate(prompt):
 
165
  if capture_interval > 0 and total_step % capture_interval == 0:
166
  intermediates.append(x.clone())
167
  total_step += 1
168
+ if tokenizer.eos_token_id is not None:
169
+ finished |= (x_blk_new == tokenizer.eos_token_id).any(dim=1)
170
+ generated += cur_len
171
  if finished.all():
172
  break
173
+
174
  if capture_interval > 0:
175
  return x, intermediates
176
  return x
 
185
  if isinstance(prompt, torch.Tensor):
186
  x = prompt.to(device).long()
187
  else:
188
+ if isinstance(prompt[0], (list, tuple)):
189
  max_len = max(len(p) for p in prompt)
190
  x = torch.full((len(prompt), max_len), pad_id, device=device, dtype=torch.long)
191
  for i, p in enumerate(prompt):
 
262
  finished |= (x_blk_new == tokenizer.eos_token_id).any(dim=1)
263
  if finished.all():
264
  break
265
+ generated += cur_len
266
 
267
  if finished.all():
268
  break
 
371
  return jsonify({"error": "Model workspace offline"}), 503
372
 
373
  data = request.get_json() or {}
374
+ if not data or 'prompt' not in data:
375
  return jsonify({"error": "Missing 'prompt' operational field"}), 400
376
 
377
  prompt = data['prompt']
 
423
  return jsonify({"error": "Model workspace offline"}), 503
424
 
425
  data = request.get_json() or {}
426
+ if not data or 'prompt' not in data:
427
  return jsonify({"error": "Missing 'prompt' operational field"}), 400
428
 
429
  prompt = data['prompt']