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

如何解决HappyTransformer中「张量a(1024)与张量b尺寸不匹配」错误

问题解决:RuntimeError张量尺寸不匹配

错误原因

  1. 按字符数拆分文本而非token数:GPT-2的上下文限制基于token数(不是字符数),当前按1024字符拆分的文本,转换为token后可能超出模型允许的长度,引发输入张量尺寸异常。
  2. max_length参数设置不当:GENSettings中的max_length指生成后文本的总token数(包含前缀、输入文本和生成内容)。当prefix + chunk的token数接近1024时,再尝试生成到1024长度会触发维度不匹配。

修复方案

1. 按token数拆分文本

用GPT-2分词器将文本转换为token,按token数拆分后转回文本,确保每个输入块(含前缀)的token数不超过模型最大上下文的安全范围。

2. 改用max_new_tokens控制生成长度

该参数直接控制新增生成的token数量,比max_length更直观,避免总长度超出模型限制。

3. 加载本地保存的模型(可选)

避免重复从远程仓库加载模型,直接使用已保存到本地的版本。

修改后的完整代码

from transformers import GPT2Tokenizer, GPT2LMHeadModel
from happytransformer import HappyGeneration, GENSettings
import torch

model_name = "gpt2"
save_path = "/home/ubuntu/storage1/various_transformer_models/gpt2"

# 加载并保存模型(本地已存在可跳过保存步骤)
tokenizer = GPT2Tokenizer.from_pretrained(model_name)
model = GPT2LMHeadModel.from_pretrained(model_name)
tokenizer.save_pretrained(save_path)
model.save_pretrained(save_path)

# 加载本地模型初始化HappyGeneration
happy_gen = HappyGeneration("GPT-2", save_path)

# 用max_new_tokens控制生成的新token数,避免总长度超限
args = GENSettings(num_beams=5, max_new_tokens=200)

mytext = "This sentence has bad grammar. This is a very long sentence that exceeds the maximum length of 512 tokens. Therefore, we need to split it into smaller chunks and process each chunk separately."
prefix = "grammar: "

# 按token数拆分文本(预留前缀的token数,确保总输入token数不超过900,给生成留空间)
max_input_tokens = 900
prefix_tokens = tokenizer.encode(prefix, add_special_tokens=False)
prefix_len = len(prefix_tokens)
text_tokens = tokenizer.encode(mytext, add_special_tokens=False)

# 拆分token块
chunks_tokens = []
for i in range(0, len(text_tokens), max_input_tokens - prefix_len):
    chunk = text_tokens[i:i + (max_input_tokens - prefix_len)]
    chunks_tokens.append(chunk)

# 将token块转回文本
chunks = [tokenizer.decode(chunk) for chunk in chunks_tokens]

# 逐块处理并收集结果
results = []
for chunk in chunks:
    input_text = prefix + chunk
    result = happy_gen.generate_text(input_text, args=args)
    # 移除前缀,只保留修正后的内容
    corrected_text = result.text.replace(prefix, "").strip()
    results.append(corrected_text)

# 拼接最终结果
output_text = " ".join(results)
print(output_text)

关键修改点说明

  • token级拆分:确保每个输入块(含前缀)的token数不超过900,预留200个token的生成空间,总长度不超过GPT-2的1024上限。
  • max_new_tokens替代max_length:明确控制生成的新token数量,避免总长度计算错误导致的张量尺寸问题。
  • 移除前缀:生成结果会包含输入的前缀,处理时将其剔除,只保留修正后的文本。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 03:05:19