#!/usr/bin/env python3 """Interactive chat with Piko-9b, with streaming output. python examples/inference_cli.py python examples/inference_cli.py --quantization none --temperature 0.7 Commands inside the session: /image attach an image to the next message /system replace the system prompt and reset the conversation /reset clear the conversation /exit quit """ from __future__ import annotations import argparse import sys from pathlib import Path from threading import Thread from typing import Any import torch from _common import add_common_arguments, generation_kwargs, load_model DEFAULT_SYSTEM = "You are Piko-9, an AI assistant. Be accurate, direct, and concise." def main() -> None: parser = argparse.ArgumentParser(description=__doc__) add_common_arguments(parser) parser.add_argument("--system", default=DEFAULT_SYSTEM) args = parser.parse_args() model, processor = load_model(args.model, args.quantization, args.dtype, args.revision) try: from transformers import TextIteratorStreamer except ImportError: sys.exit("TextIteratorStreamer unavailable; upgrade transformers.") system = args.system history: list[dict[str, Any]] = [] pending_image: str | None = None print("Piko-9b ready. /image , /system , /reset, /exit\n") while True: try: line = input(">>> ").strip() except (EOFError, KeyboardInterrupt): print() break if not line: continue if line in ("/exit", "/quit"): break if line == "/reset": history.clear() pending_image = None print("[conversation cleared]\n") continue if line.startswith("/system "): system = line[len("/system ") :].strip() history.clear() print("[system prompt set, conversation cleared]\n") continue if line.startswith("/image "): candidate = Path(line[len("/image ") :].strip()).expanduser() if not candidate.is_file(): print(f"[no such file: {candidate}]\n") continue pending_image = str(candidate.resolve()) print(f"[attached {candidate.name}; it will go with your next message]\n") continue content: list[dict[str, str]] = [] if pending_image: content.append({"type": "image", "url": pending_image}) content.append({"type": "text", "text": line}) history.append({"role": "user", "content": content}) pending_image = None messages = ([{"role": "system", "content": system}] if system else []) + history try: inputs = processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", ).to(model.device) except ImportError as exc: if "orchvision" in str(exc): print("[image input needs torchvision: pip install torchvision]\n") history.pop() continue raise streamer = TextIteratorStreamer( processor.tokenizer, skip_prompt=True, skip_special_tokens=True ) thread = Thread( target=_generate, args=(model, inputs, streamer, generation_kwargs(args)), daemon=True, ) thread.start() pieces: list[str] = [] in_reasoning = False for piece in streamer: pieces.append(piece) joined = "".join(pieces) if not args.show_reasoning: # Suppress the ... span unless asked for. if "" in joined and "" not in joined: if not in_reasoning: print("[thinking…]", end="", flush=True) in_reasoning = True continue if in_reasoning and "" in joined: in_reasoning = False print("\r" + " " * 12 + "\r", end="", flush=True) piece = joined.rsplit("", 1)[1] print(piece, end="", flush=True) thread.join() print("\n") history.append( {"role": "assistant", "content": [{"type": "text", "text": "".join(pieces)}]} ) def _generate(model: Any, inputs: Any, streamer: Any, kwargs: dict[str, Any]) -> None: try: with torch.inference_mode(): model.generate(**inputs, streamer=streamer, **kwargs) except torch.cuda.OutOfMemoryError: print("\n[CUDA out of memory — try /reset, a shorter prompt, or 4-bit]", flush=True) if __name__ == "__main__": main()