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

基于TensorFlow Datasets处理Droid数据集的两大技术问题求助

解决TensorFlow Datasets处理Droid数据集的两大问题

问题1:流式过滤时RAM持续占用直至崩溃

原因分析

你当前用Python生成器遍历tf.data.Dataset并转成numpy数组处理,这种方式会触发数据集的全量预加载——即使你只取单个样本,TensorFlow为了优化性能会提前缓存大量数据到内存,最终导致RAM溢出。测试集数据量小所以没问题,但2TB的完整数据集会直接撑爆内存。

解决方案

改用tf.data原生的图模式操作(filter+map),所有处理逻辑用TensorFlow算子实现,避免转成numpy数组,这样数据集会严格流式加载,只处理当前样本:

import tensorflow as tf
import tensorflow_datasets as tfds

# 加载数据集(不提前缓存,保持流式)
ds = tfds.load("droid", data_dir="gs://gresearch/robotics", split="train", shuffle_files=False)

# 定义过滤条件的TensorFlow函数
def filter_fn(sample):
    # 提取所有步骤的笛卡尔坐标
    ee_positions = sample['steps']['observation']['cartesian_position']
    # 取前3个维度(x,y,z)
    x = ee_positions[:, 0]
    y = ee_positions[:, 1]
    z = ee_positions[:, 2]
    
    # 检查所有步骤是否都在范围内
    x_valid = tf.logical_and(x >= -0.0, x <= 1.0)
    y_valid = tf.logical_and(y >= -1.0, y <= 1.0)
    z_valid = tf.logical_and(z >= 0.0, z <= 1.0)
    all_valid = tf.reduce_all(tf.logical_and(tf.logical_and(x_valid, y_valid), z_valid))
    
    return all_valid

# 定义转换函数,提取需要的特征
def map_fn(sample):
    # 提取初始关节位置(取第一个step)
    q_init = sample['steps']['observation']['joint_position'][0]
    # 计算时间戳(用TensorFlow实现,避免numpy)
    num_steps = tf.shape(sample['steps']['observation']['cartesian_position'])[0]
    timestamp = tf.linspace(0.0, tf.cast(num_steps-1, tf.float32)*0.1, num_steps)
    
    return sample['steps']['observation']['cartesian_position'], q_init, timestamp

# 链式操作:过滤→转换→批量/按需处理
filtered_ds = ds.filter(filter_fn).map(map_fn)

# 遍历流式处理
for ee_positions, q_init, timestamp in filtered_ds:
    # 这里处理单个样本,不会全量加载
    print(ee_positions.shape, q_init.shape, timestamp.shape)
    # 按需break或继续处理

问题2:PartialDecoding无法成功加载部分特征

原因分析

你手动指定了tfds.core.features.dataset_feature.Dataset类型,这是错误的——PartialDecoding只需要和原数据集特征结构键匹配的嵌套字典,不需要显式指定Dataset类型,TensorFlow会自动识别原数据集的嵌套结构。

解决方案

直接按照原特征结构的层级,只保留需要的键即可,不需要额外指定特征类型:

import tensorflow_datasets as tfds

# 正确的PartialDecoding定义:只保留需要的特征路径
decoder = tfds.decode.PartialDecoding({
    'steps': {
        'observation': {
            'cartesian_position': tfds.features.Tensor(),  # 可以省略shape和dtype,会自动匹配原数据集
            'joint_position': tfds.features.Tensor(),
        },
    },
})

# 加载数据集时指定解码器
ds = tfds.load(
    "droid",
    data_dir="gs://gresearch/robotics",
    split="train",
    decoders=decoder,
    shuffle_files=False
)

# 验证结果:每个样本只包含指定的特征
for sample in ds.take(1):
    print(list(sample.keys()))  # 输出: ['steps']
    print(list(sample['steps'].keys()))  # 输出: ['observation']
    print(list(sample['steps']['observation'].keys()))  # 输出: ['cartesian_position', 'joint_position']

备选方案:修改特征配置文件

如果PartialDecoding还是有问题,可以直接修改本地的features.json,然后通过tfds.load的data_dir指向本地修改后的配置,同时设置download=False避免覆盖:

  1. 从gs://gresearch/robotics下载Droid数据集的features.json到本地目录
  2. 删除不需要的特征(比如所有Image字段、action相关字段等)
  3. 加载时指定本地目录:
ds = tfds.load(
    "droid",
    data_dir="./local_droid_config",  # 指向包含修改后features.json的目录
    split="train",
    download=False,  # 禁止下载覆盖本地配置
    shuffle_files=False
)

内容的提问来源于stack exchange,提问作者user31865617

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 01:20:15