误用LinearClassifier致TensorFlow线性回归模型Assertion failed错误求助
汽车价格预测线性回归模型训练错误排查与解决
问题描述
使用TensorFlow 2.0构建汽车价格预测的线性回归模型,操作流程如下:
- 导入包含汽车特征(ID、Price、Manufacturer等)的CSV数据集并转为Pandas DataFrame
- 重命名部分列名,移除无关特征
- 分离特征数据与预测目标Price(作为标签)
- 定义类别特征与数值特征并构建feature_columns
- 创建输入函数,初始化模型并执行训练与评估
运行时触发Assertion failed错误,错误信息显示Labels must be <= n_classes - 1,标签值(如19444、56509等)远大于1,无法完成训练。
原代码
# IMPORTS import os import numpy as np import tensorflow as tf import pandas as pd # IMPORT CSV FILES AS PANDASDB dfTrain = pd.read_csv('carPriceTrain.csv') dfEval = pd.read_csv('carPriceEval.csv') # RENAME COLUMN NAMES dfTrain = dfTrain.rename(columns={"Gear box type": "gearBoxType"}) dfEval = dfTrain.rename(columns={"Gear box type": "gearBoxType"}) dfTrain = dfTrain.rename(columns={"Leather interior": "leatherInterior"}) dfEval = dfTrain.rename(columns={"Leather interior": "leatherInterior"}) # REMOVE EXTRANEOUS CAR INFORMATION dfTrain.pop('ID') dfTrain.pop('Levy') dfTrain.pop('Category') dfTrain.pop('Fuel type') dfTrain.pop('Drive wheels') dfTrain.pop('leatherInterior') dfTrain.pop('gearBoxType') dfTrain.pop('Doors') dfTrain.pop('Wheel') dfTrain.pop('Color') dfTrain.pop('Airbags') dfEval.pop('ID') dfEval.pop('Levy') dfEval.pop('Category') dfEval.pop('Fuel type') dfEval.pop('Drive wheels') dfEval.pop('leatherInterior') dfEval.pop('gearBoxType') dfEval.pop('Doors') dfEval.pop('Wheel') dfEval.pop('Color') dfEval.pop('Airbags') # POP PRICE (PREDICTING VAL) yTrain = dfTrain.pop('Price') yEval = dfEval.pop('Price') # CREATE COLUMNS NAMES CATEGORICAL_COLUMNS = ['Manufacturer', 'Model'] NUMERIC_COLUMNS = ['Prod. year', 'Mileage', 'Cylinders', 'Engine volume'] # FEATURE COLUMNS FOR CATERGORICAL featureColumns = [] for featureName in CATEGORICAL_COLUMNS: vocabulary = dfTrain[featureName].unique() featureColumns.append(tf.feature_column.categorical_column_with_vocabulary_list(featureName, vocabulary)) # FEATURE COLUMNS FOR NUMERIC for featureName in NUMERIC_COLUMNS: featureColumns.append(tf.feature_column.numeric_column(featureName, dtype=tf.float32)) # CREATING INPUT FN def makeInputFN(dfData, dfLabel, nEpochs=10, shuffle=True, batchSize=32): def inputFN(): ds = tf.data.Dataset.from_tensor_slices((dict(dfData), dfLabel)) if shuffle: ds = ds.shuffle(1000) ds = ds.batch(batchSize).repeat(nEpochs) return ds return inputFN trainInputFN = makeInputFN(dfTrain, yTrain) evalInputFN = makeInputFN(dfEval, yEval, nEpochs=1, shuffle=False) linearEst = tf.estimator.LinearClassifier(feature_columns=featureColumns) linearEst.train(trainInputFN) result = linearEst.evaluate(evalInputFN)
错误信息
WARNING:tensorflow:Using temporary folder as model directory: /var/folders/kq/823s5gqs0ds7wqckr9sxfcnm0000gn/T/tmp22bryhkx WARNING:tensorflow:From /Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/training/training_util.py:396: Variable.initialized_value (from tensorflow.python.ops.variables) is deprecated and will be removed in a future version. Instructions for updating: Use Variable.read_value. Variables in 2.X are initialized automatically both in eager and graph (inside tf.defun) contexts. WARNING:tensorflow:From /Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/keras/optimizers/optimizer_v2/ftrl.py:153: calling Constant.__init__ (from tensorflow.python.ops.init_ops) with dtype is deprecated and will be removed in a future version. Instructions for updating: Call initializer instance with the dtype argument instead of passing it to the constructor 2022-07-30 15:22:47.534048: I tensorflow/core/platform/cpu_feature_guard.cc:193] This TensorFlow binary is optimized with oneAPI Deep Neural Network Library (oneDNN) to use the following CPU instructions in performance-critical operations: AVX2 FMA To enable them in other operations, rebuild TensorFlow with the appropriate compiler flags. 2022-07-30 15:22:47.549419: I tensorflow/compiler/mlir/mlir_graph_optimization_pass.cc:354] MLIR V1 optimization pass is not enabled Traceback (most recent call last): File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/client/session.py", line 1377, in _do_call return fn(*args) File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/client/session.py", line 1360, in _run_fn return self._call_tf_sessionrun(options, feed_dict, fetch_list, File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/client/session.py", line 1453, in _call_tf_sessionrun return tf_session.TF_SessionRun_wrapper(self._session, options, feed_dict, tensorflow.python.framework.errors_impl.InvalidArgumentError: assertion failed: [Labels must be <= n_classes - 1] [Condition x <= y did not hold element-wise:] [x (head/losses/Cast:0) = ] [[19444][56509][8781]...] [y (head/losses/check_label_range/Const:0) = ] [1] [[{{function_node head_losses_check_label_range_assert_less_equal_Assert_AssertGuard_false_667}}{{node Assert}}]] During handling of the above exception, another exception occurred: Traceback (most recent call last): File "/Users/anirud/CarPricePredictor/main.py", line 75, in <module> linearEst.train(trainInputFN) File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow_estimator/python/estimator/estimator.py", line 360, in train loss = self._train_model(input_fn, hooks, saving_listeners) File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow_estimator/python/estimator/estimator.py", line 1186, in _train_model return self._train_model_default(input_fn, hooks, saving_listeners) File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow_estimator/python/estimator/estimator.py", line 1217, in _train_model_default return self._train_with_estimator_spec(estimator_spec, worker_hooks, File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow_estimator/python/estimator/estimator.py", line 1533, in _train_with_estimator_spec _, loss = mon_sess.run([estimator_spec.train_op, estimator_spec.loss]) File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/training/monitored_session.py", line 782, in run return self._sess.run( File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/training/monitored_session.py", line 1311, in run return self._sess.run( File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/training/monitored_session.py", line 1416, in run raise six.reraise(*original_exc_info) File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/six.py", line 719, in reraise raise value File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/training/monitored_session.py", line 1401, in run return self._sess.run(*args, **kwargs) File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/training/monitored_session.py", line 1469, in run outputs = _WrappedSession.run( File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/training/monitored_session.py", line 1232, in run return self._sess.run(*args, **kwargs) File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/client/session.py", line 967, in run result = self._run(None, fetches, feed_dict, options_ptr, File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/client/session.py", line 1190, in _run results = self._do_run(handle, final_targets, final_fetches, File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/client/session.py", line 1370, in _do_run return self._do_call(_run_fn, feeds, fetches, targets, options, File "/Library/Frameworks/Python.framework/Versions/3.9/lib/python3.9/site-packages/tensorflow/python/client/session.py", line 1396, in _do_call raise type(e)(node_def, op, message) # pylint: disable=no-value-for-parameter tensorflow.python.framework.errors_impl.InvalidArgumentError: Graph execution error: assertion failed: [Labels must be <= n_classes - 1] [Condition x <= y did not hold element-wise:] [x (head/losses/Cast:0) = ] [[19444][56509][8781]...] [y (head/losses/check_label_range/Const:0) = ] [1] [[{{node Assert}}]]
错误原因
- 模型类型错误:核心问题是使用了
tf.estimator.LinearClassifier(线性分类器)处理回归任务。分类器默认是二分类模式(n_classes=2),强制要求标签值为0或1,但你的标签是连续的汽车价格,远超出分类器的标签范围,触发断言错误。 - 数据泄漏问题:代码中重命名和清理dfEval时,错误地将dfTrain赋值给dfEval(如
dfEval = dfTrain.rename(...)),导致训练集和测试集完全相同,后续评估毫无意义,还会引入数据泄漏。
解决方法
- 替换为回归模型:将
LinearClassifier替换为tf.estimator.LinearRegressor,这是TensorFlow专门用于线性回归任务的估算器,不需要标签满足分类范围要求。 - 修正数据处理逻辑:确保dfEval的操作基于自身,而非复制dfTrain的结果。
- 可选优化:对数值特征做归一化/标准化处理,回归模型对特征尺度敏感,归一化后能提升模型收敛速度和预测精度。
修正后的代码
# IMPORTS import os import numpy as np import tensorflow as tf import pandas as pd # IMPORT CSV FILES AS PANDASDB dfTrain = pd.read_csv('carPriceTrain.csv') dfEval = pd.read_csv('carPriceEval.csv') # RENAME COLUMN NAMES - 修正dfEval的赋值逻辑 dfTrain = dfTrain.rename(columns={"Gear box type": "gearBoxType", "Leather interior": "leatherInterior"}) dfEval = dfEval.rename(columns={"Gear box type": "gearBoxType", "Leather interior": "leatherInterior"}) # REMOVE EXTRANEOUS CAR INFORMATION - 分别处理训练集和测试集 drop_cols = ['ID', 'Levy', 'Category', 'Fuel type', 'Drive wheels', 'leatherInterior', 'gearBoxType', 'Doors', 'Wheel', 'Color', 'Airbags'] dfTrain = dfTrain.drop(columns=drop_cols) dfEval = dfEval.drop(columns=drop_cols) # POP PRICE (PREDICTING VAL) yTrain = dfTrain.pop('Price') yEval = dfEval.pop('Price') # CREATE COLUMNS NAMES CATEGORICAL_COLUMNS = ['Manufacturer', 'Model'] NUMERIC_COLUMNS = ['Prod. year', 'Mileage', 'Cylinders', 'Engine volume'] # FEATURE COLUMNS FOR CATERGORICAL featureColumns = [] for featureName in CATEGORICAL_COLUMNS: vocabulary = dfTrain[featureName].unique() # 分类特征需要转为indicator列才能输入回归模型 cat_col = tf.feature_column.categorical_column_with_vocabulary_list(featureName, vocabulary) featureColumns.append(tf.feature_column.indicator_column(cat_col)) # FEATURE COLUMNS FOR NUMERIC - 添加归一化处理 def normalize_numeric_column(col_name): return tf.feature_column.numeric_column( col_name, dtype=tf.float32, normalizer_fn=lambda x: (x - dfTrain[col_name].mean()) / dfTrain[col_name].std() ) for featureName in NUMERIC_COLUMNS: featureColumns.append(normalize_numeric_column(featureName)) # CREATING INPUT FN def makeInputFN(dfData, dfLabel, nEpochs=10, shuffle=True, batchSize=32): def inputFN(): ds = tf.data.Dataset.from_tensor_slices((dict(dfData), dfLabel)) if shuffle: ds = ds.shuffle(1000) ds = ds.batch(batchSize).repeat(nEpochs) return ds return inputFN trainInputFN = makeInputFN(dfTrain, yTrain) evalInputFN = makeInputFN(dfEval, yEval, nEpochs=1, shuffle=False) # 替换为线性回归器 linearEst = tf.estimator.LinearRegressor(feature_columns=featureColumns) linearEst.train(trainInputFN) result = linearEst.evaluate(evalInputFN) print("评估结果:", result)
关键修改说明
- 将
LinearClassifier改为LinearRegressor,适配回归任务。 - 修正dfEval的数据处理逻辑,避免数据泄漏。
- 分类特征通过
indicator_column转为独热编码,符合回归模型的输入要求。 - 对数值特征添加归一化处理,提升模型性能。
内容的提问来源于stack exchange,提问作者AnirudL
相关产品推荐
相关产品推荐

