基于Numpy手写LSTM导出TXT时遇ValueError: probabilities contain NaN错误
解决LSTM导出文本时的NaN错误及代码逻辑问题
我帮你排查下代码里的问题,这个ValueError: probabilities contain NaN错误和几个代码逻辑问题有关,咱们一步步来修复:
核心问题分析
1. NaN的直接来源
你在计算概率时使用了output[i][j] / np.sum(output[i]),如果某一行output[i]的总和为0,就会触发除以0的操作,直接产生NaN;哪怕总和接近0,也可能因为浮点精度问题出现NaN,这会让np.random.choice直接报错。
2. 文本拼接逻辑错误
每次循环你都把outputText直接赋值为单个字符,最后写入文件的只会是最后一次循环选出来的结果,完全丢失了前面的预测内容。
3. 致命的无限递归
函数最后一行return ExportText(output, data)会让函数不断调用自己,很快就会导致栈溢出,这行完全是多余的。
4. 低效的嵌套循环
用双层循环给prob赋值没必要,Numpy支持向量级别的运算,效率更高。
修正后的代码
def ExportText(output, data): outputText = "" print(len(data)) print(output.shape[0]) for i in range(output.shape[0]): # 加入极小值1e-8避免除以0,同时用向量运算替代嵌套循环 output_sum = np.sum(output[i]) prob = output[i] / (output_sum + 1e-8) # 累加预测的字符,而不是覆盖 outputText += np.random.choice(data, p=prob) with open("output.txt", "w") as text_file: text_file.write(outputText) # 可选:返回生成的文本,方便后续调用 return outputText
额外建议
你可以检查下LSTM输出的output数组,看看是不是存在全0的行——这可能是模型训练不足、激活函数选择不当(比如sigmoid输出饱和)或者数据预处理的问题,从根源上避免异常输出会让整个流程更稳定。
内容的提问来源于stack exchange,提问作者Vaibhav
相关产品推荐
相关产品推荐

