Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 46 additions & 0 deletions docs/reading-and-writing.md
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@ Built-in readers:
| `read_json(...)` | JSON files | one row per JSON file by default |
| `read_jsonl(...)` | JSON Lines files | one row per line |
| `read_parquet(...)` | Parquet datasets or files | row views backed by Arrow columns |
| `read_webdataset(...)` | WebDataset tar archives | one row per sample, with member extensions as fields |
| `read_hf_dataset(...)` | Hugging Face datasets | rows from generated Parquet shards, with optional file path resolution |
| `read_lerobot(...)` | LeRobot robotics datasets | one row per episode, including frame/video metadata |
| `text.read_commoncrawl(...)` | Common Crawl WARC or WET files | one row per WARC/WET record with selected WARC/HTTP fields |
Expand Down Expand Up @@ -228,6 +229,50 @@ datasets or attributes default to raising an error. Set
If a selected column can be missing from every group in an input file, pass
`dtypes` for that column so the reader can emit a stable Arrow type.

## WebDataset

`read_webdataset(...)` reads `.tar`, `.tar.gz`, and `.tgz` archives using the
same fsspec-backed input handling as the other file readers. Inputs can be
archive paths, globs, directories, `DataFile` values, or mixed lists of those.

```python
import refiner as mdr

pipeline = mdr.read_webdataset(
"s3://my-bucket/shards/*.tar",
dtypes={"jpg": mdr.datatype.image_bytes()},
)
```

The reader streams each tar archive sequentially and does not load the full
archive into memory. Archives are planned as atomic files, so `num_shards`
cannot split one large archive across multiple workers. Members for a sample
must be contiguous in the archive, which is the standard WebDataset shard
layout.

Members are grouped by the path before the first dot. The suffix after that
first dot becomes the output field name:

```text
0001.jpg
0001.json
0002.jpg
0002.txt
```

emits rows like:

- `sample_key="0001"`, `jpg=<bytes>`, `json=<dict>`
- `sample_key="0002"`, `jpg=<bytes>`, `txt=<bytes>`

Dots in the basename start the field suffix, so sample keys should not contain dots.

The archive path is added as `file_path` by default. Set
`file_path_column=None` to omit it, or `sample_key_column=...` to rename the
sample key column. JSON members are parsed to Python values by default; pass
`parse_json=False` to keep `.json` members as raw bytes. All non-JSON payloads
are emitted as bytes. Members without a dot are skipped.

## Common Crawl text readers

[Common Crawl](https://commoncrawl.org/) publishes large public web crawls.
Expand Down Expand Up @@ -324,6 +369,7 @@ Reader behavior differs by format:
- CSV and line-delimited JSON usually shard by file and byte range
- HDF5 shards by file only; `num_shards` cannot exceed the input file count
- Parquet shards by file and planned row-group or row ranges
- WebDataset shards by archive file only; samples are streamed from each archive
- LeRobot shards by episode parquet metadata
- `from_items(...)` shards synthetic in-memory rows into planned chunks

Expand Down
2 changes: 2 additions & 0 deletions src/refiner/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
read_videos,
SUPPORTED_CUDA_VERSIONS,
SUPPORTED_GPU_TYPES,
read_webdataset,
task,
)
from refiner.pipeline.expressions import coalesce, col, if_else, lit
Expand Down Expand Up @@ -54,6 +55,7 @@
"read_lerobot",
"read_parquet",
"read_videos",
"read_webdataset",
"from_items",
"from_source",
"task",
Expand Down
2 changes: 2 additions & 0 deletions src/refiner/pipeline/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
read_lerobot,
read_parquet,
read_videos,
read_webdataset,
task,
)
from refiner.pipeline.resources import (
Expand Down Expand Up @@ -41,6 +42,7 @@
"read_lerobot",
"read_parquet",
"read_videos",
"read_webdataset",
"from_items",
"from_source",
"task",
Expand Down
37 changes: 37 additions & 0 deletions src/refiner/pipeline/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
Hdf5Reader,
JsonReader,
ParquetReader,
WebDatasetReader,
)
from refiner.pipeline.sources.readers.lerobot import LeRobotEpisodeReader
from refiner.pipeline.sources.readers.hdf5 import MissingPolicy, PathSelection
Expand Down Expand Up @@ -809,6 +810,42 @@ def read_hdf5(
)


def read_webdataset(
inputs: DataFileSetLike,
*,
fs: AbstractFileSystem | None = None,
storage_options: Mapping[str, Any] | None = None,
recursive: bool = False,
target_shard_bytes: int = DEFAULT_TARGET_SHARD_BYTES,
num_shards: int | None = None,
file_path_column: str | None = "file_path",
sample_key_column: str | None = "sample_key",
parse_json: bool = True,
dtypes: DTypeMapping | None = None,
) -> RefinerPipeline:
"""Create a pipeline with a WebDataset tar reader source.

WebDataset archives are planned as atomic files and read sequentially. Each
sample emits one row; the suffix after the first dot in each member basename
becomes the field name. JSON members are parsed to Python values by default,
while other member payloads are bytes.
"""
return RefinerPipeline(
source=WebDatasetReader(
inputs,
fs=fs,
storage_options=storage_options,
recursive=recursive,
target_shard_bytes=target_shard_bytes,
num_shards=num_shards,
file_path_column=file_path_column,
sample_key_column=sample_key_column,
parse_json=parse_json,
dtypes=dtypes,
)
)


def read_parquet(
inputs: DataFileSetLike,
*,
Expand Down
2 changes: 2 additions & 0 deletions src/refiner/pipeline/sources/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
JsonReader,
LeRobotEpisodeReader,
ParquetReader,
WebDatasetReader,
)

__all__ = [
Expand All @@ -20,4 +21,5 @@
"JsonReader",
"LeRobotEpisodeReader",
"ParquetReader",
"WebDatasetReader",
]
2 changes: 2 additions & 0 deletions src/refiner/pipeline/sources/readers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from refiner.pipeline.sources.readers.json import JsonReader
from refiner.pipeline.sources.readers.lerobot import LeRobotEpisodeReader
from refiner.pipeline.sources.readers.parquet import ParquetReader
from refiner.pipeline.sources.readers.webdataset import WebDatasetReader
from refiner.robotics.lerobot_format import LeRobotRow

__all__ = [
Expand All @@ -18,4 +19,5 @@
"LeRobotEpisodeReader",
"LeRobotRow",
"ParquetReader",
"WebDatasetReader",
]
156 changes: 156 additions & 0 deletions src/refiner/pipeline/sources/readers/webdataset.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,156 @@
from __future__ import annotations

from collections.abc import Iterator, Mapping
import posixpath
import tarfile
from typing import Any

from fsspec import AbstractFileSystem
import orjson

from refiner.io import DataFile
from refiner.io.fileset import DataFileSetLike
from refiner.pipeline.data.datatype import DTypeMapping, dtype_to_plan
from refiner.pipeline.data.row import DictRow
from refiner.pipeline.data.shard import FilePartsDescriptor
from refiner.pipeline.sources.readers.base import BaseReader, Shard, SourceUnit
from refiner.pipeline.sources.readers.utils import DEFAULT_TARGET_SHARD_BYTES


class WebDatasetReader(BaseReader):
"""WebDataset tar reader planned at archive granularity.

Each output row is one WebDataset sample. Members are grouped by the path
before the first dot in the basename, and the remaining suffix becomes the
output field name. JSON members are parsed to Python values by default; all
other member payloads are emitted as bytes.
"""

name = "read_webdataset"

def __init__(
self,
inputs: DataFileSetLike,
*,
fs: AbstractFileSystem | None = None,
storage_options: Mapping[str, Any] | None = None,
recursive: bool = False,
target_shard_bytes: int = DEFAULT_TARGET_SHARD_BYTES,
num_shards: int | None = None,
file_path_column: str | None = "file_path",
sample_key_column: str | None = "sample_key",
parse_json: bool = True,
dtypes: DTypeMapping | None = None,
):
super().__init__(
inputs,
fs=fs,
storage_options=storage_options,
recursive=recursive,
extensions=(".tar", ".tar.gz", ".tgz"),
target_shard_bytes=target_shard_bytes,
num_shards=num_shards,
file_path_column=file_path_column,
split_by_bytes=False,
dtypes=dtypes,
)
self.sample_key_column = sample_key_column
self.parse_json = parse_json
self._metadata_columns = frozenset(
name for name in (file_path_column, sample_key_column) if name is not None
)
if (
self.file_path_column is not None
and self.sample_key_column is not None
and self.file_path_column == self.sample_key_column
):
raise ValueError("file_path_column and sample_key_column must be distinct")

def describe(self) -> dict[str, Any]:
description = super().describe()
description.update(
{
"sample_key_column": self.sample_key_column,
"parse_json": self.parse_json,
"dtypes": (
{key: dtype_to_plan(dtype) for key, dtype in self.dtypes.items()}
if self.dtypes
else None
),
}
)
return description

def read_shard(self, shard: Shard) -> Iterator[SourceUnit]:
descriptor = shard.descriptor
assert isinstance(descriptor, FilePartsDescriptor)
for part in descriptor.parts:
source = self.fileset.resolve_file(part.source_index, part.path)
yield from self._read_archive(source)

def _read_archive(self, source: DataFile) -> Iterator[SourceUnit]:
current_key: str | None = None
current_row: dict[str, Any] = {}

def flush() -> Iterator[SourceUnit]:
if current_key is None:
return
row = dict(current_row)
if self.sample_key_column is not None:
row[self.sample_key_column] = current_key
yield DictRow(self._with_file_path(row, source))

with (
source.open(mode="rb") as raw,
tarfile.open(fileobj=raw, mode="r|*") as tar,
):
for member in tar:
if not member.isfile():
continue
member_path = posixpath.normpath(member.name).lstrip("/")
if not member_path or member_path == ".":
continue
directory, basename = posixpath.split(member_path)
sample_prefix, separator, field_name = basename.partition(".")
if not separator or not sample_prefix or not field_name:
continue
sample_key = (
f"{directory}/{sample_prefix}" if directory else sample_prefix
)
field_name = field_name.lower()
if current_key is not None and sample_key != current_key:
yield from flush()
current_row = {}
current_key = sample_key
if field_name in self._metadata_columns:
raise ValueError(
f"WebDataset member field {field_name!r} collides with a "
"metadata column; rename the metadata column or disable it"
)
if field_name in current_row:
raise ValueError(
f"Duplicate WebDataset field {field_name!r} for sample "
f"{sample_key!r} in {source.abs_path()!r}"
)
member_file = tar.extractfile(member)
if member_file is None:
current_row[field_name] = b""
continue
with member_file:
payload = member_file.read()
if self.parse_json and (
field_name == "json" or field_name.endswith(".json")
):
try:
current_row[field_name] = orjson.loads(payload)
except orjson.JSONDecodeError as exc:
raise ValueError(
f"Invalid JSON member {member_path!r} in {source.abs_path()!r}"
) from exc
continue
current_row[field_name] = payload
Comment thread
guipenedo marked this conversation as resolved.

yield from flush()


__all__ = ["WebDatasetReader"]
Loading
Loading