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
相关产品推荐
相关产品推荐

