Keras model.fit()传入class_weight数组报错:'numpy.ndarray'无'get'属性
解决Keras中class_weight传入numpy数组引发的AttributeError问题
问题说明
使用sklearn计算类别权重得到数组后,传入Keras的model.fit()时触发错误:
AttributeError: 'numpy.ndarray' object has no attribute 'get'
问题根源
Keras的class_weight参数要求接收字典类型(键为类别标签,值为对应权重),但class_weight.compute_class_weight返回的是numpy数组,直接传入会导致类型不兼容。
修复步骤
1. 将权重数组转换为字典格式
修改权重计算代码,把数组转成Keras要求的字典:
from sklearn.utils import class_weight # 获取所有类别 classes = np.unique(masks_reshaped_encoded) # 计算权重数组 class_weights_arr = class_weight.compute_class_weight( class_weight="balanced", classes=classes, y=masks_reshaped_encoded ) # 转换为字典:类别为键,权重为值 class_weights = dict(zip(classes, class_weights_arr))
2. 传入model.fit()
直接使用转换后的字典作为class_weight参数:
history=model.fit(img, msk, epochs=50, verbose=1, validation_split=0.2, shuffle=False, class_weight=class_weights)
额外提示
- 你的模型设置了
num_classes=4,但当前使用binary_crossentropy损失函数,这不符合多分类任务的要求。建议根据标签格式选择:- 若标签是独热编码格式,使用
categorical_crossentropy - 若标签是整数形式的单通道数组,使用
sparse_categorical_crossentropy
- 若标签是独热编码格式,使用
内容的提问来源于stack exchange,提问作者oliver6626
相关产品推荐
相关产品推荐

