You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于EleutherAI/gpt-j-6B的聊天机器人单角色回复截断方案问询

问题解决:GPT-J-6B聊天机器人仅生成指定角色回复并动态控制长度

问题背景

使用EleutherAI/gpt-j-6B开发聊天机器人时,模型在生成person alpha的回复后,会继续生成person beta的内容及大量无关文本,需求是让模型仅生成person alpha的回复后停止,同时支持动态调整回复长度。

解决方案

1. 配置自定义停止序列

将对话中其他角色的前缀(如"person beta:"、"Person beta:")设为停止触发条件,模型生成到该序列时立即停止,避免后续无关内容输出。需覆盖角色前缀的大小写变体,确保触发准确。

2. 用max_new_tokens动态控制回复长度

替换原固定计算max_length的方式,使用max_new_tokens直接指定模型新增生成的token数量,通过修改该参数即可灵活调整回复长度,无需计算原始prompt的token长度。

3. 移除固定min_length限制

原min_length会强制模型生成到指定长度,易导致多余内容,移除后模型会在触发停止序列或达到max_new_tokens时停止。

4. 兜底截断处理(可选)

若模型偶尔漏触发停止序列,可在解码后手动检查并截断到person alpha回复结束的位置。

修改后的完整代码

from transformers import GPTJForCausalLM, AutoTokenizer
import torch

# 对话prompt
prompt = """person alpha:
hi! how are you doing?

person beta:I am fine, thank you. What are you doing?

person alpha:
I am at home watching tv.

person beta:
That sounds like a lot of fun. What are you watching?

person alpha:
"""

# 加载模型与分词器
hf_name = "EleutherAI/gpt-j-6B"
model = GPTJForCausalLM.from_pretrained(
    hf_name,
    revision="float16",
    torch_dtype=torch.float16,
    low_cpu_mem_usage=True,
    use_cache=False,
    gradient_checkpointing=True
)
tokenizer = AutoTokenizer.from_pretrained(hf_name, pad_token='<|endoftext|>', eos_token='<|endoftext|>')

# 生成参数:动态调整resp_max_new_tokens控制回复长度
resp_max_new_tokens = 64
resp_temp = 0.72

# 处理prompt并转移到GPU
inputs = tokenizer(prompt, return_tensors='pt').to('cuda')

# 定义停止序列,匹配对话中其他角色的触发前缀
stop_sequences = ["\nperson beta:", "\nPerson beta:", "\nperson alpha:"]
# 将停止序列转换为token ID
stop_token_ids = [tokenizer.encode(seq, add_special_tokens=False) for seq in stop_sequences]

# 生成回复
out = model.generate(
    **inputs,
    max_new_tokens=resp_max_new_tokens,
    do_sample=True,
    top_k=35,
    top_p=0.90,
    temperature=resp_temp,
    no_repeat_ngram_size=4,
    clean_up_tokenization_spaces=True,
    eos_token_id=tokenizer.eos_token_id,
    stop_token_ids=stop_token_ids
)

# 解码并处理结果
res = tokenizer.decode(out[0], skip_special_tokens=True)
# 兜底截断:确保只保留person alpha的回复
for stop_seq in stop_sequences:
    if stop_seq in res:
        res = res[:res.index(stop_seq)]

print(res)

关键改动说明

  • 停止序列:stop_token_ids将其他角色的触发前缀转换为模型可识别的token ID,从根源上阻止无关内容生成。
  • 动态长度:修改resp_max_new_tokens即可直接控制person alpha回复的最大长度,操作更直观。
  • 兜底逻辑:手动截断作为备用方案,避免因模型偶尔的触发遗漏导致多余内容输出。

内容的提问来源于stack exchange,提问作者BlackHawk

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.22 09:09:18