文本生成

大语言模型的官方推理入口是 PreTrainedModel.generate()(Pipeline 的 text-generation 底层也调用它)。生成配置集中在 GenerationConfig。务必显式设置 max_new_tokens,默认往往过短。


最小 generate()

from transformers import AutoModelForCausalLM, AutoTokenizer

model_id = "distilbert/distilgpt2"
tokenizer = AutoTokenizer.from_pretrained(model_id)
model = AutoModelForCausalLM.from_pretrained(
    model_id, device_map="auto", dtype="auto"
)
if tokenizer.pad_token is None:
    tokenizer.pad_token = tokenizer.eos_token

inputs = tokenizer("Once upon a time,", return_tensors="pt").to(model.device)
out = model.generate(**inputs, max_new_tokens=40)
print(tokenizer.decode(out[0], skip_special_tokens=True))

device_map="auto"dtype="auto" 是官方推荐的省心加载方式。


贪婪 vs 采样

策略关键参数行为
贪婪do_sample=False(默认常如此)每步取最高概率 token,可复现、易重复
采样do_sample=True按分布抽样,更有变化
温度temperature(需采样)越大越随机,常用 0.7–1.0
nucleustop_p只从累积概率达 p 的集合里抽
top-ktop_k只保留分数最高的 k 个
greedy = model.generate(**inputs, max_new_tokens=30, do_sample=False)
sample = model.generate(
    **inputs,
    max_new_tokens=30,
    do_sample=True,
    temperature=0.8,
    top_p=0.9,
)
print(tokenizer.decode(greedy[0], skip_special_tokens=True))
print(tokenizer.decode(sample[0], skip_special_tokens=True))

聊天助手一般 开采样;抽题目、格式化 JSON 时用贪婪或低温度。


指令模型(Qwen)

中文场景优先 Qwen / DeepSeek 的 Instruct 权重,而不是英文 DistilGPT2:

from transformers import AutoModelForCausalLM, AutoTokenizer

model_id = "Qwen/Qwen2.5-0.5B-Instruct"
tok = AutoTokenizer.from_pretrained(model_id)
model = AutoModelForCausalLM.from_pretrained(
    model_id, device_map="auto", dtype="auto"
)
messages = [
    {"role": "user", "content": "用三句话介绍 Hugging Face Hub。"},
]
prompt = tok.apply_chat_template(
    messages, add_generation_prompt=True, return_tensors="pt"
).to(model.device)
out = model.generate(prompt, max_new_tokens=128, do_sample=True, temperature=0.7)
print(tok.decode(out[0][prompt.shape[-1]:], skip_special_tokens=True))

切片 out[0][prompt.shape[-1]:] 只打印新生成的 token,避免把整段 prompt 回显。

0.5B 在 4–8GB 显存上通常能跑;纯 CPU 会慢。更大的 Qwen / DeepSeek 先看模型卡的显存表。


GenerationConfig 与常见旋钮

from transformers import GenerationConfig

cfg = GenerationConfig(
    max_new_tokens=64,
    do_sample=True,
    temperature=0.7,
    top_p=0.9,
    repetition_penalty=1.1,
)
out = model.generate(**inputs, generation_config=cfg)
参数含义
max_new_tokens最多新生成多少 token
max_length总长(含输入),易与上一项搞混,新手优先用前者
eos_token_id遇到则停
repetition_penalty>1 惩罚重复 n-gram

也可 model.generation_config.save_pretrained("./gen-cfg") 后随模型上传。


和 Ollama 怎么选

本站 Ollama 把量化、对话模板、HTTP API 打成一条命令,适合「本机聊天」。Transformers 的 generate() 适合对照论文、改采样、接自定义 logits 处理器、以及后面的 LoRA。两条链路可以并存:日常对话用 Ollama,实验用本课。


下一步

评论