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

如何使用tf.data.Dataset.save将tf.data.Dataset保存为多个分片?

如何使用tf.data.Dataset.save将tf.data.Dataset保存为多个分片?

我来帮你搞定这个分片保存的问题!你遇到的核心卡点有两个:一是make_csv_dataset返回的数据集每个元素是**(特征字典, 标签)**的结构,直接操作字典肯定会报错;二是shard_func需要兼容TensorFlow的图模式,不能用numpy或普通Python的随机/计算逻辑,得用TensorFlow原生操作。

先给你梳理清楚shard_func的要求:它接收数据集的单个元素(也就是你这里的(features_dict, label))作为输入,必须返回一个标量Tensor,代表当前元素要存入的分片索引。而且这个函数要能在图模式下运行,所以所有计算都得用TensorFlow的API,不能用numpy那套。

先解决你遇到的几个错误场景

  1. 随机分片只生成一个分片的问题
    你之前用np.random.randint生成索引,这是Python/numpy的操作,在TensorFlow图模式下只会执行一次——也就是说所有元素都拿到同一个随机数,自然只生成一个分片。得换成TensorFlow的随机操作,让每个元素都能拿到独立的索引。

  2. 模运算报错的问题
    你直接写x % 10,但这里的x是特征字典(OrderedDict),不是单个张量!得先从字典里取出具体的特征张量,再做计算。

正确的实现示例

1. 确定性分片(每个元素固定分到某个分片)

如果你希望相同特征的元素分到同一个分片,可以用某个特征的哈希值来计算索引,比如用c1这个特征:

import pandas as pd
import numpy as np
import tensorflow as tf

# 生成测试数据(和你原来的代码一致)
n=10000
pd.DataFrame(
    {'label': np.random.randint(low=0, high=2, size=n),
     'f1': np.random.random(n),
     'f2': np.random.random(n),
     'f3': np.random.random(n),
     'c1': np.random.randint(n),
     'c2': np.random.randint(n)}
).to_csv('tmp.csv')

# 加载数据集
data_ts = tf.data.experimental.make_csv_dataset(
        'tmp.csv', 1, label_name='label', num_epochs=1)

# 确定性分片函数:基于c1特征的哈希值分到10个分片
def deterministic_shard_func(features, label):
    # 从特征字典中取出c1张量
    c1_tensor = features['c1']
    # 将张量转为字符串后计算哈希,模10得到分片索引
    shard_index = tf.strings.to_hash_bucket_fast(tf.as_string(c1_tensor), num_buckets=10)
    return shard_index

# 保存到10个分片
data_ts.save('tmp_deterministic.data', shard_func=deterministic_shard_func)

2. 随机分片(每个元素随机分到不同分片)

用TensorFlow的随机API代替numpy,确保每个元素都生成独立的随机索引:

def random_shard_func(features, label):
    # 生成0-9之间的随机整数,dtype要和分片数匹配
    shard_index = tf.random.uniform(shape=[], minval=0, maxval=10, dtype=tf.int64)
    return shard_index

data_ts.save('tmp_random.data', shard_func=random_shard_func)

额外注意点

如果你的数据集设置了大于1的batch size,shard_func的输入会是批量的特征字典和批量标签。这时候你可以选择:

  • 让整个batch分到同一个分片:返回一个标量索引即可
  • 让每个样本分到不同分片:返回和batch大小一致的索引张量

这样操作后,你就能看到生成多个分片文件啦!

备注:内容来源于stack exchange,提问作者dule arnaux

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 15:22:58