Piko-9b / docs /inference.md
Dexy2's picture
Rewrite model card around verified evidence; correct misattributed benchmarks and config path leak
0810902 verified
|
Raw
History Blame Contribute Delete
6.81 kB

Inference

Before anything else

Piko-9b must be fully resident on one device. device_map="auto" on a GPU that cannot hold the whole model offloads layers to CPU, corrupts the linear-attention state, and produces a single repeated character — with no error. Use device_map={"": 0} and pick a quantization that fits. See troubleshooting.md.

trust_remote_code is not required. torchvision is required, even for text-only use, because AutoProcessor will not construct without it.

Text generation

import torch
from transformers import AutoModelForMultimodalLM, AutoProcessor, BitsAndBytesConfig

model_id = "Dexy2/Piko-9b"

model = AutoModelForMultimodalLM.from_pretrained(
    model_id,
    dtype=torch.bfloat16,
    device_map={"": 0},
    quantization_config=BitsAndBytesConfig(
        load_in_4bit=True,
        bnb_4bit_quant_type="nf4",
        bnb_4bit_compute_dtype=torch.bfloat16,
        bnb_4bit_use_double_quant=True,
    ),
)
model.eval()
processor = AutoProcessor.from_pretrained(model_id)

messages = [
    {"role": "user", "content": [
        {"type": "text", "text": "Write a Python function that merges overlapping intervals."},
    ]},
]

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, max_new_tokens=512, do_sample=False)

text = processor.decode(output[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True)
print(text)

Reasoning traces

Piko-9b thinks before it answers, and the thinking is part of the output:

User wants Python reverse string. Simple task. Provide function. No tools needed.
</think>

```python
def reverse_string(s):
    return s[::-1]

Strip it:

```python
answer = text.rsplit("</think>", 1)[-1].strip() if "</think>" in text else text.strip()

Budget tokens for it. The trace consumes max_new_tokens. In the custom suite, one JSON-extraction case failed purely because the reasoning trace pushed the closing brace past a 512-token limit. For structured output, allow 768–1024.

Image input

messages = [
    {"role": "user", "content": [
        {"type": "image", "url": "receipt.png"},
        {"type": "text", "text": "Give the merchant and total as JSON."},
    ]},
]

url accepts a local path or an http(s) URL. Multiple images per message are supported; the chat template inserts <|vision_start|><|image_pad|><|vision_end|> for each.

Measured on the custom suite (4-bit NF4, greedy): OCR 10/10, document understanding 10/10, tables and charts 9/10. The one failure was output truncation, not misreading.

This works despite the vision tower never having been trained against this language backbone — see reports/lineage_analysis.md §4. It is an empirical result on 30 synthetic document images, not a guarantee across photographs, handwriting, or low-quality scans, none of which were tested.

Sampling

The shipped generation_config.json sets no sampling parameters, so the default is greedy. Passing temperature alone does nothing:

# no effect — do_sample is still False
model.generate(**inputs, temperature=0.7)

# correct
model.generate(**inputs, do_sample=True, temperature=0.7, top_p=0.95, max_new_tokens=512)

Greedy is the right default for extraction and evaluation, and is what every measurement here used.

Batching

The tokenizer pads left, which is what decoder-only batched generation needs. Do not change it.

texts = [
    processor.apply_chat_template(
        [{"role": "user", "content": [{"type": "text", "text": p}]}],
        add_generation_prompt=True, tokenize=False,
    )
    for p in prompts
]
inputs = processor(text=texts, return_tensors="pt", padding=True).to(model.device)

with torch.inference_mode():
    output = model.generate(**inputs, max_new_tokens=256, do_sample=False)

prompt_length = inputs["input_ids"].shape[1]
answers = [processor.decode(row[prompt_length:], skip_special_tokens=True) for row in output]

Padding makes every sequence as long as the longest, so group prompts of similar length.

Streaming

from threading import Thread
from transformers import TextIteratorStreamer

streamer = TextIteratorStreamer(processor.tokenizer, skip_prompt=True, skip_special_tokens=True)
Thread(target=model.generate, kwargs=dict(**inputs, streamer=streamer, max_new_tokens=512)).start()

for piece in streamer:
    print(piece, end="", flush=True)

Expect the reasoning trace to stream first. examples/inference_cli.py hides it behind a [thinking…] indicator.

Long context

max_position_embeddings is 262,144 with no RoPE scaling. Because 24 of 32 layers use linear attention with a fixed-size state, the KV cache grows far more slowly than in a dense transformer.

Measured: a 14,429-token prompt was processed in 3.5 s and the planted fact was retrieved correctly. Needle tests at 2K, 8K and 32K filler tokens all passed, at depths from 0.1 to 0.9. Beyond 32K is untested — treat 262K as a configuration value, not a validated capability.

System prompts

The model does not self-identify as Piko-9 without one; asked what it is, the published checkpoint answers "I am Wraith, an AI model." If you need a consistent identity, set it explicitly:

messages = [
    {"role": "system", "content": "You are Piko-9, an AI assistant. Be accurate and concise."},
    ...
]

Note that the chat template raises an exception if a system message contains an image.

Ready-made scripts

Script Purpose
examples/inference_transformers.py Text generation
examples/inference_multimodal.py Image + text, with input validation
examples/inference_batch.py Batched generation to JSONL
examples/inference_cli.py Interactive chat with streaming

All four validate VRAM before loading, refuse to enable CPU offload, and give actionable errors for missing torchvision, missing bitsandbytes, and OOM.

python examples/inference_transformers.py --prompt "Explain gradient clipping." --quantization 4bit
python examples/inference_multimodal.py --image receipt.png --prompt "Total as JSON?"

Serving

vLLM and SGLang were not tested for this release. Support depends on the engine implementing the qwen3_5 hybrid architecture and its vision tower. Verify with a short generation before trusting a served deployment, and watch specifically for the degenerate-output signature.