导出为WebDataset
更新时间:2026-07-29
简介
把训练样本按 WebDataset 规范打成 .tar shards(基于 webdataset.ShardWriter),可选把媒体文件字节一并写入;返回每条样本所在的 shard tar 路径,用于物化训练分片。
功能描述
• 标准 shards:使用 webdataset.ShardWriter,按 max_samples_per_shard / max_shard_size_mb 自动切分
• 媒体可选:include_media=True 时把 media_path 指向的文件下载后写入 tar;False 时仅写 meta.json
• 顺序写入:跨行状态保存在算子实例内,需配合 concurrency=1、batch_size=None 保证顺序
• 键兜底:sample_key 为空时用 md5 兜底生成键
算子参数
输入
| 输入 | 含义 |
|---|---|
| sample_keys | 每条样本的唯一键 str(为空时用 md5 兜底)。 |
| media_paths | 本地或 BOS 媒体路径(会自动 download 到本地再打入 tar)。 |
| meta_jsons | 结构化 caption/label/score 的 JSON 字符串,可为空。 |
输出
| 输出 | 含义 |
|---|---|
| result | string,该样本被写入的 tar shard 路径;写入失败时为空字符串。 |
参数
| 参数名称 | 类型 | 默认值 | 描述 |
|---|---|---|---|
| output_dir | str | (必填) | shard 输出目录(PFS/BOS 挂载路径,不存在会自动创建) |
| shard_prefix | str | 'shard' | shard 文件前缀 默认值:"shard" |
| max_shard_size_mb | float | 512.0 | 单 shard 目标大小上限(MB,近似值)默认值:512.0 |
| max_samples_per_shard | int | 10000 | 单 shard 样本数上限 默认值:10000 |
| media_extension | str/None | None | 媒体扩展名,为空则从 media_path 后缀自动推断 默认值:None |
| include_media | bool | True | 是否把 media 字节写入 tar(False 时仅写 meta.json)默认值:True |
| fail_on_error | bool | False | 出错是否抛异常 默认值:False |
调用示例
Python
1from __future__ import annotations
2
3import json
4import os
5
6import daft
7from daft import col
8
9from daft.aihc.common.udf import aihc_udf
10from daft.aihc.functions.export.export_to_web_dataset import ExportToWebDataset
11
12if __name__ == "__main__":
13 if os.getenv("DAFT_RUNNER", "native") == "ray":
14 import ray
15 ray.init(dashboard_host="0.0.0.0", ignore_reinit_error=True)
16 daft.set_runner_ray()
17 daft.set_execution_config(min_cpu_per_task=0)
18
19 # output_dir 为 mock 路径,请替换为实际 PFS/BOS 挂载目录
20 OUTPUT_DIR = "/mnt/pfs/webdataset_shards" # MOCK
21
22 keys = [f"sample_{i}" for i in range(5)]
23 samples = {
24 "sample_key": keys,
25 "media_path": ["" for _ in keys], # include_media=False 时无需真实媒体
26 "meta_json": [json.dumps({"idx": i, "label": "demo"}) for i in range(5)],
27 }
28 ds = daft.from_pydict(samples)
29 # 顺序写入:concurrency=1, batch_size=None
30 ds = ds.with_column(
31 "shard",
32 aihc_udf(
33 ExportToWebDataset,
34 construct_args={
35 "output_dir": OUTPUT_DIR,
36 "shard_prefix": "shard",
37 "include_media": False,
38 "max_samples_per_shard": 10000,
39 },
40 num_cpus=1,
41 concurrency=1,
42 batch_size=None,
43 )(col("sample_key"), col("media_path"), col("meta_json")),
44 )
45 ds.show()
评价此篇文章
