read my own tfreocrds encounter questions

Viewed 23

question description:

I created my own tfRecords and want to read it.When i read the tfrecords,questions happen.I only need img and label,like the famous dataset minst.It looks like error happen in iterator.I have searched similar question on websites,but i can not acquire the answer.

code:here is create code

import tensorflow as tf
import numpy as np
from PIL import Image
import os,glob
#定义基本变量
train_path_1="D:/testimages/train/AMD"
train_path_2="D:/testimages/train/normal"

test_path_1="D:/testimages/test/AMD"
test_path_2="D:/testimages/test/normal"
def create_tfrecord(path1,path2,tfname):
    writer=tf.compat.v1.python_io.TFRecordWriter(tfname)
    for file in os.listdir(path1): 
        img_path=path1+"/"+file #每一个图片的地址
        img=Image.open(img_path)
        print(img)
        img= img.resize((256,256))
        print(np.shape(img))
        img_raw=img.tobytes()#将图片转化为二进制格式
        example = tf.train.Example(features=tf.train.Features(feature={
            "label": tf.train.Feature(int64_list=tf.train.Int64List(value=[1])),
            'img_raw': tf.train.Feature(bytes_list=tf.train.BytesList(value=[img_raw]))
        })) #example对象对label和image数据进行封装
        writer.write(example.SerializeToString())  #序列化为字符串
    for file in os.listdir(path2): 
        img_path=path2+"/"+file #每一个图片的地址
        img=Image.open(img_path)
        print(img)
        img= img.resize((256,256))
        print(np.shape(img))
        img_raw=img.tobytes()#将图片转化为二进制格式
        example = tf.train.Example(features=tf.train.Features(feature={
            "label": tf.train.Feature(int64_list=tf.train.Int64List(value=[0])),
            'img_raw': tf.train.Feature(bytes_list=tf.train.BytesList(value=[img_raw]))
        })) #example对象对label和image数据进行封装
        writer.write(example.SerializeToString())  #序列化为字符串
    writer.close()
    print("create successfully!")
create_tfrecord(train_path_1,train_path_2,"train.tfRecord")
create_tfrecord(test_path_1,test_path_2,"test.tfRecord")

here is read code:

def get_data(filename):
    dataset = tf.data.TFRecordDataset(filename)
    dataset=dataset.map(read_and_decode)
#    print(dataset)
#     dataset = dataset.shuffle(buffer_size=100)  # 在缓冲区中随机打乱数据
#     dataset = dataset.batch(batch_size=4)  # 每10条数据为一个batch,生成一个新的Datasets
    #print(dataset)
    #dataset = tf.data.Dataset.from_tensor_slices(dataset)
    #print(type(dataset))
    dataset = dataset.shuffle(20).batch(20)
    #print(type(dataset))
    #print(tf.compat.v1.data.get_output_shapes(dataset))
    output_shapes=tf.compat.v1.data.get_output_shapes(dataset)
    output_types=tf.compat.v1.data.get_output_types(dataset)
    #print(tf.compat.v1.data.get_output_types(dataset))
    iterator = tf.data.Iterator.from_structure(output_types,
                                            output_shapes)
    sess.run(iterator.make_initializer(dataset))#初始化迭代器
    batch_image, batch_label = iterator.get_next()
    return batch_image, batch_label
def read_and_decode(example_string):
    
    features=tf.io.parse_single_example(example_string,features={'label':tf.io.FixedLenFeature([],tf.int64),
                                                                'img_raw':tf.io.FixedLenFeature([],tf.string)
                                                                })
    img=tf.io.decode_raw(features['img_raw'],tf.uint8)
    img = tf.reshape(img, [256, 256, 3])
    img=tf.cast(img,tf.float32)*(1./255)#要不要-0.5
    label=tf.cast(features['label'],tf.float32)
    
    return img,label

img_train,label_train = get_data("train.tfRecord")
#print(img_train[0],"sdf")
#print(label_train[0],"fasdf")
img_test,label_test = get_data("test.tfRecord")

error

AttributeError                            Traceback (most recent call last)
~\AppData\Local\Temp/ipykernel_22620/2092499530.py in <module>
     11     return img,label
     12 
---> 13 img_train,label_train = get_data("train.tfRecord")
     14 #print(img_train[0],"sdf")
     15 #print(label_train[0],"fasdf")

~\AppData\Local\Temp/ipykernel_22620/1814192877.py in get_data(filename)
     14     output_types=tf.compat.v1.data.get_output_types(dataset)
     15     #print(tf.compat.v1.data.get_output_types(dataset))
---> 16     iterator = tf.data.Iterator.from_structure(output_types,
     17                                             output_shapes)
     18     sess.run(iterator.make_initializer(dataset))#初始化迭代器

AttributeError: type object 'IteratorBase' has no attribute 'from_structure'

my error understanding:

it looks like error happen in iterator.I guess i fail to read my record.But i don't know how to solve it.I hope someone can help me.Thanks for your help!

0 Answers
Related