You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何创建包含多组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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.14 14:54:29