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

深度学习训练芒果叶数据集时Epoch运行报错求助

解决Keras训练芒果叶数据集时的InvalidArgumentError问题

我在Jupyter Notebook中用深度学习模型训练芒果叶数据集,执行以下代码:

history = model.fit(
    train_ds,
    batch_size=BATCH_SIZE,
    validation_data=val_ds,
    verbose=1,
    epochs=15,
)

运行时触发了InvalidArgumentError,已经确认数据格式没问题,但搞不懂错误原因,求解决办法。

报错栈如下:

---------------------------------------------------------------------------
InvalidArgumentError                      Traceback (most recent call last)
Cell In[32], line 1
----> 1 history = model.fit(
      2     train_ds,
      3     batch_size=BATCH_SIZE,
      4     validation_data=val_ds,
      5     verbose=1,
      6     epochs=15,
      7 )

File ~\AppData\Local\Programs\Python\Python311\Lib\site-packages\keras\utils\traceback_utils.py:70, in filter_traceback.<locals>.error_handler(*args, **kwargs)
     67     filtered_tb = _process_traceback_frames(e.__traceback__)
     68     # To get the full stack trace, call:
     69     # `tf.debugging.disable_traceback_filtering()`
---> 70     raise e.with_traceback(filtered_tb) from None
     71 finally:
     72     del filtered_tb

File ~\AppData\Local\Programs\Python\Python311\Lib\site-packages\tensorflow\python\eager\execute.py:52, in quick_execute(op_name, num_outputs, inputs, attrs, ctx, name)
     50 try:
     51   ctx.ensure_initialized()
---> 52   tensors = pywrap_tfe.TFE_Py_Execute(ctx._handle, device_name, op_name,
     53                                       inputs, attrs, num_outputs)
     54 except core._NotOkStatusException as e:
     55   if name is not None:

可行的排查与解决步骤

  • 获取完整错误详情:当前报错栈被截断,先执行以下代码关闭TensorFlow的栈过滤,重新运行训练,就能看到具体错误原因(比如维度不匹配、标签格式错误、数据类型冲突等):
    import tensorflow as tf
    tf.debugging.disable_traceback_filtering()
    
  • 核对输入与模型维度:确认train_ds的输入维度和模型输入层的shape完全一致,比如模型输入是(224,224,3),数据集的图片必须统一为该尺寸,不能存在偏差。
  • 验证标签格式:分类任务下,检查标签是否符合模型要求——多分类任务是否用了独热编码,二分类是否用0/1或对应概率值,同时标签的数据类型要和模型输出层激活函数匹配(比如sigmoid对应float类型,softmax对应int或独热编码)。
  • 检查数据集预处理一致性:确认train_ds和val_ds的预处理步骤完全同步,比如是否都做了归一化、尺寸调整,避免某一数据集漏了关键步骤导致数据异常。
  • 排查批量数据异常:手动取出一个batch的数据,检查是否存在损坏图片、NaN/Inf值或不符合预期的数据:
    for x, y in train_ds.take(1):
        print(x.shape, y.shape)
        print(tf.reduce_min(x), tf.reduce_max(x))
        print(y)
    
  • 调整batch_size参数:如果train_ds已经通过batch()方法生成批量数据,直接移除model.fit中的batch_size参数,避免参数冲突。

内容的提问来源于stack exchange,提问作者Muaaz Hasni

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 17:15:57