"""跑通这个 chat 模型 —— 只需要 torch 和 piece_tokenizer。 pip install torch pip install git+https://github.com/Ismantic/PieceTokenizer python example_load.py **这是 chat 模型,用 apply_chat_template,不是普通续写。** 用 tokenizer.encode() 直接编码问题会得到分布外输入(退化的重复内容)—— 这个模型的每一行输入都以 开头、按 // 的对话 格式训练的,apply_chat_template 已经处理好了这些约定,不用自己拼。 """ import torch from model import Qwen3ForCausalLM from tokenizer import PieceTokenizerWrapper HERE = "." tok = PieceTokenizerWrapper(HERE) model = Qwen3ForCausalLM.from_pretrained( HERE, device="cuda" if torch.cuda.is_available() else "cpu", dtype=torch.bfloat16) messages = [{"role": "user", "content": "机器翻译的基本任务是什么?"}] ids = tok.apply_chat_template(messages, tokenize=True, add_generation_prompt=True) # 贪心续写。没有 KV cache —— 每步重算前缀,短回答够用。 # **贪心只是为了演示确定性输出。实际部署建议 repetition_penalty≈1.15** # (温度 0.6),纯贪心容易陷入复读循环——这个项目自己测过,见仓库 # docs/POSTTRAIN.md 里 rp 那节的完整推导。这里为了脚本简单没接 # repetition_penalty(自己实现的 model.py 没有采样参数),想要更好的体验 # 用 example_vllm.py(vLLM 原生支持 repetition_penalty)。 x = torch.tensor([ids], device=next(model.parameters()).device) out = [] with torch.no_grad(): for _ in range(300): nxt = int(model(x)[0, -1].argmax()) if nxt in tok.stop_token_ids: break out.append(nxt) x = torch.cat([x, torch.tensor([[nxt]], device=x.device)], dim=1) print(tok.decode(out, skip_special_tokens=True))