Spaces:
Sleeping
Sleeping
Upload app.py
Browse files
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 |
-
|
| 169 |
-
|
|
|
|
| 170 |
if finished.all():
|
| 171 |
break
|
| 172 |
-
|
| 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 |
-
|
| 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']
|