Piko-9b / examples /inference_transformers.py
Dexy2's picture
Rewrite model card around verified evidence; correct misattributed benchmarks and config path leak
0810902 verified
Raw
History Blame Contribute Delete
1.64 kB
#!/usr/bin/env python3
"""Text-only generation with Piko-9b.
python examples/inference_transformers.py --prompt "Explain gradient clipping."
python examples/inference_transformers.py --model ./local-copy --quantization none
"""
from __future__ import annotations
import argparse
import sys
import torch
from _common import add_common_arguments, generation_kwargs, load_model, strip_reasoning
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("--prompt", required=True)
parser.add_argument("--system", default=DEFAULT_SYSTEM)
args = parser.parse_args()
if not args.prompt.strip():
sys.exit("--prompt must not be empty.")
model, processor = load_model(args.model, args.quantization, args.dtype, args.revision)
messages = []
if args.system:
messages.append({"role": "system", "content": args.system})
messages.append({"role": "user", "content": [{"type": "text", "text": args.prompt}]})
inputs = processor.apply_chat_template(
messages,
add_generation_prompt=True,
tokenize=True,
return_dict=True,
return_tensors="pt",
).to(model.device)
with torch.inference_mode():
output = model.generate(**inputs, **generation_kwargs(args))
text = processor.decode(
output[0][inputs["input_ids"].shape[1] :], skip_special_tokens=True
).strip()
print(text if args.show_reasoning else strip_reasoning(text))
if __name__ == "__main__":
main()