diff --git a/benchmarking/scripts/multimodal_mint1t_benchmark.py b/benchmarking/scripts/multimodal_mint1t_benchmark.py index 20897e8845..e874a7429e 100644 --- a/benchmarking/scripts/multimodal_mint1t_benchmark.py +++ b/benchmarking/scripts/multimodal_mint1t_benchmark.py @@ -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 @@ -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 @@ -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, @@ -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, }, @@ -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: diff --git a/benchmarking/scripts/utils.py b/benchmarking/scripts/utils.py index 1e187a9f32..23b17d8aa0 100644 --- a/benchmarking/scripts/utils.py +++ b/benchmarking/scripts/utils.py @@ -12,12 +12,17 @@ # See the License for the specific language governing permissions and # limitations under the License. +import itertools import json import pickle +import tarfile +from dataclasses import dataclass, field from pathlib import Path from typing import Any import pandas as pd +import pyarrow as pa +import pyarrow.compute as pc import pyarrow.parquet as pq from nemo_curator.backends.experimental.ray_actor_pool.executor import RayActorPoolExecutor @@ -98,44 +103,247 @@ def write_benchmark_results(results: dict, output_path: str | Path) -> None: (output_path / "tasks.pkl").write_bytes(pickle.dumps(results["tasks"])) -def collect_parquet_output_metrics(output_path: Path) -> dict[str, Any]: +def _collect_file_size_metrics(output_path: Path, extensions: list[str]) -> tuple[list[str], int, int]: + """Return (file_paths, num_files, total_size_bytes) for files matching extensions under output_path.""" output_files_with_size = get_all_file_paths_and_size_under( str(output_path), recurse_subdirectories=True, - keep_extensions=[".parquet"], + keep_extensions=extensions, ) - parquet_files = [path for path, _ in output_files_with_size] - num_files = len(parquet_files) + file_paths = [path for path, _ in output_files_with_size] total_size_bytes = int(sum(size for _, size in output_files_with_size)) + return file_paths, len(file_paths), total_size_bytes + + +def _resolve_paths(path: Path, extension: str) -> tuple[list[str], int, int]: + """Return (file_paths, num_files, total_size_bytes) for a single file or a directory.""" + if path.is_file(): + return [str(path)], 1, path.stat().st_size + return _collect_file_size_metrics(path, [extension]) + + +def _accumulate_modality_counts(column: pa.ChunkedArray, into: dict[str, int]) -> None: + """Accumulate value_counts from a modality column into `into`.""" + for row in column.value_counts().to_pylist(): + key = str(row["values"]) if row["values"] is not None else "None" + into[key] = into.get(key, 0) + int(row["counts"]) + + +def collect_interleaved_parquet_metrics(path: Path | str) -> dict[str, Any]: + """Collect metrics for interleaved parquet files — neutral keys, caller adds input_/output_ prefix.""" + parquet_files, num_files, total_size_bytes = _resolve_paths(Path(path), ".parquet") num_rows = 0 + num_samples = 0 modality_counts: dict[str, int] = {} materialize_error_count = 0 - for path in parquet_files: - pf = pq.ParquetFile(path) + for pq_path in parquet_files: + pf = pq.ParquetFile(pq_path) num_rows += pf.metadata.num_rows schema_names = set(pf.schema_arrow.names) - cols = [c for c in ("modality", "materialize_error") if c in schema_names] + cols = [c for c in ("sample_id", "modality", "materialize_error") if c in schema_names] if not cols: continue - table = pq.read_table(path, columns=cols) - if "modality" in table.column_names: - counts = table.column("modality").value_counts() - for row in counts.to_pylist(): - key = str(row["values"]) if row["values"] is not None else "None" - modality_counts[key] = modality_counts.get(key, 0) + int(row["counts"]) - if "materialize_error" in table.column_names: + cols_set = set(cols) + table = pf.read(columns=cols) + if "sample_id" in cols_set: + num_samples += pc.count_distinct(table.column("sample_id")).as_py() + if "modality" in cols_set: + _accumulate_modality_counts(table.column("modality"), modality_counts) + if "materialize_error" in cols_set: col = table.column("materialize_error") materialize_error_count += col.length() - col.null_count return { - "num_output_files": num_files, - "output_total_bytes": total_size_bytes, - "output_total_mb": total_size_bytes / (1024 * 1024), + "num_files": num_files, + "total_bytes": total_size_bytes, + "total_mb": total_size_bytes / 1e6, "num_rows": num_rows, + "num_samples": num_samples, + "num_metadata": modality_counts.get("metadata", 0), + "num_texts": modality_counts.get("text", 0), + "num_images": modality_counts.get("image", 0), "modality_counts": modality_counts, "materialize_error_count": materialize_error_count, } +def _collect_wds_modality_counts(tar_paths: list[str]) -> tuple[int, dict[str, int]]: + """Return (num_samples, modality_counts) by reading JSON metadata from WDS tars. + + Counts metadata (one per sample), text (non-null texts entries), and image + (non-null images entries) rows. + """ + counts: dict[str, int] = {} + for path in tar_paths: + with tarfile.open(path) as tf: + for m in tf.getmembers(): + if not m.name.endswith(".json"): + continue + raw = tf.extractfile(m) + if raw is None: + continue + payload = json.loads(raw.read()) + counts["metadata"] = counts.get("metadata", 0) + 1 + text_count = sum(1 for t in payload.get("texts", []) if t is not None) + image_count = sum(1 for img in payload.get("images", []) if img is not None) + if text_count: + counts["text"] = counts.get("text", 0) + text_count + if image_count: + counts["image"] = counts.get("image", 0) + image_count + return counts.get("metadata", 0), counts + + +def collect_interleaved_wds_metrics(path: Path | str) -> dict[str, Any]: + """Collect metrics for interleaved WebDataset tar archives — neutral keys, caller adds input_/output_ prefix.""" + tar_paths, num_files, total_size_bytes = _resolve_paths(Path(path), ".tar") + num_samples, modality_counts = _collect_wds_modality_counts(tar_paths) + total_rows = sum(modality_counts.values()) + return { + "num_files": num_files, + "total_bytes": total_size_bytes, + "total_mb": total_size_bytes / 1e6, + "num_rows": total_rows, + "num_samples": num_samples, + "num_texts": modality_counts.get("text", 0), + "num_images": modality_counts.get("image", 0), + "modality_counts": modality_counts, + } + + +_BUNCHING_RUN_THRESHOLD = 0.9 # per-sample: suspicious if longest run of one type > 90% of its count +_BUNCHING_SAMPLE_THRESHOLD = 0.7 # aggregate: ordering_valid=False if >70% of samples are suspicious +_MIN_COUNT_FOR_BUNCHING = 2 # need at least 2 elements of a type to check for bunching + + +@dataclass +class _WdsValidationAcc: + """Accumulates errors and ordering errors while scanning WDS tars.""" + + errors: list[str] = field(default_factory=list) + ordering_errors: list[str] = field(default_factory=list) + + +def _is_sample_suspicious(texts: list, images: list) -> bool: + """Return True if one type's elements are overly bunched (longest run > 90% of its total count).""" + sequence = [ + "T" if t is not None else "I" + for t, img in zip(texts, images, strict=False) + if t is not None or img is not None + ] + text_total = sequence.count("T") + image_total = len(sequence) - text_total + for type_char, total in (("T", text_total), ("I", image_total)): + if total < _MIN_COUNT_FOR_BUNCHING: + continue + max_run = max( + (sum(1 for _ in g) for ch, g in itertools.groupby(sequence) if ch == type_char), + default=0, + ) + if max_run / total > _BUNCHING_RUN_THRESHOLD: + return True + return False + + +def _check_interleaving( + sample_id: str, + texts: list, + images: list, + member_names: set[str], + acc: _WdsValidationAcc, +) -> bool: + """Check per-sample hard ordering errors. Returns True if sample is suspicious (bunching).""" + if len(texts) != len(images): + acc.ordering_errors.append(f"{sample_id}: texts length {len(texts)} != images length {len(images)}") + return False + for pos, (t, img) in enumerate(zip(texts, images, strict=False)): + if t is not None and img is not None: + acc.ordering_errors.append(f"{sample_id}: position {pos} has both text and image") + elif img is not None and img not in member_names: + acc.ordering_errors.append(f"{sample_id}: image '{img}' at position {pos} not found in tar") + return _is_sample_suspicious(texts, images) + + +def _check_wds_sample( + tf: tarfile.TarFile, + sample_id: str, + infos: list[tarfile.TarInfo], + acc: _WdsValidationAcc, + image_exts: set[str], +) -> tuple[int, bool] | None: + """Validate one WDS sample. Returns (image_count, is_suspicious), or None on hard error.""" + json_infos = [m for m in infos if m.name.endswith(".json")] + if not json_infos: + acc.errors.append(f"{sample_id}: no .json metadata file") + return None + try: + raw = tf.extractfile(json_infos[0]) + if raw is None: + acc.errors.append(f"{sample_id}: .json is not a regular file") + return None + payload = json.loads(raw.read()) + except Exception as e: + acc.errors.append(f"{sample_id}: .json parse error: {e}") + return None + + texts = payload.get("texts", []) + images = payload.get("images", []) + member_names = {m.name for m in infos} + suspicious = _check_interleaving(sample_id, texts, images, member_names, acc) + image_count = sum(1 for m in infos if Path(m.name).suffix.lower() in image_exts) + return image_count, suspicious + + +def validate_wds_ordering(tar_path: Path | str) -> dict[str, Any]: + """Validate ordering within a single WebDataset tar file. + + Checks per-sample hard errors (array length mismatch, position collision, missing image) + and aggregate bunching. ``ordering_valid`` is False if >70% of samples are suspicious. + + Returns validation fields only: 'valid', 'ordering_valid', 'errors', + 'num_samples', 'num_suspicious_samples', 'num_images'. + """ + from nemo_curator.stages.interleaved.utils.constants import DEFAULT_IMAGE_EXTENSIONS + + tar_path = Path(tar_path) + acc = _WdsValidationAcc() + num_samples = 0 + num_images = 0 + num_suspicious = 0 + image_exts = set(DEFAULT_IMAGE_EXTENSIONS) + + try: + with tarfile.open(tar_path) as tf: + members = tf.getmembers() + # Build sample→member map in a single pass while the tar is open. + # Members follow {sample_id}.json and {sample_id}.{position}.{ext} + # naming, so the sample key is the portion before the first dot. + samples: dict[str, list[tarfile.TarInfo]] = {} + for m in members: + samples.setdefault(m.name.split(".")[0], []).append(m) + + for key, infos in samples.items(): + sample_id = f"{tar_path.name}: sample '{key}'" + result = _check_wds_sample(tf, sample_id, infos, acc, image_exts) + if result is not None: + image_count, suspicious = result + num_samples += 1 + num_images += image_count + if suspicious: + num_suspicious += 1 + except Exception as e: + acc.errors.append(f"{tar_path}: failed to open tar: {e}") + + suspicious_ratio = num_suspicious / num_samples if num_samples > 0 else 0.0 + ordering_valid = len(acc.ordering_errors) == 0 and suspicious_ratio <= _BUNCHING_SAMPLE_THRESHOLD + return { + "valid": len(acc.errors) == 0 and ordering_valid, + "ordering_valid": ordering_valid, + "errors": acc.errors + acc.ordering_errors, + "num_samples": num_samples, + "num_suspicious_samples": num_suspicious, + "num_images": num_images, + } + + def validate_parquet_ordering(parquet_path: str | Path) -> dict[str, Any]: """Read a single parquet file and validate interleaved position ordering. diff --git a/nemo_curator/stages/interleaved/README.md b/nemo_curator/stages/interleaved/README.md index d696cf67ca..d84298c20c 100644 --- a/nemo_curator/stages/interleaved/README.md +++ b/nemo_curator/stages/interleaved/README.md @@ -5,26 +5,30 @@ Row-wise interleaved multimodal ingestion and write path for WebDataset tar shar ## Architecture ``` -WebDataset tar shards - | - v -┌─────────────────────────┐ -│ WebdatasetReader │ CompositeStage: FilePartitioning + InterleavedWebdatasetReaderStage -│ (io/reader.py) │ Parses tar members -> normalized interleaved rows -└────────┬────────────────┘ - | InterleavedBatch (Arrow/Pandas) - v -┌─────────────────────────┐ -│ Filter Stages │ e.g. InterleavedAspectRatioFilterStage -│ (stages.py) │ Row-wise filtering with optional materialization -└────────┬────────────────┘ - | - v -┌─────────────────────────┐ -│ InterleavedParquet- │ InterleavedParquetWriterStage -│ WriterStage │ Parquet output with optional materialize-on-write -│ (io/writers/tabular.py)│ Supports snappy/zstd compression, configurable row groups -└─────────────────────────┘ +WebDataset tar shards Parquet files + | | + v v +┌──────────────────────────┐ ┌──────────────────────────┐ +│ InterleavedWebdataset- │ │ InterleavedParquetReader │ Both are CompositeStages: +│ Reader (io/reader.py) │ │ (io/reader.py) │ FilePartitioningStage + +│ │ │ │ ReaderStage +└──────────┬───────────────┘ └────────────┬─────────────┘ + └──────────┬────────────────────┘ + | InterleavedBatch (Arrow/Pandas) + v + ┌─────────────────────────┐ + │ Filter Stages │ e.g. InterleavedAspectRatioFilterStage + │ (stages.py) │ Row-wise filtering with optional materialization + └────────┬────────────────┘ + | + ┌─────────┴──────────┐ + v v +┌───────────────┐ ┌──────────────────────────┐ +│ Interleaved- │ │ InterleavedWebdataset- │ +│ ParquetWriter │ │ WriterStage │ +│ Stage │ │ (io/writers/webdataset.py)│ +│ (tabular.py) │ │ MINT-1T-style tar shards │ +└───────────────┘ └──────────────────────────┘ ``` ## Schema (`INTERLEAVED_SCHEMA`) @@ -51,7 +55,7 @@ These are set and managed by pipeline stages. Users should not write to them dir Extra fields from the source data flow through the pipeline as additional columns. Specify them with the `fields` parameter on the reader: ```python -reader = WebdatasetReader( +reader = InterleavedWebdatasetReader( file_paths="/data/shards/", fields=("p_hash", "score", "aux"), # These become extra columns ) @@ -111,11 +115,11 @@ Materialization can happen at read time (`materialize_on_read=True`) or write ti ```python from nemo_curator.pipeline import Pipeline -from nemo_curator.stages.interleaved.io import WebdatasetReader, InterleavedParquetWriterStage +from nemo_curator.stages.interleaved.io import InterleavedWebdatasetReader, InterleavedParquetWriterStage from nemo_curator.stages.interleaved.stages import InterleavedAspectRatioFilterStage pipeline = Pipeline(name="mint1t_pipeline") -pipeline.add_stage(WebdatasetReader( +pipeline.add_stage(InterleavedWebdatasetReader( file_paths="/data/mint1t/shards/", )) pipeline.add_stage(InterleavedAspectRatioFilterStage(drop_invalid_rows=True)) @@ -135,14 +139,17 @@ stages/interleaved/ ├── stages.py # BaseInterleavedAnnotatorStage, BaseInterleavedFilterStage, │ # InterleavedAspectRatioFilterStage ├── io/ -│ ├── __init__.py # Exports WebdatasetReader, InterleavedParquetWriterStage -│ ├── reader.py # WebdatasetReader (CompositeStage) +│ ├── __init__.py # Exports InterleavedWebdatasetReader, InterleavedParquetReader, +│ │ # InterleavedParquetWriterStage, InterleavedWebdatasetWriterStage +│ ├── reader.py # InterleavedWebdatasetReader, InterleavedParquetReader (CompositeStages) │ ├── readers/ │ │ ├── base.py # BaseInterleavedReader +│ │ ├── parquet.py # InterleavedParquetReaderStage (ProcessingStage) │ │ └── webdataset.py # InterleavedWebdatasetReaderStage (ProcessingStage) │ └── writers/ │ ├── base.py # BaseInterleavedWriter (filesystem + materialization + process) -│ └── tabular.py # InterleavedParquetWriterStage +│ ├── tabular.py # InterleavedParquetWriterStage +│ └── webdataset.py # InterleavedWebdatasetWriterStage └── utils/ ├── constants.py # Default file extensions ├── materialization.py # Three-strategy materialization dispatch diff --git a/nemo_curator/stages/interleaved/io/__init__.py b/nemo_curator/stages/interleaved/io/__init__.py index c20398cb87..0dde15093d 100644 --- a/nemo_curator/stages/interleaved/io/__init__.py +++ b/nemo_curator/stages/interleaved/io/__init__.py @@ -12,10 +12,13 @@ # See the License for the specific language governing permissions and # limitations under the License. -from nemo_curator.stages.interleaved.io.reader import WebdatasetReader +from nemo_curator.stages.interleaved.io.reader import InterleavedParquetReader, InterleavedWebdatasetReader from nemo_curator.stages.interleaved.io.writers.tabular import InterleavedParquetWriterStage +from nemo_curator.stages.interleaved.io.writers.webdataset import InterleavedWebdatasetWriterStage __all__ = [ + "InterleavedParquetReader", "InterleavedParquetWriterStage", - "WebdatasetReader", + "InterleavedWebdatasetReader", + "InterleavedWebdatasetWriterStage", ] diff --git a/nemo_curator/stages/interleaved/io/reader.py b/nemo_curator/stages/interleaved/io/reader.py index 55c8355f7c..404c52fc16 100644 --- a/nemo_curator/stages/interleaved/io/reader.py +++ b/nemo_curator/stages/interleaved/io/reader.py @@ -15,10 +15,15 @@ from __future__ import annotations from dataclasses import dataclass, field -from typing import Any +from typing import TYPE_CHECKING, Any from nemo_curator.stages.base import CompositeStage + +if TYPE_CHECKING: + import pyarrow as pa + from nemo_curator.stages.file_partitioning import FilePartitioningStage +from nemo_curator.stages.interleaved.io.readers.parquet import InterleavedParquetReaderStage from nemo_curator.stages.interleaved.io.readers.webdataset import InterleavedWebdatasetReaderStage from nemo_curator.stages.interleaved.utils import ( DEFAULT_IMAGE_EXTENSIONS, @@ -30,7 +35,7 @@ @dataclass -class WebdatasetReader(CompositeStage[_EmptyTask, InterleavedBatch]): +class InterleavedWebdatasetReader(CompositeStage[_EmptyTask, InterleavedBatch]): """Composite stage for reading WebDataset shards.""" file_paths: str | list[str] @@ -49,7 +54,7 @@ class WebdatasetReader(CompositeStage[_EmptyTask, InterleavedBatch]): fields: tuple[str, ...] | None = None per_image_fields: tuple[str, ...] = () per_text_fields: tuple[str, ...] = () - name: str = "webdataset_reader" + name: str = "interleaved_webdataset_reader" def __post_init__(self): super().__init__() @@ -79,3 +84,41 @@ def decompose(self) -> list: per_text_fields=self.per_text_fields, ), ] + + +@dataclass +class InterleavedParquetReader(CompositeStage[_EmptyTask, InterleavedBatch]): + """Composite stage for reading interleaved Parquet files.""" + + file_paths: str | list[str] + files_per_partition: int | None = None + blocksize: int | str | None = None + fields: tuple[str, ...] | None = None + max_batch_bytes: int | None = None + read_kwargs: dict[str, Any] = field(default_factory=dict) + schema: pa.Schema | None = None + schema_overrides: dict[str, pa.DataType] | None = None + file_extensions: list[str] = field(default_factory=lambda: [".parquet"]) + name: str = "interleaved_parquet_reader" + + def __post_init__(self): + super().__init__() + self.storage_options = resolve_storage_options(io_kwargs=self.read_kwargs) + + def decompose(self) -> list: + return [ + FilePartitioningStage( + file_paths=self.file_paths, + files_per_partition=self.files_per_partition, + blocksize=self.blocksize, + file_extensions=self.file_extensions, + storage_options=self.storage_options, + ), + InterleavedParquetReaderStage( + read_kwargs=self.read_kwargs, + fields=self.fields, + max_batch_bytes=self.max_batch_bytes, + schema=self.schema, + schema_overrides=self.schema_overrides, + ), + ] diff --git a/nemo_curator/stages/interleaved/io/readers/__init__.py b/nemo_curator/stages/interleaved/io/readers/__init__.py index ab4a85faf9..0d1bb0f4a1 100644 --- a/nemo_curator/stages/interleaved/io/readers/__init__.py +++ b/nemo_curator/stages/interleaved/io/readers/__init__.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +from nemo_curator.stages.interleaved.io.readers.parquet import InterleavedParquetReaderStage from nemo_curator.stages.interleaved.io.readers.webdataset import InterleavedWebdatasetReaderStage -__all__ = ["InterleavedWebdatasetReaderStage"] +__all__ = ["InterleavedParquetReaderStage", "InterleavedWebdatasetReaderStage"] diff --git a/nemo_curator/stages/interleaved/io/readers/base.py b/nemo_curator/stages/interleaved/io/readers/base.py index 8006d2c160..aa75be25fb 100644 --- a/nemo_curator/stages/interleaved/io/readers/base.py +++ b/nemo_curator/stages/interleaved/io/readers/base.py @@ -16,6 +16,7 @@ from typing import Any import pyarrow as pa +from loguru import logger from nemo_curator.stages.base import ProcessingStage from nemo_curator.stages.interleaved.utils.schema import align_table, reconcile_schema, resolve_schema @@ -63,3 +64,31 @@ def _align_output(self, table: pa.Table) -> pa.Table: if self.schema is not None: return align_table(table, self.schema) return table.cast(reconcile_schema(table.schema)) + + @staticmethod + def _source_files_for_split( + split: pa.Table, + idx: int, + sample_id_to_path: dict[str, str], + all_paths: list[str], + ) -> list[str]: + """Return source_files for one split, annotated with the split index for lineage tracking. + + The ``::split_NNN`` suffix is appended so that downstream consumers can correlate + each output batch back to the exact split of its source file(s), even when a single + source file is split into multiple batches by ``max_batch_bytes``. + """ + seen: set[str] = set() + for sid in split["sample_id"].unique().to_pylist(): + path = sample_id_to_path.get(sid) + if path is not None: + seen.add(path) + contributing = [p for p in all_paths if p in seen] + if not contributing: + logger.warning( + "_source_files_for_split: no source path found for any sample_id in this split " + "(possible null sample_ids); falling back to all {} source path(s).", + len(all_paths), + ) + contributing = all_paths + return [f"{p}::split_{idx:05d}" for p in contributing] diff --git a/nemo_curator/stages/interleaved/io/readers/parquet.py b/nemo_curator/stages/interleaved/io/readers/parquet.py new file mode 100644 index 0000000000..0e8260ce3b --- /dev/null +++ b/nemo_curator/stages/interleaved/io/readers/parquet.py @@ -0,0 +1,136 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +import pyarrow as pa +import pyarrow.parquet as pq +from fsspec.core import url_to_fs +from pyarrow.fs import FSSpecHandler, PyFileSystem + +from nemo_curator.core.utils import split_table_by_group_max_bytes +from nemo_curator.stages.interleaved.utils import resolve_storage_options +from nemo_curator.tasks import FileGroupTask, InterleavedBatch +from nemo_curator.tasks.interleaved import INTERLEAVED_SCHEMA, RESERVED_COLUMNS + +from .base import BaseInterleavedReader + + +@dataclass +class InterleavedParquetReaderStage(BaseInterleavedReader): + """Read interleaved Parquet files into an ``InterleavedBatch``. + + *fields* lists extra (passthrough) column names to read beyond the reserved + schema columns. Any *fields* entry that is absent from a given file is + null-filled, consistent with how the WebDataset reader handles ``fields``. + Reserved columns are always read regardless of *fields*. + + When *max_batch_bytes* is set, the combined table is split into multiple + batches so that no single batch exceeds the byte limit. Each split's + ``source_files`` metadata lists only the parquet files that contributed + rows to that batch. + """ + + fields: tuple[str, ...] | None = None + max_batch_bytes: int | None = None + name: str = "interleaved_parquet_reader" + + def __post_init__(self) -> None: + super().__post_init__() + self._storage_options = resolve_storage_options(io_kwargs=self.read_kwargs) + + def _columns_to_read(self, file_schema: pa.Schema) -> list[str] | None: + """Return the column list to pass to ``pq.read_table``. + + When *fields* is ``None`` (the default) returns ``None``, which tells + PyArrow to read all columns — non-lossy by default, consistent with the + WebDataset reader. + + When *fields* is set, returns reserved columns plus those extra columns + that exist in the file; missing declared fields are null-filled after + the read. + """ + if self.fields is None: + return None + file_col_set = set(file_schema.names) + cols = [c for c in RESERVED_COLUMNS if c in file_col_set] + for f in self.fields: + if f in file_col_set and f not in RESERVED_COLUMNS: + cols.append(f) + return cols + + def _null_fill_missing_columns(self, table: pa.Table) -> pa.Table: + """Null-fill any reserved or extra *fields* columns absent from *table*. + + Handles both reserved columns (typed from INTERLEAVED_SCHEMA) and user-requested + passthrough fields (pa.null() typed, resolved later by _align_output). + A single set() pass avoids duplicate schema introspection. + """ + existing = set(table.schema.names) + for schema_field in INTERLEAVED_SCHEMA: + if schema_field.name not in existing: + table = table.append_column(schema_field, pa.nulls(len(table), type=schema_field.type)) + existing.add(schema_field.name) + if self.fields: + for f in self.fields: + if f not in existing: + table = table.append_column(pa.field(f, pa.null()), pa.nulls(len(table), type=pa.null())) + return table + + def process(self, task: FileGroupTask) -> InterleavedBatch | list[InterleavedBatch]: + tables: list[pa.Table] = [] + sample_id_to_path: dict[str, str] = {} + + for path in task.data: + fs, _ = url_to_fs(path, **(self._storage_options or {})) + pa_fs = PyFileSystem(FSSpecHandler(fs)) + file_schema = pq.read_schema(path, filesystem=pa_fs) + columns = self._columns_to_read(file_schema) + table = pq.read_table(path, columns=columns, filesystem=pa_fs) + for sid in table["sample_id"].unique().to_pylist(): + sample_id_to_path.setdefault(sid, path) + tables.append(table) + + if tables: + combined = pa.concat_tables(tables, promote_options="default") + combined = self._null_fill_missing_columns(combined) + combined = self._align_output(combined) + else: + base = self.schema if self.schema is not None else INTERLEAVED_SCHEMA + combined = pa.Table.from_pylist([], schema=base) + + splits = split_table_by_group_max_bytes(combined, "sample_id", self.max_batch_bytes) + batches: list[InterleavedBatch] = [] + for idx, split in enumerate(splits): + task_id = f"{task.task_id}_processed" if len(splits) == 1 else f"{task.task_id}_processed_{idx:05d}" + metadata: dict[str, Any] = dict(task._metadata) + if len(splits) == 1: + metadata["source_files"] = list(task.data) + else: + metadata["source_files"] = self._source_files_for_split(split, idx, sample_id_to_path, task.data) + if self._storage_options: + metadata["source_storage_options"] = self._storage_options + batches.append( + InterleavedBatch( + task_id=task_id, + dataset_name=task.dataset_name, + data=split, + _metadata=metadata, + _stage_perf=task._stage_perf, + ) + ) + return batches if len(batches) > 1 else batches[0] diff --git a/nemo_curator/stages/interleaved/io/readers/webdataset.py b/nemo_curator/stages/interleaved/io/readers/webdataset.py index 8becb652a4..767d442c7b 100644 --- a/nemo_curator/stages/interleaved/io/readers/webdataset.py +++ b/nemo_curator/stages/interleaved/io/readers/webdataset.py @@ -84,6 +84,7 @@ class InterleavedWebdatasetReaderStage(BaseInterleavedReader): def __post_init__(self) -> None: super().__post_init__() + self._storage_options = resolve_storage_options(io_kwargs=self.read_kwargs) # -- source_ref construction -- @@ -148,7 +149,7 @@ def _apply_per_modality_fields( for field_name, values in passthrough.items(): if index < len(values): val = values[index] - row[field_name] = json.dumps(val, ensure_ascii=True) if isinstance(val, (dict, list)) else val + row[field_name] = json.dumps(val, ensure_ascii=False) if isinstance(val, (dict, list)) else val @staticmethod def _warn_per_modality_length_mismatch( @@ -422,11 +423,10 @@ def _source_files_for_split( def process(self, task: FileGroupTask) -> InterleavedBatch | list[InterleavedBatch]: rows: list[dict[str, Any]] = [] sample_id_to_tar: dict[str, str] = {} - storage_options = resolve_storage_options(io_kwargs=self.read_kwargs) for tar_path in task.data: with ( - fsspec.open(tar_path, mode="rb", **storage_options) as fobj, + fsspec.open(tar_path, mode="rb", **self._storage_options) as fobj, tarfile.open(fileobj=fobj, mode="r:*") as tf, ): members = [m for m in tf.getmembers() if m.isfile()] @@ -435,7 +435,7 @@ def process(self, task: FileGroupTask) -> InterleavedBatch | list[InterleavedBat tar_path=tar_path, member_names=member_names, member_info={m.name: m for m in members}, - storage_options=storage_options, + storage_options=self._storage_options, byte_cache={}, ) for member in members: @@ -462,8 +462,8 @@ def process(self, task: FileGroupTask) -> InterleavedBatch | list[InterleavedBat metadata["source_files"] = list(task.data) else: metadata["source_files"] = self._source_files_for_split(split, idx, sample_id_to_tar, task.data) - if storage_options: - metadata["source_storage_options"] = storage_options + if self._storage_options: + metadata["source_storage_options"] = self._storage_options batches.append( InterleavedBatch( task_id=task_id, diff --git a/nemo_curator/stages/interleaved/io/writers/__init__.py b/nemo_curator/stages/interleaved/io/writers/__init__.py index 953557e310..073dceec2e 100644 --- a/nemo_curator/stages/interleaved/io/writers/__init__.py +++ b/nemo_curator/stages/interleaved/io/writers/__init__.py @@ -13,7 +13,9 @@ # limitations under the License. from nemo_curator.stages.interleaved.io.writers.tabular import InterleavedParquetWriterStage +from nemo_curator.stages.interleaved.io.writers.webdataset import InterleavedWebdatasetWriterStage __all__ = [ "InterleavedParquetWriterStage", + "InterleavedWebdatasetWriterStage", ] diff --git a/nemo_curator/stages/interleaved/io/writers/webdataset.py b/nemo_curator/stages/interleaved/io/writers/webdataset.py new file mode 100644 index 0000000000..efc59ef6e8 --- /dev/null +++ b/nemo_curator/stages/interleaved/io/writers/webdataset.py @@ -0,0 +1,217 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import json +import mimetypes +import tarfile +import urllib.parse +from dataclasses import dataclass +from io import BytesIO +from typing import Any, ClassVar + +import fsspec +import pandas as pd + +from nemo_curator.tasks.interleaved import RESERVED_COLUMNS + +from .base import BaseInterleavedWriter + +# --------------------------------------------------------------------------- +# Module-level helpers (importable in tests) +# --------------------------------------------------------------------------- + +_CONTENT_TYPE_TO_EXT: dict[str, str] = { + "image/jpeg": "jpg", + "image/png": "png", + "image/tiff": "tiff", + "image/webp": "webp", + "image/gif": "gif", + "image/bmp": "bmp", + "image/avif": "avif", +} + + +def _escape_key(sample_id: str) -> str: + """Percent-encode a sample_id so it is safe as a tar member name stem.""" + return urllib.parse.quote(sample_id, safe="") + + +def _ext_from_content_type(content_type: str | None) -> str: + """Return a file extension for *content_type*, falling back to ``"bin"``.""" + if content_type: + ext = _CONTENT_TYPE_TO_EXT.get(content_type) + if ext: + return ext + guessed = mimetypes.guess_extension(content_type, strict=False) + if guessed: + return guessed.lstrip(".") + return "bin" + + +def _is_null(value: object) -> bool: + """Return True for Python/pandas null-ish scalars.""" + if value is None: + return True + if isinstance(value, float) and pd.isna(value): + return True + try: + # Handles pd.NA, pd.NaT, and ArrowDtype scalars + return pd.isna(value) + except (TypeError, ValueError): + return False + + +def _series_to_nullable_list(series: pd.Series) -> list: + """Return a Python list from *series*, replacing any null-like value with ``None``.""" + na_mask = series.isna() + return [None if is_na else v for is_na, v in zip(na_mask, series, strict=True)] + + +def _per_modality_passthrough( + passthrough_cols: list[str], + meta_keys: set[str], + image_rows: pd.DataFrame, + text_rows: pd.DataFrame, + content_rows: pd.DataFrame, +) -> dict[str, Any]: + """Build per-modality passthrough lists from content rows. + + For each passthrough column not already in *meta_keys*: + + * **Pure per-image** (non-null only in image rows): emitted as a list with + one entry per image row (None preserved for sparse nulls). + * **Pure per-text** (non-null only in text rows): emitted as a list with + one entry per text row (None preserved for sparse nulls). + * **Mixed** (non-null in both image and text rows): emitted as a + position-aligned list with one entry per content row in position order + (None where the value is absent). When read back without declaring the + field in ``per_image_fields`` / ``per_text_fields``, the reader treats + the entire list as a sample-level passthrough on the metadata row. + """ + result: dict[str, Any] = {} + for col in passthrough_cols: + if col in meta_keys: + continue + image_has_values = not image_rows.empty and not image_rows[col].isna().all() + text_has_values = not text_rows.empty and not text_rows[col].isna().all() + + if image_has_values and not text_has_values: + result[col] = _series_to_nullable_list(image_rows[col]) + elif text_has_values and not image_has_values: + result[col] = _series_to_nullable_list(text_rows[col]) + elif image_has_values and text_has_values: + result[col] = _series_to_nullable_list(content_rows[col]) + return result + + +def _write_sample( + tf: tarfile.TarFile, + sample_df: pd.DataFrame, + sample_id: str, + passthrough_cols: list[str], +) -> None: + """Write one sample (JSON + image binaries) into *tf*.""" + escaped = _escape_key(sample_id) + content_rows = sample_df[sample_df["position"] >= 0].sort_values("position") + max_pos = int(content_rows["position"].max()) if not content_rows.empty else -1 + texts: list[str | None] = [None] * (max_pos + 1) + images: list[str | None] = [None] * (max_pos + 1) + + for row in content_rows.itertuples(index=False): + pos = int(row.position) + if row.modality == "text": + texts[pos] = row.text_content + elif row.modality == "image": + ext = _ext_from_content_type(row.content_type) + member_name = f"{escaped}.{pos}.{ext}" + images[pos] = member_name + raw = row.binary_content + if not _is_null(raw): + img_bytes = bytes(raw) + info = tarfile.TarInfo(name=member_name) + info.size = len(img_bytes) + tf.addfile(info, BytesIO(img_bytes)) + + passthrough: dict[str, Any] = {} + meta_rows = sample_df[sample_df["position"] == -1] + if not meta_rows.empty: + meta_row = meta_rows.iloc[0] + for col in passthrough_cols: + val = meta_row[col] + if not _is_null(val): + passthrough[col] = val + + image_rows = content_rows[content_rows["modality"] == "image"] + text_rows = content_rows[content_rows["modality"] == "text"] + passthrough.update( + _per_modality_passthrough(passthrough_cols, set(passthrough), image_rows, text_rows, content_rows) + ) + + payload = {"sample_id": sample_id, "texts": texts, "images": images, **passthrough} + json_bytes = json.dumps(payload, ensure_ascii=False).encode("utf-8") + info = tarfile.TarInfo(name=f"{escaped}.json") + info.size = len(json_bytes) + tf.addfile(info, BytesIO(json_bytes)) + + +# --------------------------------------------------------------------------- +# Stage +# --------------------------------------------------------------------------- + + +@dataclass +class InterleavedWebdatasetWriterStage(BaseInterleavedWriter): + """Write an ``InterleavedBatch`` as a MINT-1T-style WebDataset tar shard. + + Each sample is reconstructed from its row-based representation: + + * ``metadata`` rows supply passthrough fields embedded in the JSON. + * ``text`` rows are assembled into the ``"texts"`` list (``None`` at gaps). + * ``image`` rows are assembled into the ``"images"`` list and written as + individual tar members; ``binary_content`` must be populated (either by + the upstream pipeline or via ``materialize_on_write=True``). + + The JSON member key is ``urllib.parse.quote(sample_id, safe="")`` so that + roundtripping via :class:`InterleavedWebdatasetReaderStage` with + ``sample_id_field="sample_id"`` recovers the original sample_id. + + Only ``"metadata"``, ``"text"``, and ``"image"`` modalities are supported. + Any other modality raises ``ValueError`` at write time. + """ + + file_extension: str = "tar" + name: str = "interleaved_webdataset_writer" + + _SUPPORTED_MODALITIES: ClassVar[frozenset[str]] = frozenset({"metadata", "text", "image"}) + + def _write_dataframe(self, df: pd.DataFrame, file_path: str, _write_kwargs: dict[str, Any]) -> None: + unsupported = set(df["modality"].dropna().unique()) - self._SUPPORTED_MODALITIES + if unsupported: + msg = f"Unsupported modality {sorted(unsupported)!r}. Supported: {sorted(self._SUPPORTED_MODALITIES)}" + raise ValueError(msg) + + passthrough_cols = [c for c in df.columns if c not in RESERVED_COLUMNS] + + with ( + fsspec.open(file_path, "wb", **self.storage_options) as fobj, + tarfile.open(fileobj=fobj, mode="w:") as tf, + ): + sample_count = 0 + for sample_id, sample_df in df.groupby("sample_id", sort=False): + _write_sample(tf, sample_df, sample_id, passthrough_cols) + sample_count += 1 + + self._log_metric("samples_written", float(sample_count)) diff --git a/nemo_curator/stages/interleaved/utils/validation_utils.py b/nemo_curator/stages/interleaved/utils/validation_utils.py index d2ab2f2522..1ad77f052a 100644 --- a/nemo_curator/stages/interleaved/utils/validation_utils.py +++ b/nemo_curator/stages/interleaved/utils/validation_utils.py @@ -55,5 +55,5 @@ def validate_and_project_source_fields( result[field] = None else: value = sample[field] - result[field] = json.dumps(value, ensure_ascii=True) if isinstance(value, (dict, list)) else value + result[field] = json.dumps(value, ensure_ascii=False) if isinstance(value, (dict, list)) else value return result diff --git a/tests/stages/interleaved/test_base_writer.py b/tests/stages/interleaved/test_base_writer.py index a491297352..f089160792 100644 --- a/tests/stages/interleaved/test_base_writer.py +++ b/tests/stages/interleaved/test_base_writer.py @@ -220,11 +220,12 @@ def test_process_no_source_files_uses_uuid(tmp_path: Path, metadata: dict) -> No def test_base_writer_inputs_and_outputs(tmp_path: Path) -> None: + """inputs() and outputs() must satisfy the ProcessingStage contract.""" writer = _writer(tmp_path) task_attrs, data_attrs = writer.inputs() assert task_attrs == ["data"] assert data_attrs == [] - task_attrs, data_attrs = writer.outputs() - assert task_attrs == ["data"] - assert data_attrs == [] + out_task_attrs, out_data_attrs = writer.outputs() + assert out_task_attrs == ["data"] + assert out_data_attrs == [] diff --git a/tests/stages/interleaved/test_materialization.py b/tests/stages/interleaved/test_materialization.py index ad3abe8538..f2944fb1e2 100644 --- a/tests/stages/interleaved/test_materialization.py +++ b/tests/stages/interleaved/test_materialization.py @@ -177,6 +177,18 @@ def test_fill_tar_extract_rows_errors( # --- _scatter_range_blobs --- +def _make_range_setup( + filename: str, offset: int, size: int, frame_index: int | None = None +) -> tuple[ + list[tuple[str, int, int]], + dict[tuple[str, int, int], list[tuple[int, str, int | None]]], + list[object], + list[str | None], +]: + key = (filename, offset, size) + return [key], {key: [(0, filename, frame_index)]}, [None], [None] + + @pytest.mark.parametrize( ("blob", "expected_error_substr"), [ @@ -186,24 +198,14 @@ def test_fill_tar_extract_rows_errors( ], ) def test_scatter_range_blobs_error_cases(blob: object, expected_error_substr: str) -> None: - range_keys = [(0, 10)] - unique_ranges: dict[tuple[int, int], list[tuple[int, str, int | None]]] = { - (0, 10): [(0, "img.jpg", None)], - } - binary_values: list[object] = [None] - error_values: list[str | None] = [None] + range_keys, unique_ranges, binary_values, error_values = _make_range_setup("img.jpg", 0, 10) _scatter_range_blobs([blob], range_keys, unique_ranges, binary_values, error_values) assert error_values[0] is not None assert expected_error_substr in error_values[0] def test_scatter_range_blobs_bytearray_conversion() -> None: - range_keys = [(0, 10)] - unique_ranges: dict[tuple[int, int], list[tuple[int, str, int | None]]] = { - (0, 10): [(0, "img.jpg", None)], - } - binary_values: list[object] = [None] - error_values: list[str | None] = [None] + range_keys, unique_ranges, binary_values, error_values = _make_range_setup("img.jpg", 0, 10) _scatter_range_blobs([bytearray(b"image-data")], range_keys, unique_ranges, binary_values, error_values) assert binary_values[0] == b"image-data" assert isinstance(binary_values[0], bytes) @@ -224,12 +226,9 @@ def test_scatter_range_blobs_tiff_frame( expected_error_substr: str | None, ) -> None: tiff_bytes = build_multi_frame_tiff(n_frames) - range_keys = [(0, len(tiff_bytes))] - unique_ranges: dict[tuple[int, int], list[tuple[int, str, int | None]]] = { - (0, len(tiff_bytes)): [(0, "doc.tiff", frame_index)], - } - binary_values: list[object] = [None] - error_values: list[str | None] = [None] + range_keys, unique_ranges, binary_values, error_values = _make_range_setup( + "doc.tiff", 0, len(tiff_bytes), frame_index + ) _scatter_range_blobs([tiff_bytes], range_keys, unique_ranges, binary_values, error_values) if expect_success: assert binary_values[0] is not None diff --git a/tests/stages/interleaved/test_multimodal_core.py b/tests/stages/interleaved/test_multimodal_core.py index cd798e893a..b1d13adf58 100644 --- a/tests/stages/interleaved/test_multimodal_core.py +++ b/tests/stages/interleaved/test_multimodal_core.py @@ -22,7 +22,7 @@ import pytest from nemo_curator.core.utils import split_table_by_group_max_bytes -from nemo_curator.stages.interleaved.io.reader import WebdatasetReader +from nemo_curator.stages.interleaved.io.reader import InterleavedWebdatasetReader from nemo_curator.stages.interleaved.stages import ( BaseInterleavedAnnotatorStage, BaseInterleavedFilterStage, @@ -732,7 +732,7 @@ def test_count_with_pandas_data() -> None: def test_webdataset_reader_composite_decompose(tmp_path: Path) -> None: - reader = WebdatasetReader(file_paths=str(tmp_path)) + reader = InterleavedWebdatasetReader(file_paths=str(tmp_path)) stages = reader.decompose() assert len(stages) == 2 assert stages[0].name == "file_partitioning" @@ -1130,6 +1130,7 @@ def content_keep_mask(self, task: InterleavedBatch, df: pd.DataFrame) -> pd.Seri def _materialized_bytes(binary_content: object) -> list[tuple[int, bytes | None]]: + """Run iter_materialized_bytes with a fake materialization returning binary_content for the image row.""" task = make_image_task([make_image_row(path=None)]) df = task.to_pandas() if binary_content is _DROP_COL: diff --git a/tests/stages/interleaved/test_multimodal_reader.py b/tests/stages/interleaved/test_multimodal_reader.py index 7afb8fbe9d..416698bfe0 100644 --- a/tests/stages/interleaved/test_multimodal_reader.py +++ b/tests/stages/interleaved/test_multimodal_reader.py @@ -13,18 +13,32 @@ # limitations under the License. import json +import logging import tarfile from io import BytesIO from pathlib import Path import pandas as pd +import pyarrow as pa +import pyarrow.parquet as pq import pytest from PIL import Image +from nemo_curator.stages.file_partitioning import FilePartitioningStage +from nemo_curator.stages.interleaved.io.reader import InterleavedParquetReader +from nemo_curator.stages.interleaved.io.readers.base import BaseInterleavedReader +from nemo_curator.stages.interleaved.io.readers.parquet import InterleavedParquetReaderStage from nemo_curator.stages.interleaved.io.readers.webdataset import InterleavedWebdatasetReaderStage +from nemo_curator.stages.interleaved.io.writers.tabular import InterleavedParquetWriterStage +from nemo_curator.stages.interleaved.io.writers.webdataset import InterleavedWebdatasetWriterStage from nemo_curator.tasks import FileGroupTask, InterleavedBatch +from nemo_curator.tasks.interleaved import INTERLEAVED_SCHEMA -from .conftest import build_multi_frame_tiff, task_for_tar, write_tar +from .conftest import build_multi_frame_tiff, make_interleaved_batch, make_row, task_for_tar, write_tar + +# --------------------------------------------------------------------------- +# Shared helpers +# --------------------------------------------------------------------------- def _as_df(task_or_tasks: InterleavedBatch | list[InterleavedBatch]) -> pd.DataFrame: @@ -40,23 +54,76 @@ def _write_tar_sample( image_name: str = "image.jpg", image_bytes: bytes = b"abc", ) -> None: - with tarfile.open(tar_path, "w") as tf: - json_blob = json.dumps(payload).encode("utf-8") - json_info = tarfile.TarInfo(name=json_name) - json_info.size = len(json_blob) - tf.addfile(json_info, BytesIO(json_blob)) - img_info = tarfile.TarInfo(name=image_name) - img_info.size = len(image_bytes) - tf.addfile(img_info, BytesIO(image_bytes)) + write_tar(tar_path, {json_name: json.dumps(payload).encode("utf-8"), image_name: image_bytes}) -def _task_for_tar(tar_path: Path, task_id: str) -> FileGroupTask: - return FileGroupTask( - task_id=task_id, - dataset_name="custom_dataset", - data=[str(tar_path)], - _metadata={"source_files": [str(tar_path)]}, - ) +def _write_parquet_task(batch: InterleavedBatch, out_dir: Path) -> str: + """Write *batch* to parquet and return the written file path.""" + writer = InterleavedParquetWriterStage(path=str(out_dir), materialize_on_write=False, mode="overwrite") + write_task = writer.process(batch) + return write_task.data[0] + + +def _make_aligned_rows(fake_jpg: bytes, num_images: int = 2) -> list[dict]: + """Build a standard s1 sample with sample_metadata/text_metadata/image_metadata on the correct rows. + + sample_metadata lives on the metadata row, text_metadata on the text row, image_metadata on each + image row — three distinct classes used to verify per-row field alignment in round-trip tests. + """ + rows = [ + make_row("s1", -1, "metadata", sample_metadata="doc A", text_metadata=None, image_metadata=None), + make_row( + "s1", + 0, + "text", + text_content="hello world", + sample_metadata=None, + text_metadata="conf:0.9", + image_metadata=None, + ), + ] + for i in range(num_images): + rows.append( + make_row( + "s1", + i + 1, + "image", + content_type="image/jpeg", + binary_content=fake_jpg, + sample_metadata=None, + text_metadata=None, + image_metadata=f'{{"page": {i}}}', + ) + ) + return rows + + +def _assert_field_alignment(df: pd.DataFrame, *, image_vals: list, text_vals: list, sample_val: object) -> None: + """Verify sample/text/image extra fields land on the correct rows and are null everywhere else.""" + meta = df[df["modality"] == "metadata"].sort_values("position") + text = df[df["modality"] == "text"].sort_values("position") + imgs = df[df["modality"] == "image"].sort_values("position") + + assert len(meta) == 1, f"expected 1 metadata row, got {len(meta)}" + assert len(text) == len(text_vals), f"expected {len(text_vals)} text rows, got {len(text)}" + assert len(imgs) == len(image_vals), f"expected {len(image_vals)} image rows, got {len(imgs)}" + + assert meta["sample_metadata"].iloc[0] == sample_val + assert text["sample_metadata"].isna().all() + assert imgs["sample_metadata"].isna().all() + + assert text["text_metadata"].tolist() == text_vals + assert meta["text_metadata"].isna().all() + assert imgs["text_metadata"].isna().all() + + assert imgs["image_metadata"].tolist() == image_vals + assert meta["image_metadata"].isna().all() + assert text["image_metadata"].isna().all() + + +# --------------------------------------------------------------------------- +# InterleavedWebdatasetReaderStage +# --------------------------------------------------------------------------- def test_reader_supports_custom_field_mapping(tmp_path: Path) -> None: @@ -77,7 +144,7 @@ def test_reader_supports_custom_field_mapping(tmp_path: Path) -> None: image_name="custom-image.jpg", image_bytes=image_bytes, ) - task = _task_for_tar(tar_path, "file_group_custom") + task = task_for_tar(str(tar_path), "file_group_custom") reader = InterleavedWebdatasetReaderStage( sample_id_field="doc_id", texts_field="captions", @@ -113,7 +180,7 @@ def test_reader_reads_all_fields_by_default(tmp_path: Path) -> None: "aux": {"page": 3}, } _write_tar_sample(tar_path, payload, json_name="sample.meta.json") - task = _task_for_tar(tar_path, "all_fields") + task = task_for_tar(str(tar_path), "all_fields") reader = InterleavedWebdatasetReaderStage( sample_id_field="doc_id", texts_field="captions", @@ -125,7 +192,7 @@ def test_reader_reads_all_fields_by_default(tmp_path: Path) -> None: meta_row = df[df["modality"] == "metadata"].iloc[0] assert meta_row["p_hash"] == "phash-1" assert meta_row["score"] == 0.91 - assert meta_row["aux"] == json.dumps({"page": 3}, ensure_ascii=True) + assert meta_row["aux"] == json.dumps({"page": 3}, ensure_ascii=False) image_row = df[df["modality"] == "image"].iloc[0] assert pd.isna(image_row["p_hash"]) assert "captions" not in df.columns @@ -153,7 +220,7 @@ def test_reader_uses_resolved_content_key_for_content_type(tmp_path: Path) -> No jpg_info.size = 3 tf.addfile(jpg_info, BytesIO(b"jpg")) - task = _task_for_tar(tar_path, "content_type_resolve") + task = task_for_tar(str(tar_path), "content_type_resolve") reader = InterleavedWebdatasetReaderStage( sample_id_field="doc_id", texts_field="captions", @@ -175,7 +242,7 @@ def test_reader_image_tokens_with_frame_index(tmp_path: Path) -> None: "images": [None, "page_0_image_15", "page_1_image_22"], } _write_tar_sample(tar_path, payload, json_name="sample.json", image_name="doc.pdf.tiff", image_bytes=b"TIFF_DATA") - task = _task_for_tar(tar_path, "sub_image_test") + task = task_for_tar(str(tar_path), "sub_image_test") reader = InterleavedWebdatasetReaderStage( sample_id_field="pdf_name", image_extensions=(".tiff",), @@ -210,7 +277,7 @@ def test_reader_interleaved_positions_do_not_overlap(tmp_path: Path) -> None: "images": [None, "page_img", None, "chart_img", None], } _write_tar_sample(tar_path, payload, image_name="interleaved.pdf.jpg", image_bytes=b"\xff\xd8\xff") - task = _task_for_tar(tar_path, "interleaved_test") + task = task_for_tar(str(tar_path), "interleaved_test") reader = InterleavedWebdatasetReaderStage(sample_id_field="pdf_name") df = _as_df(reader.process(task)) @@ -237,7 +304,7 @@ def test_reader_empty_output_schema_includes_requested_passthrough_fields(tmp_pa img_info.size = 3 tf.addfile(img_info, BytesIO(b"abc")) - task = _task_for_tar(tar_path, "empty_schema") + task = task_for_tar(str(tar_path), "empty_schema") reader = InterleavedWebdatasetReaderStage(fields=("p_hash",)) df = _as_df(reader.process(task)) assert "p_hash" in df.columns @@ -247,7 +314,7 @@ def test_reader_fields_reserved_key_raises(tmp_path: Path) -> None: tar_path = tmp_path / "reserved_key.tar" payload = {"pdf_name": "doc.pdf", "texts": ["t"], "images": []} _write_tar_sample(tar_path, payload) - task = _task_for_tar(tar_path, "reserved_key") + task = task_for_tar(str(tar_path), "reserved_key") reader = InterleavedWebdatasetReaderStage(fields=("sample_id",)) with pytest.raises(ValueError, match="fields contains reserved keys"): _ = reader.process(task) @@ -257,7 +324,7 @@ def test_reader_fields_missing_key_warns_and_fills_none(tmp_path: Path, caplog: tar_path = tmp_path / "missing_key.tar" payload = {"pdf_name": "doc.pdf", "texts": ["t"], "images": []} _write_tar_sample(tar_path, payload) - task = _task_for_tar(tar_path, "missing_key") + task = task_for_tar(str(tar_path), "missing_key") reader = InterleavedWebdatasetReaderStage(fields=("p_hash",)) with caplog.at_level("WARNING"): result = reader.process(task) @@ -278,7 +345,7 @@ def test_reader_per_image_fields_distributed_to_image_rows(tmp_path: Path) -> No "image_metadata": [{"height": 100, "width": 200}], } _write_tar_sample(tar_path, payload) - task = _task_for_tar(tar_path, "per_image") + task = task_for_tar(str(tar_path), "per_image") reader = InterleavedWebdatasetReaderStage( per_image_fields=("image_metadata",), ) @@ -288,7 +355,7 @@ def test_reader_per_image_fields_distributed_to_image_rows(tmp_path: Path) -> No image_rows = df[df["modality"] == "image"] assert len(image_rows) == 1 - assert image_rows.iloc[0]["image_metadata"] == json.dumps({"height": 100, "width": 200}) + assert image_rows.iloc[0]["image_metadata"] == json.dumps({"height": 100, "width": 200}, ensure_ascii=False) text_rows = df[df["modality"] == "text"] assert all(pd.isna(v) for v in text_rows["image_metadata"]) @@ -307,7 +374,7 @@ def test_reader_per_text_fields_distributed_to_text_rows(tmp_path: Path) -> None "text_scores": [0.95, 0.42], } _write_tar_sample(tar_path, payload) - task = _task_for_tar(tar_path, "per_text") + task = task_for_tar(str(tar_path), "per_text") reader = InterleavedWebdatasetReaderStage( per_text_fields=("text_scores",), ) @@ -339,7 +406,7 @@ def test_reader_per_image_and_per_text_fields_together(tmp_path: Path) -> None: "url": "https://example.com", } _write_tar_sample(tar_path, payload) - task = _task_for_tar(tar_path, "both_per_modality") + task = task_for_tar(str(tar_path), "both_per_modality") reader = InterleavedWebdatasetReaderStage( per_image_fields=("image_metadata",), per_text_fields=("text_lang",), @@ -347,7 +414,7 @@ def test_reader_per_image_and_per_text_fields_together(tmp_path: Path) -> None: df = _as_df(reader.process(task)) image_rows = df[df["modality"] == "image"] - assert image_rows.iloc[0]["image_metadata"] == json.dumps({"page": 1, "width": 640}) + assert image_rows.iloc[0]["image_metadata"] == json.dumps({"page": 1, "width": 640}, ensure_ascii=False) assert pd.isna(image_rows.iloc[0]["text_lang"]) text_rows = df[df["modality"] == "text"].sort_values("position") @@ -373,7 +440,7 @@ def test_reader_per_modality_fields_excluded_from_sample_passthrough(tmp_path: P "url": "https://example.com", } _write_tar_sample(tar_path, payload) - task = _task_for_tar(tar_path, "exclude_pt") + task = task_for_tar(str(tar_path), "exclude_pt") reader = InterleavedWebdatasetReaderStage( per_image_fields=("image_metadata",), per_text_fields=("text_scores",), @@ -395,7 +462,7 @@ def test_reader_per_modality_field_missing_warns(tmp_path: Path, caplog: pytest. "images": [], } _write_tar_sample(tar_path, payload) - task = _task_for_tar(tar_path, "missing_per_field") + task = task_for_tar(str(tar_path), "missing_per_field") reader = InterleavedWebdatasetReaderStage( per_image_fields=("image_metadata",), ) @@ -416,7 +483,7 @@ def test_reader_raises_on_non_list_per_modality_field(tmp_path: Path) -> None: "image_metadata": "not-a-list", } _write_tar_sample(tar_path, payload) - task = _task_for_tar(tar_path, "non_list_field") + task = task_for_tar(str(tar_path), "non_list_field") reader = InterleavedWebdatasetReaderStage( per_image_fields=("image_metadata",), ) @@ -747,9 +814,11 @@ def test_reader_source_files_per_split_only_contributing_tars(tmp_path: Path) -> src = batch._metadata["source_files"] assert len(src) == 1, f"expected 1 source file for split, got {src}" if "doc1" in sample_ids: - assert tar1 + "::split_" in src[0], f"doc1 split should point to {tar1}, got {src}" + assert src[0].startswith(tar1), f"doc1 split should point to {tar1}, got {src}" + assert "::split_" in src[0], f"doc1 split should have ::split_ suffix, got {src}" elif "doc2" in sample_ids: - assert tar2 + "::split_" in src[0], f"doc2 split should point to {tar2}, got {src}" + assert src[0].startswith(tar2), f"doc2 split should point to {tar2}, got {src}" + assert "::split_" in src[0], f"doc2 split should have ::split_ suffix, got {src}" else: pytest.fail(f"unexpected sample_ids in split: {sample_ids}") @@ -827,3 +896,304 @@ def test_reader_unknown_fields_pass_through_by_default(tmp_path: Path) -> None: meta = df[df["modality"] == "metadata"].iloc[0] assert meta["pdf_name"] == "doc.pdf" assert meta["url"] == "https://example.com" + + +# --------------------------------------------------------------------------- +# InterleavedParquetReaderStage +# --------------------------------------------------------------------------- + + +def test_parquet_reader_roundtrip(tmp_path: Path) -> None: + """Write a batch with the parquet writer, read it back; data matches.""" + batch = make_interleaved_batch(num_samples=2, include_images=False) + pq_path = _write_parquet_task(batch, tmp_path / "out") + + task = FileGroupTask(task_id="pq_rt", dataset_name="d", data=[pq_path]) + reader = InterleavedParquetReaderStage() + result = reader.process(task) + assert isinstance(result, InterleavedBatch) + df = result.to_pandas() + + assert set(df["sample_id"].tolist()) == {"sample_0", "sample_1"} + text_rows = df[df["modality"] == "text"] + assert set(text_rows["text_content"].tolist()) == {"Hello 0", "Hello 1"} + assert result._metadata.get("source_files") == [pq_path] + + +def test_parquet_reader_missing_columns_filled_with_null(tmp_path: Path) -> None: + """A parquet file with only 3 columns; all other schema cols become null.""" + minimal = pa.Table.from_pylist( + [{"sample_id": "s1", "position": 0, "modality": "text"}], + schema=pa.schema( + [ + pa.field("sample_id", pa.string()), + pa.field("position", pa.int32()), + pa.field("modality", pa.string()), + ] + ), + ) + pq_path = tmp_path / "minimal.parquet" + pq.write_table(minimal, pq_path) + + task = FileGroupTask(task_id="minimal", dataset_name="d", data=[str(pq_path)]) + result = InterleavedParquetReaderStage().process(task) + assert isinstance(result, InterleavedBatch) + df = result.to_pandas() + assert len(df) == 1 + assert pd.isna(df.loc[0, "text_content"]) + assert pd.isna(df.loc[0, "binary_content"]) + + +def test_parquet_reader_fields_subset(tmp_path: Path) -> None: + """fields=(...) reads only reserved cols + requested extras; others absent.""" + batch = make_interleaved_batch(num_samples=1, include_images=False) + pq_path = _write_parquet_task(batch, tmp_path / "out") + + task = FileGroupTask(task_id="fields_sub", dataset_name="d", data=[pq_path]) + result = InterleavedParquetReaderStage(fields=("text_content",)).process(task) + assert isinstance(result, InterleavedBatch) + df = result.to_pandas() + assert "text_content" in df.columns + assert "binary_content" in df.columns # reserved — always present + + +def test_parquet_reader_fields_null_fill_missing(tmp_path: Path) -> None: + """A field in fields= that is absent from disk is null-filled, not errored.""" + batch = make_interleaved_batch(num_samples=1, include_images=False) + pq_path = _write_parquet_task(batch, tmp_path / "out") + + task = FileGroupTask(task_id="null_fill", dataset_name="d", data=[pq_path]) + result = InterleavedParquetReaderStage(fields=("nonexistent_field",)).process(task) + assert isinstance(result, InterleavedBatch) + df = result.to_pandas() + assert "nonexistent_field" in df.columns + assert df["nonexistent_field"].isna().all() + + +def test_parquet_reader_extra_column_passthrough_by_default(tmp_path: Path) -> None: + """fields=None reads ALL columns in the file — sample/text/image extra fields are preserved + and land on the correct rows after read.""" + fake_jpg = b"\xff\xd8\xff\xe0fake" + pq_path = str(tmp_path / "extra.parquet") + pq.write_table(pa.Table.from_pylist(_make_aligned_rows(fake_jpg)), pq_path) + + result = InterleavedParquetReaderStage().process( + FileGroupTask(task_id="extra_col", dataset_name="d", data=[pq_path]) + ) + assert isinstance(result, InterleavedBatch) + _assert_field_alignment( + result.to_pandas(), + image_vals=['{"page": 0}', '{"page": 1}'], + text_vals=["conf:0.9"], + sample_val="doc A", + ) + + +def test_parquet_reader_extra_column_excluded_when_fields_set(tmp_path: Path) -> None: + """When fields= is explicit, only listed extras are read; unlisted ones are dropped.""" + fake_jpg = b"\xff\xd8\xff\xe0fake" + pq_path = str(tmp_path / "extra.parquet") + pq.write_table(pa.Table.from_pylist(_make_aligned_rows(fake_jpg, num_images=1)), pq_path) + + # fields=("text_metadata",) → only text_metadata + reserved cols; others are NOT read + result = InterleavedParquetReaderStage(fields=("text_metadata",)).process( + FileGroupTask(task_id="fields_excl", dataset_name="d", data=[pq_path]) + ) + assert isinstance(result, InterleavedBatch) + df = result.to_pandas() + assert "text_metadata" in df.columns + assert "image_metadata" not in df.columns + assert "sample_metadata" not in df.columns + + +def test_parquet_reader_wds_to_pq_to_wds_roundtrip(tmp_path: Path) -> None: + """WDS→PQ→WDS: sample/text/image extra fields survive the full round-trip through parquet. + + Primary regression test for the image_metadata loss bug: InterleavedParquetReaderStage + with fields=None must read ALL columns, not just RESERVED_COLUMNS. + """ + fake_jpg = b"\xff\xd8\xff\xe0fake" + payload = { + "sample_id": "s1", + "sample_metadata": "doc A", + "texts": ["hello world", None, None], + "text_metadata": ["conf:0.9"], + "images": [None, "s1.1.jpg", "s1.2.jpg"], + "image_metadata": [{"page": 1}, {"page": 2}], + } + tar_path = write_tar( + tmp_path / "input.tar", + {"s1.json": json.dumps(payload).encode(), "s1.1.jpg": fake_jpg, "s1.2.jpg": fake_jpg}, + ) + + wds_reader = InterleavedWebdatasetReaderStage( + sample_id_field="sample_id", + per_image_fields=("image_metadata",), + per_text_fields=("text_metadata",), + ) + + batch_wds = wds_reader.process(task_for_tar(tar_path)) + assert isinstance(batch_wds, InterleavedBatch) + _assert_field_alignment( + batch_wds.to_pandas(), + image_vals=['{"page": 1}', '{"page": 2}'], + text_vals=["conf:0.9"], + sample_val="doc A", + ) + + pq_task = InterleavedParquetWriterStage( + path=str(tmp_path / "pq_out"), materialize_on_write=False, mode="overwrite" + ).process(batch_wds) + + batch_pq = InterleavedParquetReaderStage().process( + FileGroupTask(task_id="pq_rt", dataset_name="d", data=[pq_task.data[0]]) + ) + assert isinstance(batch_pq, InterleavedBatch) + df_pq = batch_pq.to_pandas() + for col in ("sample_metadata", "text_metadata", "image_metadata"): + assert col in df_pq.columns, f"{col} lost in PQ read — bug regression" + _assert_field_alignment( + df_pq, image_vals=['{"page": 1}', '{"page": 2}'], text_vals=["conf:0.9"], sample_val="doc A" + ) + + wds2_task = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "wds_out"), materialize_on_write=False, mode="overwrite" + ).process(batch_pq) + batch_final = wds_reader.process(FileGroupTask(task_id="final", dataset_name="d", data=wds2_task.data)) + assert isinstance(batch_final, InterleavedBatch) + _assert_field_alignment( + batch_final.to_pandas(), + image_vals=['{"page": 1}', '{"page": 2}'], + text_vals=["conf:0.9"], + sample_val="doc A", + ) + + +def test_parquet_reader_pq_to_wds_to_pq_roundtrip(tmp_path: Path) -> None: + """PQ→WDS→PQ: sample/text/image extra columns survive a full write-to-WDS and back.""" + fake_jpg = b"\xff\xd8\xff\xe0fake" + batch0 = InterleavedBatch( + task_id="pq0", + dataset_name="d", + data=pa.Table.from_pylist(_make_aligned_rows(fake_jpg)), + _metadata={"source_files": ["source.parquet"]}, + ) + + pq_path0 = ( + InterleavedParquetWriterStage(path=str(tmp_path / "pq0_out"), materialize_on_write=False, mode="overwrite") + .process(batch0) + .data[0] + ) + + batch1 = InterleavedParquetReaderStage().process(FileGroupTask(task_id="pq1", dataset_name="d", data=[pq_path0])) + assert isinstance(batch1, InterleavedBatch) + _assert_field_alignment( + batch1.to_pandas(), + image_vals=['{"page": 0}', '{"page": 1}'], + text_vals=["conf:0.9"], + sample_val="doc A", + ) + + wds_task = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "wds1_out"), materialize_on_write=False, mode="overwrite" + ).process(batch1) + + wds_reader = InterleavedWebdatasetReaderStage( + sample_id_field="sample_id", + per_image_fields=("image_metadata",), + per_text_fields=("text_metadata",), + ) + batch2 = wds_reader.process(FileGroupTask(task_id="wds1", dataset_name="d", data=wds_task.data)) + assert isinstance(batch2, InterleavedBatch) + _assert_field_alignment( + batch2.to_pandas(), + image_vals=['{"page": 0}', '{"page": 1}'], + text_vals=["conf:0.9"], + sample_val="doc A", + ) + + pq_path1 = ( + InterleavedParquetWriterStage(path=str(tmp_path / "pq1_out"), materialize_on_write=False, mode="overwrite") + .process(batch2) + .data[0] + ) + batch_final = InterleavedParquetReaderStage().process( + FileGroupTask(task_id="pq_final", dataset_name="d", data=[pq_path1]) + ) + assert isinstance(batch_final, InterleavedBatch) + _assert_field_alignment( + batch_final.to_pandas(), + image_vals=['{"page": 0}', '{"page": 1}'], + text_vals=["conf:0.9"], + sample_val="doc A", + ) + + +def test_parquet_reader_max_batch_bytes_splits(tmp_path: Path) -> None: + """Two parquet files, one sample each; max_batch_bytes=1 → 2 splits, + each split's source_files lists only its contributing file.""" + batch_a = make_interleaved_batch(num_samples=1, task_id="a", include_images=False) + batch_b = make_interleaved_batch(num_samples=1, task_id="b", include_images=False) + # Give distinct sample_ids + rows_a = batch_a.to_pandas().copy() + rows_a["sample_id"] = "doc_a" + rows_b = batch_b.to_pandas().copy() + rows_b["sample_id"] = "doc_b" + + out_a = tmp_path / "a" + out_b = tmp_path / "b" + writer = InterleavedParquetWriterStage(path=str(out_a), materialize_on_write=False, mode="overwrite") + pq_a = writer.process(InterleavedBatch(task_id="a", dataset_name="d", data=rows_a)).data[0] + writer2 = InterleavedParquetWriterStage(path=str(out_b), materialize_on_write=False, mode="overwrite") + pq_b = writer2.process(InterleavedBatch(task_id="b", dataset_name="d", data=rows_b)).data[0] + + task = FileGroupTask(task_id="split_test", dataset_name="d", data=[pq_a, pq_b]) + result = InterleavedParquetReaderStage(max_batch_bytes=1).process(task) + + assert isinstance(result, list) + assert len(result) == 2 + for batch in result: + sample_ids = set(batch.to_pandas()["sample_id"].tolist()) + src = batch._metadata["source_files"] + assert len(src) == 1 + if "doc_a" in sample_ids: + assert pq_a in src[0] + elif "doc_b" in sample_ids: + assert pq_b in src[0] + + +def test_parquet_reader_empty_file(tmp_path: Path) -> None: + """An empty parquet file produces an empty InterleavedBatch with correct schema.""" + empty = pa.Table.from_pylist([], schema=INTERLEAVED_SCHEMA) + pq_path = tmp_path / "empty.parquet" + pq.write_table(empty, pq_path) + + task = FileGroupTask(task_id="empty", dataset_name="d", data=[str(pq_path)]) + result = InterleavedParquetReaderStage().process(task) + assert isinstance(result, InterleavedBatch) + assert len(result.to_pandas()) == 0 + + +def test_parquet_reader_composite_decompose(tmp_path: Path) -> None: + """InterleavedParquetReader.decompose() returns [FilePartitioningStage, InterleavedParquetReaderStage].""" + reader = InterleavedParquetReader(file_paths=str(tmp_path)) + stages = reader.decompose() + assert len(stages) == 2 + assert isinstance(stages[0], FilePartitioningStage) + assert isinstance(stages[1], InterleavedParquetReaderStage) + + +def test_parquet_reader_empty_file_list_returns_empty_batch() -> None: + result = InterleavedParquetReaderStage().process(FileGroupTask(task_id="t", dataset_name="d", data=[])) + assert isinstance(result, InterleavedBatch) + df = result.to_pandas() + assert len(df) == 0 + assert set(INTERLEAVED_SCHEMA.names) <= set(df.columns) + + +def test_source_files_for_split_null_sample_ids_fallback(caplog) -> None: # noqa: ANN001 + split = pa.table({"sample_id": pa.array([None, None], type=pa.string())}) + with caplog.at_level(logging.WARNING, logger="nemo_curator"): + result = BaseInterleavedReader._source_files_for_split(split, 2, {}, ["/a.parquet", "/b.parquet"]) + assert result == ["/a.parquet::split_00002", "/b.parquet::split_00002"] + assert any("falling back" in r.message for r in caplog.records) diff --git a/tests/stages/interleaved/test_multimodal_writer.py b/tests/stages/interleaved/test_multimodal_writer.py index 8729529348..03e1596353 100644 --- a/tests/stages/interleaved/test_multimodal_writer.py +++ b/tests/stages/interleaved/test_multimodal_writer.py @@ -24,10 +24,18 @@ from nemo_curator.stages.interleaved.io.readers.webdataset import InterleavedWebdatasetReaderStage from nemo_curator.stages.interleaved.io.writers.tabular import InterleavedParquetWriterStage +from nemo_curator.stages.interleaved.io.writers.webdataset import ( + InterleavedWebdatasetWriterStage, + _escape_key, + _ext_from_content_type, + _is_null, +) from nemo_curator.stages.interleaved.stages import BaseInterleavedFilterStage from nemo_curator.tasks import FileGroupTask, InterleavedBatch from nemo_curator.tasks.interleaved import INTERLEAVED_SCHEMA, RESERVED_COLUMNS +from .conftest import make_row + def _read_batch(input_task: FileGroupTask) -> InterleavedBatch: batch = InterleavedWebdatasetReaderStage().process(input_task) @@ -411,3 +419,457 @@ def test_writer_custom_compression(tmp_path: Path, compression: str) -> None: meta = pq.read_metadata(write_task.data[0]) actual = meta.row_group(0).column(0).compression.lower() assert actual == compression + + +# --------------------------------------------------------------------------- +# Helper utilities +# --------------------------------------------------------------------------- + + +def _make_wds_batch( + sample_id: str = "s1", + text: str = "hello", + image_bytes: bytes | None = b"fake-img", + extra_cols: dict | None = None, + source_files: list[str] | None = None, +) -> InterleavedBatch: + """Build a minimal metadata+text+image InterleavedBatch for WDS writer tests.""" + rows: list[dict] = [ + { + "sample_id": sample_id, + "position": -1, + "modality": "metadata", + "content_type": "application/json", + "text_content": None, + "binary_content": None, + "source_ref": None, + "materialize_error": None, + **(extra_cols or {}), + }, + { + "sample_id": sample_id, + "position": 0, + "modality": "text", + "content_type": "text/plain", + "text_content": text, + "binary_content": None, + "source_ref": None, + "materialize_error": None, + **(extra_cols or {}), + }, + ] + if image_bytes is not None: + rows.append( + { + "sample_id": sample_id, + "position": 1, + "modality": "image", + "content_type": "image/png", + "text_content": None, + "binary_content": image_bytes, + "source_ref": None, + "materialize_error": None, + **(extra_cols or {}), + } + ) + return InterleavedBatch( + task_id=f"wds_{sample_id}", + dataset_name="test", + data=pd.DataFrame(rows), + _metadata={"source_files": source_files or ["test.tar"]}, + ) + + +def _read_wds_tar(tar_path: str) -> tuple[dict, dict[str, bytes]]: + """Read a WDS tar: returns (json_payload, {member_name: bytes}).""" + payload: dict = {} + members: dict[str, bytes] = {} + with tarfile.open(tar_path, "r") as tf: + for m in tf.getmembers(): + f = tf.extractfile(m) + if f is None: + continue + data = f.read() + members[m.name] = data + if m.name.endswith(".json"): + payload = json.loads(data) + return payload, members + + +# --------------------------------------------------------------------------- +# Unit tests for helper functions +# --------------------------------------------------------------------------- + + +def test_escape_key_encodes_special_chars() -> None: + assert "/" not in _escape_key("a/b") + assert ":" not in _escape_key("a:b") + assert _escape_key("simple") == "simple" + assert _escape_key("a/b:c") == "a%2Fb%3Ac" + + +def test_ext_from_content_type_known() -> None: + assert _ext_from_content_type("image/jpeg") == "jpg" + assert _ext_from_content_type("image/png") == "png" + assert _ext_from_content_type("image/tiff") == "tiff" + + +def test_ext_from_content_type_fallback() -> None: + assert _ext_from_content_type(None) == "bin" + assert _ext_from_content_type("application/octet-stream") == "bin" + + +# --------------------------------------------------------------------------- +# InterleavedWebdatasetWriterStage tests +# --------------------------------------------------------------------------- + + +def test_wds_writer_roundtrip(tmp_path: Path) -> None: + """Write a WDS batch, read it back; text and image content are preserved.""" + batch = _make_wds_batch(sample_id="doc1", text="roundtrip text", image_bytes=b"img-data") + writer = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "wds_out"), + materialize_on_write=False, + mode="overwrite", + ) + write_task = writer.process(batch) + tar_path = write_task.data[0] + assert tar_path.endswith(".tar") + + read_task = FileGroupTask(task_id="rt", dataset_name="test", data=[tar_path]) + result = InterleavedWebdatasetReaderStage(sample_id_field="sample_id").process(read_task) + assert isinstance(result, InterleavedBatch) + df = result.to_pandas() + + assert "doc1" in df["sample_id"].tolist() + text_rows = df[df["modality"] == "text"] + assert text_rows["text_content"].tolist() == ["roundtrip text"] + image_rows = df[df["modality"] == "image"] + assert len(image_rows) == 1 + + +def test_wds_writer_text_only_sample(tmp_path: Path) -> None: + """No image rows → tar has JSON with all-None images list.""" + batch = _make_wds_batch(sample_id="text_only", image_bytes=None) + writer = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "text_only_out"), + materialize_on_write=False, + mode="overwrite", + ) + write_task = writer.process(batch) + payload, members = _read_wds_tar(write_task.data[0]) + + assert all(v is None for v in payload["images"]) + assert not any(k.endswith((".png", ".jpg")) for k in members) + + +def test_wds_writer_unsupported_modality_raises(tmp_path: Path) -> None: + """A modality='video' row raises ValueError naming the bad modality.""" + df = pd.DataFrame( + [ + { + "sample_id": "s1", + "position": 0, + "modality": "video", + "content_type": "video/mp4", + "text_content": None, + "binary_content": None, + "source_ref": None, + "materialize_error": None, + } + ] + ) + task = InterleavedBatch(task_id="vid", dataset_name="t", data=df, _metadata={"source_files": ["x.tar"]}) + writer = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "vid_out"), materialize_on_write=False, mode="overwrite" + ) + with pytest.raises(ValueError, match="video"): + writer.process(task) + + +def test_wds_writer_key_escaping(tmp_path: Path) -> None: + """sample_id with special chars → tar members have safe names; roundtrip recovers original id.""" + sample_id = "a/b:c.d" + batch = _make_wds_batch(sample_id=sample_id, image_bytes=None) + writer = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "escape_out"), materialize_on_write=False, mode="overwrite" + ) + write_task = writer.process(batch) + _, members = _read_wds_tar(write_task.data[0]) + + for name in members: + assert "/" not in name.split(".json")[0], f"unescaped slash in member name: {name}" + assert ":" not in name.split(".json")[0], f"unescaped colon in member name: {name}" + + read_task = FileGroupTask(task_id="esc_rt", dataset_name="test", data=[write_task.data[0]]) + result = InterleavedWebdatasetReaderStage(sample_id_field="sample_id").process(read_task) + assert isinstance(result, InterleavedBatch) + assert sample_id in result.to_pandas()["sample_id"].tolist() + + +def test_wds_writer_passthrough_columns_in_json(tmp_path: Path) -> None: + """Extra 'url' col in metadata row appears in the written JSON payload.""" + rows = [ + { + "sample_id": "s1", + "position": -1, + "modality": "metadata", + "content_type": "application/json", + "text_content": None, + "binary_content": None, + "source_ref": None, + "materialize_error": None, + "url": "https://example.com", + }, + { + "sample_id": "s1", + "position": 0, + "modality": "text", + "content_type": "text/plain", + "text_content": "hi", + "binary_content": None, + "source_ref": None, + "materialize_error": None, + "url": None, + }, + ] + task = InterleavedBatch( + task_id="pt", dataset_name="t", data=pd.DataFrame(rows), _metadata={"source_files": ["x.tar"]} + ) + writer = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "pt_out"), materialize_on_write=False, mode="overwrite" + ) + write_task = writer.process(task) + payload, _ = _read_wds_tar(write_task.data[0]) + assert payload.get("url") == "https://example.com" + + +def test_wds_writer_null_binary_skips_member(tmp_path: Path) -> None: + """null binary_content with materialize_on_write=False → no image member written.""" + table = pa.Table.from_pylist( + [ + { + "sample_id": "s1", + "position": -1, + "modality": "metadata", + "content_type": "application/json", + "text_content": None, + "binary_content": None, + "source_ref": None, + "materialize_error": None, + }, + { + "sample_id": "s1", + "position": 0, + "modality": "image", + "content_type": "image/png", + "text_content": None, + "binary_content": None, + "source_ref": InterleavedBatch.build_source_ref(path="/fake/img.png", member=None), + "materialize_error": None, + }, + ], + schema=INTERLEAVED_SCHEMA, + ) + task = InterleavedBatch( + task_id="null_bin", dataset_name="t", data=table, _metadata={"source_files": ["/fake/img.png"]} + ) + writer = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "null_bin_out"), materialize_on_write=False, mode="overwrite" + ) + write_task = writer.process(task) + payload, members = _read_wds_tar(write_task.data[0]) + + assert any(k.endswith(".json") for k in members) + assert not any(k.endswith(".png") for k in members) + assert any(v is not None for v in payload.get("images", [])) + + +def test_wds_writer_deterministic_filename(tmp_path: Path) -> None: + """Same source_files + task_id → same output filename across two writer instances.""" + source = ["shard-00000.tar"] + task_id = "fixed_task" + batch = InterleavedBatch( + task_id=task_id, + dataset_name="test", + data=_make_wds_batch(sample_id="s1", source_files=source).to_pandas(), + _metadata={"source_files": source}, + ) + + out1 = tmp_path / "det_out1" + out2 = tmp_path / "det_out2" + t1 = InterleavedWebdatasetWriterStage(path=str(out1), materialize_on_write=False, mode="overwrite").process(batch) + t2 = InterleavedWebdatasetWriterStage(path=str(out2), materialize_on_write=False, mode="overwrite").process(batch) + assert Path(t1.data[0]).name == Path(t2.data[0]).name + + +def test_wds_writer_per_image_fields_roundtrip(tmp_path: Path) -> None: + """per_image_fields are written as a list in JSON and survive a write→read round-trip.""" + fake_png = b"\x89PNG\r\n" + rows = [ + # image_alt_text is null on the metadata row — it's a per-image field, not sample-level + make_row("s1", -1, "metadata", url="http://example.com", image_alt_text=None), + make_row("s1", 0, "text", text_content="intro text", url=None, image_alt_text=None), + make_row( + "s1", + 1, + "image", + content_type="image/png", + binary_content=fake_png, + url=None, + image_alt_text="a cat on a chair", + ), + make_row("s1", 2, "image", content_type="image/png", binary_content=fake_png, url=None, image_alt_text=None), + make_row( + "s1", + 3, + "image", + content_type="image/png", + binary_content=fake_png, + url=None, + image_alt_text="a dog in a park", + ), + ] + task = InterleavedBatch( + task_id="per_img", dataset_name="t", data=pd.DataFrame(rows), _metadata={"source_files": ["x.tar"]} + ) + writer = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "per_img_out"), materialize_on_write=False, mode="overwrite" + ) + write_task = writer.process(task) + + payload, _ = _read_wds_tar(write_task.data[0]) + assert "image_alt_text" in payload, "image_alt_text missing from JSON payload" + assert payload["image_alt_text"] == ["a cat on a chair", None, "a dog in a park"] + # sample-level passthrough still present + assert payload["url"] == "http://example.com" + + reader = InterleavedWebdatasetReaderStage( + sample_id_field="sample_id", + per_image_fields=("image_alt_text",), + ) + result = reader.process(FileGroupTask(task_id="rt", dataset_name="t", data=write_task.data)) + assert isinstance(result, InterleavedBatch) + df = result.to_pandas() + + image_rows = df[df["modality"] == "image"].sort_values("position") + alt_texts = [None if pd.isna(v) else v for v in image_rows["image_alt_text"]] + assert alt_texts == ["a cat on a chair", None, "a dog in a park"] + + +def test_wds_writer_per_text_fields_roundtrip(tmp_path: Path) -> None: + """per_text_fields are written as a list in JSON and survive a write→read round-trip.""" + rows = [ + make_row("s1", -1, "metadata", text_confidence=None), + make_row("s1", 0, "text", text_content="first paragraph", text_confidence=0.95), + make_row("s1", 1, "text", text_content="second paragraph", text_confidence=None), # gap + make_row("s1", 2, "text", text_content="third paragraph", text_confidence=0.72), + ] + task = InterleavedBatch( + task_id="per_txt", dataset_name="t", data=pd.DataFrame(rows), _metadata={"source_files": ["x.tar"]} + ) + writer = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "per_txt_out"), materialize_on_write=False, mode="overwrite" + ) + write_task = writer.process(task) + + payload, _ = _read_wds_tar(write_task.data[0]) + assert "text_confidence" in payload, "text_confidence missing from JSON payload" + assert payload["text_confidence"] == [0.95, None, 0.72] + + reader = InterleavedWebdatasetReaderStage( + sample_id_field="sample_id", + per_text_fields=("text_confidence",), + ) + result = reader.process(FileGroupTask(task_id="rt2", dataset_name="t", data=write_task.data)) + assert isinstance(result, InterleavedBatch) + df = result.to_pandas() + + text_rows = df[df["modality"] == "text"].sort_values("position") + confidences = [None if pd.isna(v) else v for v in text_rows["text_confidence"]] + assert confidences == [0.95, None, 0.72] + + +def test_wds_writer_non_ascii_per_image_field_roundtrip(tmp_path: Path) -> None: + """Non-ASCII per_image_fields are written as raw UTF-8, not \\uXXXX escapes.""" + alt_texts = ["北极熊", None, "الدب القطبي"] # Chinese, gap, Arabic + fake_png = b"\x89PNG" + rows = [ + make_row("s1", -1, "metadata", image_alt_text=None), + make_row("s1", 0, "image", content_type="image/png", binary_content=fake_png, image_alt_text=alt_texts[0]), + make_row("s1", 1, "image", content_type="image/png", binary_content=fake_png, image_alt_text=alt_texts[1]), + make_row("s1", 2, "image", content_type="image/png", binary_content=fake_png, image_alt_text=alt_texts[2]), + ] + task = InterleavedBatch( + task_id="non_ascii", dataset_name="t", data=pd.DataFrame(rows), _metadata={"source_files": ["x.tar"]} + ) + writer = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "non_ascii_out"), materialize_on_write=False, mode="overwrite" + ) + write_task = writer.process(task) + + _, members = _read_wds_tar(write_task.data[0]) + raw_bytes = next(v for k, v in members.items() if k.endswith(".json")) + assert "\\u" not in raw_bytes.decode("utf-8"), "Non-ASCII chars should not be \\u-escaped" + assert "北极熊".encode() in raw_bytes + assert "الدب القطبي".encode() in raw_bytes + + reader = InterleavedWebdatasetReaderStage( + sample_id_field="sample_id", + per_image_fields=("image_alt_text",), + ) + result = reader.process(FileGroupTask(task_id="rt3", dataset_name="t", data=write_task.data)) + assert isinstance(result, InterleavedBatch) + df = result.to_pandas() + image_rows = df[df["modality"] == "image"].sort_values("position") + recovered = [None if pd.isna(v) else v for v in image_rows["image_alt_text"]] + assert recovered == alt_texts + + +def test_wds_writer_mixed_modality_field_written_as_position_aligned_list(tmp_path: Path) -> None: + """A column with non-null values in BOTH image and text rows is written as a + position-aligned list (one entry per content row). When read back without + per_image_fields / per_text_fields, the entire list lands on the metadata row + as a passthrough JSON string.""" + fake_png = b"\x89PNG" + rows = [ + make_row("s1", -1, "metadata", confidence=None), + make_row("s1", 0, "text", text_content="intro", confidence="text_conf_A"), + make_row("s1", 1, "image", content_type="image/png", binary_content=fake_png, confidence="img_conf_B"), + make_row("s1", 2, "text", text_content="outro", confidence=None), + ] + task = InterleavedBatch( + task_id="mixed", dataset_name="t", data=pd.DataFrame(rows), _metadata={"source_files": ["x.tar"]} + ) + write_task = InterleavedWebdatasetWriterStage( + path=str(tmp_path / "mixed_out"), materialize_on_write=False, mode="overwrite" + ).process(task) + + payload, _ = _read_wds_tar(write_task.data[0]) + assert payload["confidence"] == ["text_conf_A", "img_conf_B", None] + + # Without per_image/per_text declaration, reader treats it as a sample-level passthrough + reader = InterleavedWebdatasetReaderStage(sample_id_field="sample_id") + result = reader.process(FileGroupTask(task_id="rt_mixed", dataset_name="t", data=write_task.data)) + assert isinstance(result, InterleavedBatch) + df = result.to_pandas() + meta_row = df[df["modality"] == "metadata"].iloc[0] + assert json.loads(meta_row["confidence"]) == ["text_conf_A", "img_conf_B", None] + + +def test_is_null_null_variants() -> None: + assert _is_null(None) is True + assert _is_null(float("nan")) is True + assert _is_null(pd.NA) is True + + +def test_is_null_non_null_and_uncomparable() -> None: + class _Uncomparable: + __hash__ = None # type: ignore[assignment] + + def __eq__(self, other: object) -> bool: + msg = "cannot compare" + raise TypeError(msg) + + for v in ("hello", 42, 0, b"bytes", _Uncomparable()): + assert _is_null(v) is False diff --git a/tutorials/multimodal/mint1t_mvp_pipeline.py b/tutorials/multimodal/mint1t_mvp_pipeline.py index 735d615bf1..d5761d890f 100644 --- a/tutorials/multimodal/mint1t_mvp_pipeline.py +++ b/tutorials/multimodal/mint1t_mvp_pipeline.py @@ -17,7 +17,7 @@ from nemo_curator.core.client import RayClient from nemo_curator.pipeline import Pipeline -from nemo_curator.stages.interleaved.io import InterleavedParquetWriterStage, WebdatasetReader +from nemo_curator.stages.interleaved.io import InterleavedParquetWriterStage, InterleavedWebdatasetReader from nemo_curator.stages.interleaved.stages import InterleavedAspectRatioFilterStage @@ -31,7 +31,7 @@ def build_pipeline(args: argparse.Namespace) -> Pipeline: pipe = Pipeline(name="mint1t_mvp_multimodal", description="WebDataset MINT1T -> multimodal rows -> parquet") pipe.add_stage( - WebdatasetReader( + InterleavedWebdatasetReader( file_paths=args.input_path, files_per_partition=args.files_per_partition, blocksize=args.input_blocksize,