You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.24 10:13:10