使用TensorFlow生成器构建数据集时遇as_list()未定义ValueError
问题:TensorFlow拟合模型时报错ValueError: as_list() is not defined on an unknown TensorShape
通过生成器创建TensorFlow数据集并拟合简单API模型,此前曾通过将张量重塑为shape=(1,)解决同类错误,但修改生成器逻辑后错误复现,且无法定位数据集中的未知形状张量。
数据集样例及输出
运行代码查看单条数据:
for example in ds.take(1): print(example[0], example[1])
输出结果:
{'LoB': <tf.Tensor: shape=(1,), dtype=int32, numpy=array([5], dtype=int32)>, 'cc': <tf.Tensor: shape=(1,), dtype=int32, numpy=array([17], dtype=int32)>, 'inj_part': <tf.Tensor: shape=(1,), dtype=int32, numpy=array([41], dtype=int32)>, 'age': <tf.Tensor: shape=(1,), dtype=float32, numpy=array([2.3495796], dtype=float32)>, 'RepDel': <tf.Tensor: shape=(1,), dtype=float32, numpy=array([-0.26196158], dtype=float32)>, 'dev_year_predictor': <tf.Tensor: shape=(1,), dtype=float32, numpy=array([-1.2747549], dtype=float32)>, 'cum_loss': <tf.Tensor: shape=(1,), dtype=float32, numpy=array([1.8005615], dtype=float32)>} tf.Tensor([5.], shape=(1,), dtype=float32)
报错的模型代码
运行以下简单模型时触发错误:
inputs = layers.Input(shape=(1, ), name="LoB") output = layers.Dense(1, activation="linear")(inputs) test = models.Model(inputs=inputs, outputs=output) test.compile(loss="mse", optimizer="sgd") test.fit(ds, epochs=5, verbose=True)
错误信息:ValueError: as_list() is not defined on an unknown TensorShape
数据集构建简化代码
def create_tensor(seq): LoB=seq["LoB"].values[0] target=seq[target].values[-1] return {"LoB": LoB}, target def pad_and_format(seq): x, y = seq x["LoB"] = tf.reshape(x["LoB"], shape=(1,)) y = tf.reshape(y, shape=(1,)) return x,y def generator(): for i in range(train_df["ClNr_sub"].max()+1): seq=train_df[train_df["ClNr_sub"] == i] seq=create_tensor(seq) seq=pad_and_format(seq) yield seq ds = tf.data.Dataset.from_generator(generator, output_types=({'LoB': tf.int32}, tf.float32))
解决方案
问题根源在于tf.data.Dataset.from_generator仅指定了output_types,未明确output_shapes,导致TensorFlow无法确认输入张量的固定形状,进而在模型拟合时触发错误。
修改数据集构建代码,补充output_shapes参数,明确输入输出张量的形状:
ds = tf.data.Dataset.from_generator( generator, output_types=({'LoB': tf.int32}, tf.float32), output_shapes=({'LoB': (1,)}, (1,)) )
额外注意点:
- 检查
create_tensor函数中的target变量,确保其为数据集的有效列名字符串,避免因变量未定义或拼写错误导致的隐性问题。 - 若后续需使用输出中显示的其他特征(如
cc、age等),需同步更新模型的输入层结构,以及数据集的output_types和output_shapes,保证输入特征与模型输入完全匹配。
内容的提问来源于stack exchange,提问作者AdamS
相关产品推荐
相关产品推荐

