TensorFlow文本生成代码运行触发KeyError: '3'报错如何解决
问题原因
报错KeyError: '3'本质是传入的查询字符不在char2int字典的键范围内。
你的char2int和int2char是从训练语料shrek.txt统计生成的,字典里只存了训练文本里实际出现过的字符和对应索引。但你生成初始seed的逻辑是从大写字母+数字的字符池里随机选字符再转小写,只要随机到数字(比如这次报错的'3'),而你的训练语料里根本没出现过数字,查字典自然找不到对应值抛错。
之前代码能跑纯是概率问题——之前随机生成的seed刚好全是训练语料里存在的字母,没抽到数字而已,和你改没改代码没关系。
修复方法
选下面任意一种改就行,推荐第二种:
- 改随机字符池,去掉训练语料里不存在的数字:
把原代码里生成随机字符串的行
替换成ran = ''.join(random.choices(string.ascii_uppercase + string.digits, k = S))
这种写法仅适合你确定训练语料只有小写字母的场景,如果训练语料包含标点、特殊字符,还是可能触发同类KeyError。ran = ''.join(random.choices(string.ascii_lowercase, k = S)) - (最稳妥)直接从训练得到的词表字符里随机采样生成初始seed,从根源保证所有字符都在字典里:
删掉原来生成ran、转seed的三行代码:
在加载完S = random.randint(5, 20) ran = ''.join(random.choices(string.ascii_uppercase + string.digits, k = S)) seed = str(ran).lower()char2int字典的代码后面,加上如下seed生成逻辑:
这种写法不管你训练语料里有什么字符(标点、数字、特殊符号、大小写),生成的seed百分百能在S = random.randint(5, 20) # 直接取训练词表里的所有字符作为采样池 seed = ''.join(random.choices(list(char2int.keys()), k = S))char2int里查到索引,不会再触发同类KeyError。
补充说明
后续生成循环里的next_char是模型输出索引通过int2char映射回来的,本身就属于词表内字符,只要初始seed合法,后续迭代不会再出现同类键错误,不用修改后面的预测、拼接逻辑。
内容的提问来源于stack exchange,提问作者somethingidk
相关产品推荐
相关产品推荐

