Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
25 commits
Select commit Hold shift + click to select a range
34d2f34
feat(interleaved): foundation IO improvements — schema utilities, bas…
VibhuJawa Mar 24, 2026
7702bfd
simplify(interleaved): move resolve_schema to shared utils, remove de…
VibhuJawa Mar 24, 2026
a1eeccc
Merge remote-tracking branch 'origin/main' into feat/interleaved-io-f…
VibhuJawa Mar 24, 2026
33088b1
chore: remove unrelated nixl-cu12 from override-dependencies
VibhuJawa Mar 24, 2026
6f0e655
chore: restore nixl-cu12 override-dependency
VibhuJawa Mar 24, 2026
57263df
Remove source_id_field from WebDataset reader
VibhuJawa Mar 24, 2026
27cc69a
simplify: fix align_table, type hints, and add dict-type comment
VibhuJawa Mar 24, 2026
987f5d4
fix(tutorial): update quickstart notebook and set source_files in rea…
VibhuJawa Mar 24, 2026
575e505
Apply suggestion from @claude[bot]
VibhuJawa Mar 24, 2026
7df12f3
ci: update secrets baseline for notebook image outputs (false positives)
VibhuJawa Mar 24, 2026
1a4dce4
fix: make source_files unique per split in webdataset reader
VibhuJawa Mar 24, 2026
472b21c
fix: source_files per split lists only contributing tars in webdatase…
VibhuJawa Mar 24, 2026
be9e7c5
fix: source_files per split lists only contributing tars in webdatase…
VibhuJawa Mar 24, 2026
86de41b
test: add schema_utils tests and expand interleaved coverage
VibhuJawa Mar 31, 2026
9c35387
Merge remote-tracking branch 'origin/main' into feat/interleaved-io-f…
VibhuJawa Apr 1, 2026
9e6f70a
feat(interleaved): add InterleavedParquetReaderStage and InterleavedW…
VibhuJawa Apr 1, 2026
f1c874c
refactor(interleaved): rename WebdatasetReader to InterleavedWebdatas…
VibhuJawa Apr 1, 2026
9d47188
feat(benchmarking): add input row/size tracking to interleaved IO ben…
VibhuJawa Apr 1, 2026
7eee1bf
chore: untrack local-only benchmark runner script
VibhuJawa Apr 1, 2026
cf7abd0
Merge remote-tracking branch 'origin/main' into feat/interleaved-parq…
VibhuJawa Apr 1, 2026
5e7a6eb
feat(benchmarking): count samples in WDS input/output metrics
VibhuJawa Apr 1, 2026
31c450a
feat(benchmarking): refactor interleaved IO metrics and validation
VibhuJawa Apr 3, 2026
05aec31
Merge branch 'main' into feat/interleaved-parquet-reader-wds-writer
VibhuJawa Apr 3, 2026
c53a957
Merge branch 'main' into feat/interleaved-parquet-reader-wds-writer
VibhuJawa Apr 3, 2026
c9dbbcc
perf(interleaved): cache storage_options in __post_init__ and replace…
VibhuJawa Apr 3, 2026
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
215 changes: 163 additions & 52 deletions benchmarking/scripts/multimodal_mint1t_benchmark.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.

"""Benchmark for multimodal MINT1T workflow: WebDataset -> filter -> parquet."""
"""Benchmark for multimodal MINT1T workflow: WebDataset/Parquet -> filter -> Parquet/WebDataset."""

import argparse
import time
Expand All @@ -21,57 +21,143 @@
from typing import Any

from loguru import logger
from utils import collect_parquet_output_metrics, setup_executor, validate_parquet_ordering, write_benchmark_results
from utils import (
collect_interleaved_parquet_metrics,
collect_interleaved_wds_metrics,
setup_executor,
validate_parquet_ordering,
validate_wds_ordering,
write_benchmark_results,
)

from nemo_curator.pipeline import Pipeline
from nemo_curator.stages.interleaved.io import InterleavedParquetWriterStage, WebdatasetReader
from nemo_curator.stages.interleaved.io import (
InterleavedParquetReader,
InterleavedParquetWriterStage,
InterleavedWebdatasetReader,
InterleavedWebdatasetWriterStage,
)
from nemo_curator.stages.interleaved.stages import InterleavedAspectRatioFilterStage
from nemo_curator.tasks.utils import TaskPerfUtils


def create_pipeline(args: argparse.Namespace) -> Pipeline:
read_kwargs = {}
write_kwargs = {}
if args.parquet_row_group_size is not None:
write_kwargs["row_group_size"] = args.parquet_row_group_size
if args.parquet_compression is not None:
write_kwargs["compression"] = args.parquet_compression
pipeline = Pipeline(
name="multimodal_mint1t_benchmark",
description="Benchmark: WebDataset MINT1T to multimodal parquet",
description="Benchmark: multimodal interleaved IO pipeline",
)
pipeline.add_stage(
WebdatasetReader(
file_paths=args.input_path,
files_per_partition=args.files_per_partition,
blocksize=args.input_blocksize,
max_batch_bytes=args.output_max_batch_bytes,
read_kwargs=read_kwargs,
materialize_on_read=args.materialize_on_read,
per_image_fields=tuple(args.per_image_fields) if args.per_image_fields else (),
per_text_fields=tuple(args.per_text_fields) if args.per_text_fields else (),

# ── Reader ────────────────────────────────────────────────────────────────
if args.reader_type == "wds":
pipeline.add_stage(
InterleavedWebdatasetReader(
file_paths=args.input_path,
files_per_partition=args.files_per_partition,
blocksize=args.input_blocksize,
max_batch_bytes=args.output_max_batch_bytes,
materialize_on_read=args.materialize_on_read,
per_image_fields=tuple(args.per_image_fields) if args.per_image_fields else (),
per_text_fields=tuple(args.per_text_fields) if args.per_text_fields else (),
)
)
)
pipeline.add_stage(
InterleavedAspectRatioFilterStage(drop_invalid_rows=True, min_aspect_ratio=1.0, max_aspect_ratio=2.0)
)
pipeline.add_stage(
InterleavedParquetWriterStage(
path=args.output_path,
materialize_on_write=args.materialize_on_write,
write_kwargs=write_kwargs,
mode=args.mode,
else: # parquet
pipeline.add_stage(
InterleavedParquetReader(
file_paths=args.input_path,
files_per_partition=args.files_per_partition,
blocksize=args.input_blocksize,
max_batch_bytes=args.output_max_batch_bytes,
fields=tuple(args.reader_fields) if args.reader_fields else None,
)
)
)

# ── Optional filter ───────────────────────────────────────────────────────
if not args.no_filter:
pipeline.add_stage(
InterleavedAspectRatioFilterStage(drop_invalid_rows=True, min_aspect_ratio=1.0, max_aspect_ratio=2.0)
)

# ── Writer ────────────────────────────────────────────────────────────────
if args.writer_format == "wds":
pipeline.add_stage(
InterleavedWebdatasetWriterStage(
path=args.output_path,
materialize_on_write=args.materialize_on_write,
on_materialize_error=args.on_materialize_error,
mode=args.mode,
)
)
else:
write_kwargs: dict[str, Any] = {}
if args.parquet_row_group_size is not None:
write_kwargs["row_group_size"] = args.parquet_row_group_size
if args.parquet_compression is not None:
write_kwargs["compression"] = args.parquet_compression
pipeline.add_stage(
InterleavedParquetWriterStage(
path=args.output_path,
materialize_on_write=args.materialize_on_write,
on_materialize_error=args.on_materialize_error,
write_kwargs=write_kwargs,
mode=args.mode,
)
)

return pipeline


def _validate_output(writer_format: str, output_path: Path) -> tuple[bool | None, bool | None]:
"""Spot-check one output file for the given format.

Returns (ordering_valid, wds_valid). None = not applicable or no files found.
"""
if writer_format == "wds":
tar = next(output_path.glob("*.tar"), None)
if tar is None:
logger.warning("WDS validation skipped: no output tars found")
return None, None
result = validate_wds_ordering(tar)
ordering_valid = result["ordering_valid"]
wds_valid = result["valid"]
if not wds_valid:
logger.error("WDS output validation failed: {}", result["errors"])
else:
logger.info(
"WDS validation passed on {}: {} samples, {} images, ordering_valid={}",
tar.name,
result["num_samples"],
result["num_images"],
ordering_valid,
)
return ordering_valid, wds_valid
else:
pq_file = next(output_path.glob("*.parquet"), None)
if pq_file is None:
logger.warning("Parquet ordering validation skipped: no output files found")
return None, None
result = validate_parquet_ordering(pq_file)
ordering_valid = result["valid"]
if not ordering_valid:
logger.error("Parquet ordering validation failed on {}: {}", pq_file.name, result["errors"])
else:
logger.info("Parquet ordering validation passed on {}", pq_file.name)
return ordering_valid, None # wds_valid=None: not applicable for parquet


def run_benchmark(args: argparse.Namespace) -> dict[str, Any]:
executor = setup_executor(args.executor)
input_path = str(Path(args.input_path).absolute())
output_path = Path(args.output_path).absolute()
output_path.mkdir(parents=True, exist_ok=True)

input_metrics_start = time.perf_counter()
collect_fn = (
collect_interleaved_parquet_metrics if args.reader_type == "parquet" else collect_interleaved_wds_metrics
)
input_metrics = {f"input_{k}": v for k, v in collect_fn(args.input_path).items()}
input_metrics_elapsed = time.perf_counter() - input_metrics_start
logger.info("collect_input_metrics took {:.3f}s", input_metrics_elapsed)

start = time.perf_counter()
output_tasks = []
success = False
Expand All @@ -86,35 +172,36 @@ def run_benchmark(args: argparse.Namespace) -> dict[str, Any]:

elapsed = time.perf_counter() - start
metrics_start = time.perf_counter()
output_metrics = collect_parquet_output_metrics(output_path)
collect_fn = (
collect_interleaved_parquet_metrics if args.writer_format == "parquet" else collect_interleaved_wds_metrics
)
output_metrics = {f"output_{k}": v for k, v in collect_fn(output_path).items()}
ordering_valid, wds_valid = None, None
if success:
ordering_valid, wds_valid = _validate_output(args.writer_format, output_path)
metrics_elapsed = time.perf_counter() - metrics_start
logger.info("collect_parquet_output_metrics took {:.3f}s", metrics_elapsed)
logger.info("collect_output_metrics took {:.3f}s", metrics_elapsed)
task_metrics = TaskPerfUtils.aggregate_task_metrics(output_tasks, prefix="task")
writer_stats = {k: v for k, v in task_metrics.items() if "multimodal_" in k and "_writer" in k}
writer_stats = {k: v for k, v in task_metrics.items() if "interleaved_" in k and "_writer" in k}
logger.info("Writer stage stats: {}", writer_stats)

ordering_valid = False
if success:
parquet_files = sorted(output_path.glob("*.parquet"))
if parquet_files:
result = validate_parquet_ordering(parquet_files[0])
ordering_valid = result["valid"]
if not ordering_valid:
logger.error("Ordering validation failed on {}: {}", parquet_files[0].name, result["errors"])
else:
logger.info("Ordering validation passed on {}", parquet_files[0].name)

rows = output_metrics["num_rows"]
rows = output_metrics["output_num_rows"]
samples = output_metrics["output_num_samples"]
return {
"params": {
"executor": args.executor,
"input_path": input_path,
"output_path": str(output_path),
"reader_type": args.reader_type,
"writer_format": args.writer_format,
"files_per_partition": args.files_per_partition,
"input_blocksize": args.input_blocksize,
"output_max_batch_bytes": args.output_max_batch_bytes,
"materialize_on_read": args.materialize_on_read,
"materialize_on_write": args.materialize_on_write,
"on_materialize_error": args.on_materialize_error,
"reader_fields": list(args.reader_fields),
"no_filter": args.no_filter,
"per_image_fields": list(args.per_image_fields) if args.per_image_fields else [],
"per_text_fields": list(args.per_text_fields) if args.per_text_fields else [],
"parquet_row_group_size": args.parquet_row_group_size,
Expand All @@ -124,8 +211,11 @@ def run_benchmark(args: argparse.Namespace) -> dict[str, Any]:
"metrics": {
"is_success": success,
"ordering_valid": ordering_valid,
"wds_valid": wds_valid,
"time_taken_s": elapsed,
"throughput_rows_per_sec": (rows / elapsed) if elapsed > 0 else 0.0,
"throughput_rows_per_sec": (rows / elapsed) if (elapsed > 0 and rows > 0) else 0.0,
"throughput_samples_per_sec": (samples / elapsed) if (elapsed > 0 and samples > 0) else 0.0,
**input_metrics,
**task_metrics,
**output_metrics,
},
Expand All @@ -134,24 +224,45 @@ def run_benchmark(args: argparse.Namespace) -> dict[str, Any]:


def main() -> int:
parser = argparse.ArgumentParser(description="Multimodal MINT1T benchmark")
parser = argparse.ArgumentParser(description="Multimodal interleaved IO benchmark")
parser.add_argument("--benchmark-results-path", type=Path, required=True)
parser.add_argument("--executor", default="xenna", choices=["xenna", "ray_data"])
parser.add_argument("--executor", default="ray_data", choices=["xenna", "ray_data"])
parser.add_argument("--input-path", type=str, required=True)
parser.add_argument("--output-path", type=str, required=True)
parser.add_argument("--reader-type", default="wds", choices=["wds", "parquet"])
parser.add_argument("--writer-format", default="parquet", choices=["parquet", "wds"])
parser.add_argument("--files-per-partition", type=int, default=1)
parser.add_argument("--input-blocksize", type=str, default=None)
parser.add_argument("--output-max-batch-bytes", type=int, default=None)
parser.add_argument(
"--reader-fields",
nargs="*",
default=[
"image_metadata",
"url",
"language_id_whole_page_fasttext",
"pdf_name",
"previous_word_count",
"bff_contained_ngram_count_before_dedupe",
],
)
parser.add_argument("--materialize-on-read", action="store_true", dest="materialize_on_read")
parser.add_argument("--no-materialize-on-read", action="store_false", dest="materialize_on_read")
parser.add_argument("--parquet-row-group-size", type=int, default=None)
parser.add_argument("--parquet-compression", type=str, default=None)
parser.add_argument("--materialize-on-write", action="store_true", dest="materialize_on_write")
parser.add_argument("--no-materialize-on-write", action="store_false", dest="materialize_on_write")
parser.add_argument(
"--on-materialize-error",
default="error",
choices=["error", "warn", "drop_row", "drop_sample"],
dest="on_materialize_error",
)
parser.add_argument("--no-filter", action="store_true", default=False)
parser.add_argument("--parquet-row-group-size", type=int, default=None)
parser.add_argument("--parquet-compression", type=str, default=None)
parser.add_argument("--mode", type=str, default="overwrite", choices=["ignore", "overwrite", "append", "error"])
parser.add_argument("--per-image-fields", nargs="*", default=["image_metadata"])
parser.add_argument("--per-text-fields", nargs="*", default=[])
parser.set_defaults(materialize_on_write=False, materialize_on_read=False)
parser.set_defaults(materialize_on_write=True, materialize_on_read=True)
args = parser.parse_args()

try:
Expand Down
Loading
Loading