Summer-0.5B-Chat / example_load.py
tf-bao's picture
Upload Summer-0.5B-Chat
d05f7e5 verified
Raw History Blame Contribute Delete
1.84 kB
"""跑通这个 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() 直接编码问题会得到分布外输入(退化的重复内容)——
这个模型的每一行输入都以 <bos> 开头、按 <user>/<assistant>/<end> 的对话
格式训练的,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))