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

如何优化JSON存储numpy数组并通过TensorFlow高效加载?

适配TensorFlow的多维数组保存与高效加载方案

当前加载代码的问题

  1. json.loads(file_path)错误:在tf.data.Dataset.map中,file_path是Tensor类型,而json.loads仅接受字符串/字节对象,类型不匹配导致报错。
  2. tf.io.decode_json_example使用错误:该API是用于解析TensorFlow Example protobuf序列化后的JSON格式,并非普通嵌套列表JSON,完全不适用你的场景。

方案一:使用Numpy原生格式(.npy/.npz)—— 简单高效

Numpy的.npy格式专门用于存储多维数组,体积小、读写快,TensorFlow可直接兼容。

保存代码

import numpy as np
# 保存单个三维数组
np.save("dataset.npy", dataset)

# 若需同时保存数据和标签,用.npz格式
np.savez("dataset.npz", data=dataset, labels=your_labels_array)

TF加载代码

import tensorflow as tf
from io import BytesIO

def load_npy_file(file_path):
    # 读取文件字节并解码为numpy数组
    def _load_npy(bytes_data):
        return np.load(BytesIO(bytes_data.numpy()))
    
    data_bytes = tf.io.read_file(file_path)
    data = tf.numpy_function(_load_npy, [data_bytes], tf.float32)
    # 手动指定数组形状(根据你的实际数据调整,示例为(2,2,N))
    data.set_shape((2, 2, None))
    
    # 替换为你的标签获取逻辑
    label = tf.py_function(lambda path: get_label(path.numpy().decode()), [file_path], tf.int32)
    return data, label

# 构建数据集并并行加载
train_ds = tf.data.Dataset.list_files("/path/to/*.npy")
train_ds = train_ds.map(load_npy_file, num_parallel_calls=tf.data.AUTOTUNE)

方案二:使用TFRecord格式—— TensorFlow原生大规模数据方案

TFRecord是TensorFlow官方推荐的大规模存储格式,支持高效并行读取和序列化,适合数据量较大的场景。

保存为TFRecord

import tensorflow as tf

def array_to_tfexample(arr, label):
    # 将数组转为字节,记录形状
    arr_bytes = arr.tobytes()
    shape = arr.shape
    example = tf.train.Example(features=tf.train.Features(feature={
        "data": tf.train.Feature(bytes_list=tf.train.BytesList(value=[arr_bytes])),
        "shape": tf.train.Feature(int64_list=tf.train.Int64List(value=shape)),
        "label": tf.train.Feature(int64_list=tf.train.Int64List(value=[label]))
    }))
    return example.SerializeToString()

# 遍历数据写入TFRecord
with tf.io.TFRecordWriter("dataset.tfrecord") as writer:
    # 假设你有数据集列表和对应标签列表
    for arr, label in zip(dataset_list, label_list):
        writer.write(array_to_tfexample(arr, label))

TF加载代码

def parse_tfexample(example_proto):
    # 定义Feature解析规则
    feature_desc = {
        "data": tf.io.FixedLenFeature([], tf.string),
        "shape": tf.io.FixedLenFeature([3], tf.int64),  # 三维数组对应3个形状参数
        "label": tf.io.FixedLenFeature([], tf.int64)
    }
    features = tf.io.parse_single_example(example_proto, feature_desc)
    # 恢复数组形状
    data = tf.io.decode_raw(features["data"], tf.float32)
    data = tf.reshape(data, features["shape"])
    return data, features["label"]

# 构建并行加载的数据集
train_ds = tf.data.TFRecordDataset("dataset.tfrecord")
train_ds = train_ds.map(parse_tfexample, num_parallel_calls=tf.data.AUTOTUNE)

方案三:修正JSON加载逻辑(仅推荐小数据场景)

如果坚持使用JSON,需用tf.py_function包装Python的JSON读取逻辑,避开Tensor类型限制,但JSON对大数组的读写效率远低于前两种方案。

修正后的加载代码

import json
import tensorflow as tf
import numpy as np

def load_json_file(file_path):
    def _load_json(path):
        path_str = path.numpy().decode("utf-8")
        with open(path_str, "r") as f:
            data_list = json.load(f)
        return np.array(data_list, dtype=np.float32)
    
    data = tf.py_function(_load_json, [file_path], tf.float32)
    data.set_shape((2, 2, None))  # 指定数组形状
    label = tf.py_function(lambda path: get_label(path.numpy().decode()), [file_path], tf.int32)
    return data, label

train_ds = tf.data.Dataset.list_files("/path/to/*.json")
train_ds = train_ds.map(load_json_file, num_parallel_calls=tf.data.AUTOTUNE)

总结

  • 数据量较大时优先选择TFRecord或**.npy格式**,二者都比JSON更适配TensorFlow,读写效率更高。
  • JSON仅适合存储小规模的嵌套结构数据,不适合大规模多维数组存储。

内容的提问来源于stack exchange,提问作者Life is full of Learning

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 07:03:30