Keras启用return_sequences时class_weights报错问题排查
问题原因与解决办法
核心原因
- 标签值超出类别范围:你定义
num_classes=3,合法类别索引应为0、1、2,但y_train中存在值为4的标签,导致Keras在通过class_weight查找对应权重时,找不到索引4的键,触发越界错误。 - class_weights字典不匹配:
class_weights中可能包含了键为4的权重条目,或者缺少部分合法类别的权重,与实际标签的类别范围不对应。 - 时序模式下的权重逻辑冲突:当设置
sample_weight_mode='temporal'时,模型会为每个时间步的标签应用类别权重,但如果标签本身不合法,就会直接触发索引错误。
解决步骤
- 检查并修正标签数据:使用
np.unique(y_train)查看y_train中的所有唯一标签值,确认是否存在超出[0, num_classes-1]的数值,将这些异常值修正为合法类别索引,或者根据实际标签数量调整num_classes的取值。 - 修正class_weights字典:确保
class_weights的键仅包含0、1、2这三个合法类别索引,删除多余的键(比如4),同时保证每个合法类别都有对应的权重值。 - 验证数据一致性:重新确认
y_train的标签范围与num_classes、class_weights的匹配性,避免因数据标注错误或参数定义失误导致的不兼容问题。
内容的提问来源于stack exchange,提问作者ChaddersCheese
相关产品推荐
相关产品推荐

