#!/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()