Piko-9b / examples /inference_batch.py
Dexy2's picture
Rewrite model card around verified evidence; correct misattributed benchmarks and config path leak
0810902 verified
Raw
History Blame Contribute Delete
4.12 kB
#!/usr/bin/env python3
"""Batched text generation with Piko-9b.
python examples/inference_batch.py --prompts prompts.txt --batch-size 4
python examples/inference_batch.py --prompt "2+2?" --prompt "Capital of Peru?"
Reads one prompt per line from --prompts, and/or repeated --prompt flags.
Results are written as JSONL so they can be diffed between runs.
The tokenizer pads on the left, which is what batched decoder-only generation
needs; this script does not override it.
"""
from __future__ import annotations
import argparse
import json
import sys
import time
from pathlib import Path
import torch
from _common import add_common_arguments, generation_kwargs, load_model, strip_reasoning
def collect_prompts(args: argparse.Namespace) -> list[str]:
prompts: list[str] = list(args.prompt or [])
if args.prompts:
path = Path(args.prompts)
if not path.is_file():
sys.exit(f"Prompt file not found: {path}")
prompts += [line.strip() for line in path.read_text(encoding="utf-8").splitlines()]
prompts = [p for p in prompts if p]
if not prompts:
sys.exit("No prompts given. Use --prompt and/or --prompts.")
return prompts
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
add_common_arguments(parser)
parser.add_argument("--prompt", action="append", help="Repeatable.")
parser.add_argument("--prompts", help="File with one prompt per line.")
parser.add_argument("--batch-size", type=int, default=2)
parser.add_argument("--system", default="You are Piko-9, an AI assistant.")
parser.add_argument("--output", type=Path, default=None)
args = parser.parse_args()
if args.batch_size < 1:
sys.exit("--batch-size must be >= 1")
prompts = collect_prompts(args)
model, processor = load_model(args.model, args.quantization, args.dtype, args.revision)
records = []
for start in range(0, len(prompts), args.batch_size):
chunk = prompts[start : start + args.batch_size]
texts = [
processor.apply_chat_template(
(
([{"role": "system", "content": args.system}] if args.system else [])
+ [{"role": "user", "content": [{"type": "text", "text": prompt}]}]
),
add_generation_prompt=True,
tokenize=False,
)
for prompt in chunk
]
inputs = processor(text=texts, return_tensors="pt", padding=True).to(model.device)
began = time.perf_counter()
try:
with torch.inference_mode():
output = model.generate(**inputs, **generation_kwargs(args))
except torch.cuda.OutOfMemoryError:
sys.exit(
f"CUDA OOM at batch size {args.batch_size}. Retry with a smaller "
"--batch-size, or a stronger --quantization."
)
elapsed = time.perf_counter() - began
prompt_length = inputs["input_ids"].shape[1]
generated = output.shape[1] - prompt_length
for index, prompt in enumerate(chunk):
answer = processor.decode(
output[index][prompt_length:], skip_special_tokens=True
).strip()
record = {
"prompt": prompt,
"response": answer if args.show_reasoning else strip_reasoning(answer),
}
records.append(record)
print(f"--- {prompt}\n{record['response']}\n")
print(
f"[batch {start // args.batch_size + 1}] {len(chunk)} prompts, "
f"{generated} new tokens each, {elapsed:.1f}s, "
f"{len(chunk) * generated / elapsed:.1f} tok/s aggregate\n",
file=sys.stderr,
)
if args.output:
args.output.parent.mkdir(parents=True, exist_ok=True)
with args.output.open("w", encoding="utf-8") as handle:
for record in records:
handle.write(json.dumps(record, ensure_ascii=False) + "\n")
print(f"wrote {args.output}", file=sys.stderr)
if __name__ == "__main__":
main()