如何将多层嵌套JSON文件转换为TensorFlow Tensor或TensorDataset
处理多层嵌套JSON转TensorFlow Tensor/TensorDataset的方法
针对多层嵌套的JSON结构,不用先转成Pandas DataFrame(反而会因为嵌套结构卡壳),可以用下面两种实用方法直接处理:
方法一:扁平化嵌套结构后转Dataset
如果你的模型不需要保留嵌套结构,先把嵌套JSON转成扁平的键值对,再生成Dataset:
- 写一个递归函数扁平化嵌套字典:
import json def flatten_dict(nested_dict, parent_key='', sep='_'): items = [] for k, v in nested_dict.items(): new_key = f"{parent_key}{sep}{k}" if parent_key else k if isinstance(v, dict): items.extend(flatten_dict(v, new_key, sep=sep).items()) elif isinstance(v, list): # 如果是列表,可根据需求转成逗号分隔字符串或保留为数组 items.append((new_key, json.dumps(v))) else: items.append((new_key, v)) return dict(items)
- 加载JSON并处理:
# 加载JSON文件(假设每个样本是列表里的元素) with open('your_data.json', 'r') as f: raw_data = json.load(f) # 扁平化每个样本 flattened_data = [flatten_dict(sample) for sample in raw_data] # 转成DataFrame(可选,方便查看) import pandas as pd df = pd.DataFrame(flattened_data) # 转成TensorFlow Dataset import tensorflow as tf dataset = tf.data.Dataset.from_tensor_slices(dict(df)) # 后续可以做shuffle、batch等操作 dataset = dataset.shuffle(1000).batch(32)
方法二:直接构建嵌套TensorDataset
如果你的模型需要保留嵌套结构(比如多输入模型,不同嵌套分支对应不同输入),TensorFlow支持直接处理嵌套数据:
- 加载JSON后直接生成嵌套Dataset:
import json import tensorflow as tf # 加载JSON(假设数据是列表格式,每个元素是嵌套字典) with open('your_data.json', 'r') as f: raw_data = json.load(f) # 直接转成嵌套Dataset dataset = tf.data.Dataset.from_tensor_slices(raw_data) # 查看嵌套结构 for sample in dataset.take(1): print(sample) # 输出类似:{'user': {'id': <tf.Tensor: ...>, 'name': <tf.Tensor: ...>}, 'features': {'f1': <tf.Tensor: ...>}}
- 适配模型训练:
如果你的模型接受嵌套输入,可以直接把Dataset喂进去。比如定义一个多输入模型:
input_user_id = tf.keras.Input(shape=(), name='user_id') input_user_name = tf.keras.Input(shape=(None,), dtype=tf.string, name='user_name') input_f1 = tf.keras.Input(shape=(), name='features_f1') # 后续构建模型结构...
处理大JSON文件(无法全量加载)
如果JSON文件太大,不能一次性加载到内存,可以逐行读取处理:
import tensorflow as tf import json def parse_json_line(line): # 把字符串转成Python字典 sample = json.loads(line.numpy().decode('utf-8')) # 这里可以按需扁平化或直接返回嵌套结构 return sample # 读取每行JSON dataset = tf.data.TextLineDataset('large_data.json') # 用tf.py_function处理每行 dataset = dataset.map(lambda x: tf.py_function(parse_json_line, [x], tf.types.experimental.TensorStructure(raw_data[0]))) # 后续操作 dataset = dataset.shuffle(1000).batch(32)
注意事项
- 确保所有样本的嵌套结构一致,否则TensorFlow无法生成规整的张量;
- 如果嵌套结构里有可变长度的序列(比如不同样本的列表长度不同),需要用Ragged Tensor,可以在map函数中把列表转成tf.RaggedTensor;
- 字符串类型的字段会被自动转成tf.string张量,数值类型自动转成对应的数值张量。
内容的提问来源于stack exchange,提问作者Jens Voorpyl
相关产品推荐
相关产品推荐

