使用TensorFlow Dataset API读取多CSV时遇list无get_shape属性错误
嘿,我来帮你搞定这个'list' object has no attribute 'get_shape'的问题~
从你给出的代码片段来看,之前的decode_csv逻辑是没问题的,这个错误大概率是你在构建Dataset的后续操作中,不小心把普通Python列表当成TensorFlow的Tensor对象来处理了。TensorFlow的Dataset API只认Tensor或Dataset对象,一旦你传入了原生列表,就会触发这个属性错误。
我给你梳理几个最常见的坑和对应的解决方案:
1. 检查interleave/flat_map的返回值
如果你是用多文件读取,肯定会用到interleave或者flat_map来遍历每个CSV文件。这里一定要确保你传入的lambda函数返回的是**TextLineDataset(或其他Dataset对象)**,而不是把Dataset转换成了Python列表。
❌ 错误示例(会触发问题):
# 错误:把Dataset转成了列表 dataset5 = tf.data.Dataset.from_tensor_slices(filenames) dataset5 = dataset5.interleave(lambda fn: list(tf.data.TextLineDataset(fn).map(decode_csv)))
✅ 正确写法:
record_defaults = [[""], [0.0], [0.0], [0.0], [0.0], [0.0], [0.0]] def decode_csv(line): col1, col2, col3, col4, col5, col6, col7 = tf.decode_csv(line, record_defaults) features = tf.stack([col2, col3, col4, col5, col6]) labels = tf.stack([col7]) return features, labels filenames = tf.placeholder(tf.string, shape=[None]) # 从文件名构建基础Dataset dataset5 = tf.data.Dataset.from_tensor_slices(filenames) # 用interleave读取每个文件,返回的是Dataset对象 dataset5 = dataset5.interleave( lambda filename: tf.data.TextLineDataset(filename).skip(1).map(decode_csv), # skip跳过表头 cycle_length=4 # 并行处理的文件数 ) # 后续的洗牌、批量操作 dataset5 = dataset5.shuffle(buffer_size=1000).batch(32)
2. 确保map操作返回的是Tensor
如果你的map函数里不小心返回了Python列表(比如手动把Tensor转成了列表),也会触发这个错误。比如不要在decode_csv里做类似return list(features), list(labels)的操作——你现在的tf.stack已经正确返回Tensor了,保持这个逻辑就好。
3. 喂数据时的小细节
当你给filenames placeholder喂数据时,尽量传入numpy数组而不是普通Python列表(虽然TensorFlow能兼容,但避免潜在的类型混淆):
# 比如你有多个CSV路径 csv_files = ["data1.csv", "data2.csv", "data3.csv"] # 转成numpy数组再喂入 sess.run(..., feed_dict={filenames: np.array(csv_files)})
按照上面的步骤排查一遍,应该就能解决这个问题啦~
内容的提问来源于stack exchange,提问作者edn
相关产品推荐
相关产品推荐

