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

TensorFlow训练随机报错:期望int64张量却得到float类型求助

问题

使用M1 MacBook Pro,搭配Python 3.11与TensorFlow 2.15.0构建带L2正则的全连接分类模型,模型代码如下:

from tensorflow.keras.models import Sequential, load_model
from tensorflow.keras.layers import Input, Dense, Dropout, BatchNormalization
from tensorflow.keras.optimizers import Adam
import tensorflow.keras.regularizers as regularizers
from tensorflow.keras.callbacks import ModelCheckpoint

regularizer = regularizers.l2(REGULARIZATION_RATE)

model = Sequential(
    [
        Input(shape=[input_shape]),
        
        Dense(64, activation="relu", kernel_regularizer=regularizer),
        Dense(64, activation="relu", kernel_regularizer=regularizer),
        Dense(32, activation="relu", kernel_regularizer=regularizer),

        Dense(units=output_shape, activation="softmax"),
    ]
)

optimizer = Adam(learning_rate=LEARNING_RATE)

checkpointer_callback = ModelCheckpoint(
    filepath=f"models/{model_name}/{model_name}.hdf5",
    monitor="val_loss",
    verbose=False,
    save_best_only=True,
)

model.summary()

model.compile(
    optimizer=optimizer, loss="categorical_crossentropy", metrics=["accuracy"]
)

训练代码:

history = model.fit(
    x_train,
    y_train,
    batch_size=BATCH_SIZE,
    epochs=NUM_EPOCHS,
    validation_split=VALIDATION_SPLIT,
    verbose=VERBOSE,
    callbacks=[
        checkpointer_callback,
        csv_logger_callback,
        early_stopping_callback,
        tensorboard_callback,
    ],
)

调用model.fit时会随机在某一epoch抛出错误:

Expected tensor of type int64 but got type float [[{{node Equal}}]] [Op:__inference_train_function_1276164]

完整报错栈:

---------------------------------------------------------------------------
InvalidArgumentError                      Traceback (most recent call last)
Cell In[47], line 9
      5 print(f"Training model {model_name}...\n")
      6 start_time = time.time()
----> 9 history = model.fit(
     10     x_train,
     11     y_train,
     12     batch_size=BATCH_SIZE,
     13     epochs=NUM_EPOCHS,
     14     validation_split=VALIDATION_SPLIT,
     15     verbose=VERBOSE,
     16     callbacks=[
     17         checkpointer_callback,
     18         csv_logger_callback,
     19         early_stopping_callback,
     20         tensorboard_callback,
     21     ],
     22 )
     25 training_duration = get_timestamp(time.time() - start_time)
     27 print("\n-------------\n")
File ~/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/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 ~/miniconda3/envs/kinected/lib/python3.11/site-packages/tensorflow/python/eager/execute.py:53, in quick_execute(op_name, num_outputs, inputs, attrs, ctx, name)
     51 try:
     52   ctx.ensure_initialized()
---> 53   tensors = pywrap_tfe.TFE_Py_Execute(ctx._handle, device_name, op_name,
     54                                       inputs, attrs, num_outputs)
     55 except core._NotOkStatusException as e:
     56   if name is not None:

InvalidArgumentError: Graph execution error:

Detected at node Equal defined at (most recent call last):
  File "<frozen runpy>", line 198, in _run_module_as_main

  File "<frozen runpy>", line 88, in _run_code

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/ipykernel_launcher.py", line 17, in <module>

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/traitlets/config/application.py", line 1077, in launch_instance

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/ipykernel/kernelapp.py", line 739, in start

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/tornado/platform/asyncio.py", line 195, in start

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/asyncio/base_events.py", line 607, in run_forever

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/asyncio/base_events.py", line 1922, in _run_once

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/asyncio/events.py", line 80, in _run

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/ipykernel/kernelbase.py", line 529, in dispatch_queue

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/ipykernel/kernelbase.py", line 518, in process_one

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/ipykernel/kernelbase.py", line 424, in dispatch_shell

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/ipykernel/kernelbase.py", line 766, in execute_request

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/ipykernel/ipkernel.py", line 429, in do_execute

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/ipykernel/zmqshell.py", line 549, in run_cell

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/IPython/core/interactiveshell.py", line 3048, in run_cell

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/IPython/core/interactiveshell.py", line 3103, in _run_cell

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/IPython/core/async_helpers.py", line 129, in _pseudo_sync_runner

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/IPython/core/interactiveshell.py", line 3308, in run_cell_async

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/IPython/core/interactiveshell.py", line 3490, in run_ast_nodes

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/IPython/core/interactiveshell.py", line 3550, in run_code

  File "/var/folders/85/04dxvd3x7wsbz88m2jf0z4000000gn/T/ipykernel_6585/1245095393.py", line 9, in <module>

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/utils/traceback_utils.py", line 65, in error_handler

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/engine/training.py", line 1807, in fit

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/engine/training.py", line 1401, in train_function

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/engine/training.py", line 1384, in step_function

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/engine/training.py", line 1373, in run_step

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/engine/training.py", line 1155, in train_step

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/engine/training.py", line 1249, in compute_metrics

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/engine/compile_utils.py", line 620, in update_state

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/utils/metrics_utils.py", line 77, in decorated

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/metrics/base_metric.py", line 140, in update_state_fn

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/metrics/base_metric.py", line 723, in update_state

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/metrics/accuracy_metrics.py", line 426, in categorical_accuracy

  File "/Users/louislecouturier/miniconda3/envs/kinected/lib/python3.11/site-packages/keras/src/utils/metrics_utils.py", line 969, in sparse_categorical_matches

Expected tensor of type int64 but got type float
     [[{{node Equal}}]] [Op:__inference_train_function_1276164]

输入数据均为float32类型,标签已做独热编码,不清楚错误原因,想了解哪里需要int64类型数据及解决方法。

解决方案

错误原因

从报错栈的sparse_categorical_matches可以看出,Keras在计算准确率时误将独热编码的标签当成了稀疏标签(整数类型的类别索引)处理,导致类型不匹配。虽然使用的是categorical_crossentropy损失,但默认的accuracy指标在M1芯片的TensorFlow环境中可能存在类型推断bug,自动切换为稀疏模式。

解决步骤

  • 显式指定匹配的准确率指标
    不要用默认的["accuracy"],明确使用CategoricalAccuracy,确保和独热编码标签兼容:

    model.compile(
        optimizer=optimizer, 
        loss="categorical_crossentropy", 
        metrics=[tf.keras.metrics.CategoricalAccuracy()]
    )
    
  • 统一标签数据类型
    确认独热编码后的标签类型与输入数据一致(float32),避免混入整数类型:

    y_train = y_train.astype("float32")
    # 若有独立验证集,同样处理
    y_val = y_val.astype("float32")
    
  • 临时关闭动态图优化(可选)
    若上述方法无效,可强制关闭TensorFlow的动态图执行,避免类型推断错误:

    tf.config.run_functions_eagerly(True)
    

    注意:此操作会降低训练速度,仅作为临时排查手段。

  • 升级TensorFlow版本
    TensorFlow 2.15.0在M1芯片上存在部分类型匹配已知问题,升级到最新稳定版(如2.16+)可能直接修复:

    pip install --upgrade tensorflow
    

内容的提问来源于Stack Exchange,提问作者Louis Lecouturier

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 14:44:49