Keras训练模型报错:ValueError: as_list() is not defined on unknown TensorShape
解决TensorFlow中
ValueError: as_list() is not defined on an unknown TensorShape错误 问题描述
使用Keras构建并训练模型时触发如下错误:
ValueError: as_list() is not defined on an unknown TensorShape.
可复现的最简代码:
import numpy as np import tensorflow as tf def func(x): return x def create_dataset(vect): a = [vect] b = [vect] dataset = tf.data.Dataset.from_tensor_slices((a, b)) py_func = lambda x, y: (tf.numpy_function(func, [x], tf.float32), y) ##LINE 1 ! dataset = dataset.map(py_func) ##LINE 2 ! dataset = dataset.map(lambda data, label: (tf.expand_dims(data, axis=0), tf.expand_dims(label, axis=0))) return dataset vect = np.array([-0.18441772, -0.17321777, -0.16046143, -0.14782715, -0.13504028, -0.12179565, -0.10858154, -0.09503174, -0.08117676, -0.06964111], dtype=np.float32) example_dataset = create_dataset(vect) example_list = list(example_dataset) input_shape_m = example_list[0][0].shape #which is TensorShape([1, 10]) def DeepModel(input_shape): X_input = tf.keras.Input(input_shape) X = tf.keras.layers.Dense(10)(X_input) model = tf.keras.Model(inputs=X_input, outputs=X) return model model = DeepModel(input_shape_m[1:]) model.compile(optimizer='adam', loss='mse', metrics=['accuracy']) history = model.fit(example_dataset, epochs=5)
注:注释掉LINE1和LINE2后代码可正常运行。
错误原因
核心问题出在tf.numpy_function的特性上:
tf.numpy_function是用来包装纯Python函数的API,但它会丢失张量的静态形状信息——TensorFlow无法提前推断Python函数内部操作对张量形状的影响,所以返回的张量形状会变成unknown(即TensorShape(None))。- 后续的
tf.expand_dims虽然能在运行时给张量添加维度,但数据集里的张量依然没有可静态推断的形状。 - Keras的
model.fit在处理输入时,需要明确的静态形状来匹配模型输入层的定义,当遇到形状未知的张量时,内部会调用as_list()方法尝试将形状转为列表,而未知形状无法执行这个操作,因此触发错误。
解决方案
有两种实用的修复方式:
方法1:手动恢复张量形状
在tf.numpy_function之后,显式指定返回张量的形状,让TensorFlow重新获取静态形状信息:
# 修改LINE1和LINE2的代码 py_func = lambda x, y: (tf.reshape(tf.numpy_function(func, [x], tf.float32), x.shape), y) dataset = dataset.map(py_func)
这里用tf.reshape把返回的张量强制设置为输入x的原有形状,补全静态形状信息。
方法2:改用tf.py_function(推荐)
tf.py_function相比tf.numpy_function能更好地保留张量的元信息,只需在包装后显式确认形状:
# 替换LINE1和LINE2的代码 def func_wrapper(x, y): def py_func(x): return x result = tf.py_function(py_func, [x], tf.float32) result.set_shape(x.shape) # 显式设置形状,确保静态信息不丢失 return result, y dataset = dataset.map(func_wrapper)
通过set_shape手动确认形状后,数据集的张量形状回到可静态推断的状态,Keras就能正常处理输入了。
内容的提问来源于stack exchange,提问作者RandomUser
相关产品推荐
相关产品推荐

