如何创建包含多组2D张量字典的TFRecords并无解析错误地读取?
如何创建包含多组2D张量字典的TFRecords并无解析错误地读取?
看起来你遇到的问题核心是TFRecord写入和读取时的feature格式不匹配!咱们先拆解问题,再给你修正后的代码。
错误原因分析
- 写入阶段:你把每个4x4的feats张量和2x2的labels张量都通过
tf.io.serialize_tensor转成了单个字节串,然后用tf.train.Feature(bytes_list=tf.train.BytesList(value=[featsByte.numpy()]))存储——这里每个feature本质是单个字节值,不是多维字符串数组。 - 读取阶段:你错误地用了
FixedLenFeature([4,4], tf.string),这相当于告诉TensorFlow“我要解析出一个4x4的字符串数组”,但实际存储的是单个字节串,自然会触发解析错误。
修正方案
只需要调整读取时的feature schema,把FixedLenFeature的shape改成空数组[],表示每个feature是单个字节串,之后再通过tf.io.parse_tensor还原成原来的2D张量即可。另外也可以优化一下写入时的代码,不用绕tf.constant转,直接用numpy处理更简洁。
修正后的完整代码
import numpy as np import tensorflow as tf import random import os # options are "gen", "consume" or "both" task = "both" def createSerialExample(shard): # 直接用numpy处理,更简洁 feats_np = np.array(shard["feats"], dtype=np.uint8) labels_np = np.array(shard["labels"], dtype=np.uint8) # Convert dataset tensor element to serialised bytes featsByte = tf.io.serialize_tensor(feats_np).numpy() labelsByte = tf.io.serialize_tensor(labels_np).numpy() # Build feature from byte lists featsFeature = tf.train.Feature(bytes_list=tf.train.BytesList(value=[featsByte])) labelsFeature = tf.train.Feature(bytes_list=tf.train.BytesList(value=[labelsByte])) # Build the feature map featureMap = { "feats": featsFeature, "labels": labelsFeature } # Build a collection of features defined by the feature map, followed by building an example # from the features and serialising this example example = tf.train.Example(features=tf.train.Features(feature=featureMap)) serialisedExample = example.SerializeToString() return serialisedExample # end createSerialExample #------------------------ Generation -------------------------------------------- if task == "gen" or task == "both": vdata = { "feats":[], "labels":[] } # Create random data in 2 x TFRecords with 2 x shards each of 4x4 and 2x2 data features cnt = 0 for _ in range(2): for _ in range(2): feat4x4 = [[random.randint(1, 10) for _ in range(4)] for _ in range(4)] lbl2x2 = [[random.randint(1, 10) for _ in range(2)] for _ in range(2)] vdata["feats"].append(feat4x4) vdata["labels"].append(lbl2x2) path = os.path.join("datasets//", f"{cnt:06d}.tfrec") dataset = tf.data.Dataset.from_tensor_slices(vdata) # Create a dataset of shards # Write the shards into the TFRecord with tf.io.TFRecordWriter(path) as writer: for shard in dataset: serialisedExample = createSerialExample(shard=shard) writer.write(serialisedExample) # Clear old data vdata = { "feats":[], "labels":[] } cnt = cnt + 1 print("Generation done...") #------------------------ Consumption -------------------------------------------- # Map function for: train_dataset = tf.data.TFRecordDataset(ds_train_files).map(loadDataset) below def loadDataset(ds): # 重点修正:把shape改成[],表示单个字节串 featuresSchema = { "feats": tf.io.FixedLenFeature([], tf.string), "labels": tf.io.FixedLenFeature([], tf.string) } # Extract the dict from the serialised data parsed_ds = tf.io.parse_single_example(ds, featuresSchema) # Get the tensors from the parsed dict X = tf.io.parse_tensor(parsed_ds["feats"], tf.dtypes.uint8) # input training data X.set_shape([4,4]) X = tf.cast(X, tf.float32)/255.0 Y = tf.io.parse_tensor(parsed_ds["labels"], tf.dtypes.uint8) # output target data Y.set_shape([2,2]) Y = tf.cast(Y, tf.float32)/255.0 return X, Y #end loadDataset if task == "consume" or task == "both": ds_train_files = tf.data.Dataset.list_files("datasets\\*", seed=42) # Load the datasets from the list of filenames train_dataset = tf.data.TFRecordDataset(ds_train_files).map(loadDataset).batch(1) # Train the model tf.random.set_seed(42) # Make the results reproducible with consistent random weight matrix model = tf.keras.models.Sequential([ tf.keras.layers.Input(shape=[4,4], dtype=tf.float32), tf.keras.layers.Flatten(), tf.keras.layers.Dense(10, activation="relu"), tf.keras.layers.Dense(10, activation="relu"), tf.keras.layers.Dense(4, activation="sigmoid"), # 2x2 output prediction tf.keras.layers.Reshape((2,2)) ]) model.compile(loss="mse", optimizer="sgd") history = model.fit(train_dataset, epochs=5) print("Consumption done...")
验证说明
修改后,代码就能正常生成TFRecord文件,并且读取后可以顺利训练模型,不会再出现解析错误。核心就是保证写入的feature格式和读取时的schema完全匹配——单个字节串对应FixedLenFeature([], tf.string),之后再还原成原形状的张量。
备注:内容来源于stack exchange,提问作者Keith
相关产品推荐
相关产品推荐

