如何优化JSON存储numpy数组并通过TensorFlow高效加载?
适配TensorFlow的多维数组保存与高效加载方案
当前加载代码的问题
json.loads(file_path)错误:在tf.data.Dataset.map中,file_path是Tensor类型,而json.loads仅接受字符串/字节对象,类型不匹配导致报错。tf.io.decode_json_example使用错误:该API是用于解析TensorFlowExampleprotobuf序列化后的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
相关产品推荐
相关产品推荐

