# 配套文章：04-context-and-inference.md
# 本轮新增示例未执行；环境和运行方式见上级 GUIDE.md。

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

model_id = "Qwen/Qwen2.5-0.5B-Instruct"
tokenizer = AutoTokenizer.from_pretrained(model_id)
model = AutoModelForCausalLM.from_pretrained(
    model_id, torch_dtype=torch.float32
).eval()
messages = [
    {"role": "system", "content": "根据给定资料回答；缺少依据时说明无法判断。"},
    {"role": "user", "content": "我的订单号是 A100，请复述这个订单号。"},
]
inputs = tokenizer.apply_chat_template(
    messages, tokenize=True, add_generation_prompt=True,
    return_dict=True, return_tensors="pt",
)

with torch.inference_mode():
    # Prefill：完整输入只处理一次，返回每层缓存和所有输入位置的 logits。
    first = model(**inputs, use_cache=True)
    cache = first.past_key_values
    next_id = first.logits[:, -1, :].argmax(dim=-1, keepdim=True)
    print("Prefill 后缓存位置数：", cache.get_seq_length())
    # 此时 next_id 只是选出来了，还未送入模型，所以缓存不包含它。
    generated = [next_id.item()]
    mask = inputs["attention_mask"]
    eos = model.generation_config.eos_token_id
    stop_ids = set(eos if isinstance(eos, list) else [eos])

    # 最多生成 16 个 Token。每轮只提交一个未处理的新 ID；旧前缀由缓存提供。
    for _ in range(15):
        if generated[-1] in stop_ids:
            break
        mask = torch.cat([mask, mask.new_ones((1, 1))], dim=-1)
        step = model(input_ids=next_id, attention_mask=mask,
                     past_key_values=cache, use_cache=True)
        cache = step.past_key_values
        next_id = step.logits[:, -1, :].argmax(dim=-1, keepdim=True)
        generated.append(next_id.item())
        print("追加处理一个位置后：", cache.get_seq_length())

# 对新生成 ID 整体解码，不把原始提示一起当作回答打印。
print("生成文本：", tokenizer.decode(generated, skip_special_tokens=True))
print("结束类型：", "结束标记" if generated[-1] in stop_ids else "达到教学长度上限")
