如何解决HappyTransformer中「张量a(1024)与张量b尺寸不匹配」错误
问题解决:RuntimeError张量尺寸不匹配
错误原因
- 按字符数拆分文本而非token数:GPT-2的上下文限制基于token数(不是字符数),当前按1024字符拆分的文本,转换为token后可能超出模型允许的长度,引发输入张量尺寸异常。
- 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
相关产品推荐
相关产品推荐

