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

在Ray Tune Tuner中集成大规模机器学习数据集的正确方法

Ray Tune集成PyTorch数据集的最佳实践与核心规则

一、集成数据集的首选方法

1. 大数据集:Worker独立加载(首推)

直接把数据集的下载、初始化逻辑写到trainable函数里,每个调优Worker自己处理数据加载。这种方式彻底避免跨进程传递大对象的问题,还能利用分布式环境的并行下载能力,效率更高。示例代码:

def trainable(config):
    # 每个Worker自行加载数据集
    dataset = MyDataset(root="./data")  # 本地无数据自动下载
    dataloader = DataLoader(dataset, batch_size=config["batch_size"])
    
    # 后续训练流程
    model = MyModel()
    # ...训练循环代码

2. 小数据集:借助Ray对象存储

如果数据集不大,先在主进程用ray.put()把数据集句柄存入Ray对象存储,再在trainable函数里用ray.get()获取引用。这种方式避免了重复序列化大对象,减少传输开销:

# 主进程中加载并存储数据集
dataset = MyDataset(root="./data")
dataset_ref = ray.put(dataset)

def trainable(config):
    # 从对象存储拉取数据集
    dataset = ray.get(dataset_ref)
    dataloader = DataLoader(dataset, batch_size=config["batch_size"])
    # ...训练逻辑

3. 中等数据集:用Ray Data做分布式预处理

如果需要对数据集做分布式预处理,推荐用Ray Data封装数据集,它会自动处理分布式存储与加载,和Ray Tune无缝兼容:

# 主进程创建Ray Data数据集
ray_dataset = ray.data.from_torch(MyDataset(root="./data"))

def trainable(config):
    # 转换为PyTorch DataLoader
    dataloader = ray_dataset.iter_torch_batches(batch_size=config["batch_size"])
    # ...训练逻辑

二、Ray跨模块传递数据的核心规则

  • 别直接传大对象:Ray默认会序列化函数的所有闭包变量(比如你用partial传入的数据集),大对象序列化后容易超过阈值触发报错,必须用ray.put()存入对象存储,只传递引用而非原始对象。
  • 优先让Worker自主加载:分布式场景下,Worker自己加载数据能减轻主进程的带宽压力,避免序列化开销,还能利用多节点的存储资源。
  • 对象存储全局共享:ray.put()的对象会存在Ray的分布式对象存储里,所有Worker都能通过引用快速获取,不用重复传输。
  • 注意序列化兼容性:Ray默认用Pickle序列化对象,要是遇到不支持Pickle的对象(比如某些自定义PyTorch数据集),要么换支持的格式,要么让Worker本地加载。

三、绕过Python Pickle序列化的方案

1. Worker独立加载数据(最直接)

就像前面说的,把数据集的加载逻辑放到trainable函数内部,每个Worker自己创建数据集句柄,完全不涉及对象序列化和跨进程传递。这种方式不仅绕开Pickle,还从根源上解决了大对象传输的问题。

2. 用Ray Data替代原生数据集

Ray Data内部用Apache Arrow这类高效序列化格式代替Pickle,处理大规模数据时性能更好,而且不用手动管理对象存储,适配大部分PyTorch数据集场景。

3. 注册自定义序列化器(进阶)

如果必须传递特定对象,可以通过ray.register_serializer()注册自定义的序列化/反序列化函数,替换默认的Pickle。但这种方式需要针对特定对象实现,复杂度较高,一般不推荐作为首选方案。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 05:40:29