Files
DataPipeline/pipeline/io/export.py
T

168 lines
5.6 KiB
Python

"""Dataset export and caching utilities."""
import json
import logging
import os
from pathlib import Path
from typing import Any, Callable, Dict, List, Optional, Union
from datasets import Dataset
from tqdm import tqdm
from pipeline.io.file_scanner import FileScanner
from pipeline.io.hdf5_handler import HDF5Handler
from pipeline.processors import BaseProcessor
from pipeline.packing import pack_tensors
from pipeline.utils import error_handler
logger = logging.getLogger(__name__)
@error_handler()
def export_dataset(
dataset: Dataset,
output_dir: str,
output_prefix: str,
*,
chunk_size: int = 1_000_000,
max_chunks: Optional[int] = None,
process_func: Optional[
Callable[[Dict[str, Any]], Union[Dict[str, Any], List[Dict[str, Any]]]]
] = None,
column: str = "text",
) -> List[str]:
"""Export HuggingFace Dataset to JSONL files in chunks.
Args:
dataset: HuggingFace Dataset object.
output_dir: Output directory.
output_prefix: Output file name prefix, e.g., "chinese-c4-pretrain".
chunk_size: Maximum number of samples per file.
max_chunks: Maximum number of chunks to process (for debugging).
process_func: Single sample transformation function (dict) -> dict | list[dict].
column: Default text column name (only used when process_func is None).
Returns:
List of generated file paths.
"""
os.makedirs(output_dir, exist_ok=True)
total = len(dataset)
num_chunks = (total + chunk_size - 1) // chunk_size
lim = min(max_chunks, num_chunks) if max_chunks else num_chunks
output_files: List[str] = []
for i in range(lim):
start = i * chunk_size
end = min(start + chunk_size, total)
chunk = dataset.select(range(start, end))
path = os.path.join(output_dir, f"{output_prefix}_chunk_{i}.jsonl")
try:
with open(path, "w", encoding="utf-8") as f:
for example in chunk:
processed = (
process_func(example)
if process_func
else {column: example[column]}
)
items = processed if isinstance(processed, list) else [processed]
for item in items:
f.write(json.dumps(item, ensure_ascii=False) + "\n")
output_files.append(path)
logger.info(f"[{i + 1}/{lim}] Saved {path}")
except (OSError, IOError) as e:
logger.error(f"Failed to write chunk {i} to {path}: {e}")
return output_files
@error_handler()
def cache_jsonl(
files: List[str],
output_dir: str,
processor: BaseProcessor,
*,
pack_size: int = -1,
pad_value: int = 0,
batch_size: int = 256,
) -> List[str]:
"""Tokenize JSONL files and pack them into HDF5 storage.
Args:
files: List of JSONL file paths.
output_dir: H5 output directory.
processor: Initialized Processor instance.
pack_size: Packing length, <=0 means no packing.
pad_value: Padding value.
batch_size: Number of records passed to the processor at once.
Returns:
List of generated H5 file paths.
"""
os.makedirs(output_dir, exist_ok=True)
output_files: List[str] = []
output_keys = processor.output_keys
for file_path in files:
file_name = Path(file_path).stem
arrows: Dict[str, List] = {key: [] for key in output_keys}
def append_batch(batch):
items = [item for _, item in batch]
try:
results = processor.process_batch(items)
if len(results) != len(items):
raise RuntimeError(
"Batch processor returned a different number of results"
)
except Exception:
results = []
for line_num, item in batch:
try:
results.append(processor.process(item))
except Exception as e:
logger.warning(
f"Unexpected error processing line {line_num} "
f"in {file_path}: {e}. Skipping line."
)
results.append(None)
for result in results:
if result is not None:
for key in output_keys:
arrows[key].append(result[key])
batch = []
batch_size = max(1, batch_size)
with open(file_path, "r", encoding="utf-8") as f:
for line_num, line in enumerate(
tqdm(f, desc=f"Processing {file_name}", leave=False), start=1
):
try:
batch.append((line_num, json.loads(line)))
if len(batch) >= batch_size:
append_batch(batch)
batch = []
except json.JSONDecodeError as e:
logger.warning(
f"JSON decode error in {file_path} line {line_num}: {e}. Skipping line."
)
if batch:
append_batch(batch)
if pack_size > 0:
dtypes = (
dict(processor.schema.output_fields)
if processor.schema is not None
else None
)
output = pack_tensors(arrows, pack_size, pad_value, dtypes)
else:
output = arrows
h5_path = HDF5Handler.save(output_dir, file_name, output)
output_files.append(h5_path)
logger.info(f"Saved {h5_path}")
return output_files