"""File, HDF5, JSONL I/O operations.""" import json import os import logging from pathlib import Path from typing import Dict, List, Optional, Callable, Union, Any import h5py import torch from torch import Tensor from tqdm import tqdm from datasets import Dataset from pipeline.utils import error_handler from pipeline.processors import BaseProcessor from pipeline.packing import SequencePacker logger = logging.getLogger(__name__) class IOHandler: """File and HDF5 read/write operations.""" @staticmethod def fetch_files(directory: str, suffix: Optional[str] = None) -> List[str]: files = [ os.path.join(root, f) for root, _, files in os.walk(directory) for f in files ] if suffix: files = [f for f in files if f.endswith(suffix)] return sorted(files) @staticmethod def fetch_folders( root_dir: str, filter_func: Optional[Callable[[str], bool]] = None ) -> List[str]: folders = [] for root, dirs, _ in os.walk(root_dir): for dir_name in dirs: folder_path = os.path.join(root, dir_name) if filter_func is None or filter_func(folder_path): folders.append(folder_path) return folders @staticmethod @error_handler() def save_h5( output_dir: str, file_name: str, tensor_group: Dict[str, List[Tensor]] ) -> None: os.makedirs(output_dir, exist_ok=True) full_path = os.path.join(output_dir, f"{file_name}.h5") with h5py.File(full_path, "w") as f: for key, tensors in tensor_group.items(): grp = f.create_group(key) for idx, tensor in enumerate(tensors): grp.create_dataset(f"data_{idx}", data=tensor.cpu().numpy()) @staticmethod @error_handler() def load_h5(file_path: str, share_memory: bool = True) -> Dict[str, List[Tensor]]: tensor_group: Dict[str, List[Tensor]] = {} root_path = Path(file_path) h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5")) for h5_file in h5_files: with h5py.File(h5_file, "r") as f: for key in f.keys(): grp = f[key] dsets = [] for dset_name in grp.keys(): dset = grp[dset_name] tensor = torch.from_numpy(dset[:]) if share_memory: tensor = tensor.share_memory_() dsets.append(tensor) if tensor_group.get(key) is None: tensor_group[key] = [] tensor_group[key].extend(dsets) return tensor_group # ── Stage 1: Export HuggingFace Dataset to JSONL ────────────────────────── @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 # ── Stage 2: Tokenize JSONL and cache to HDF5 ──────────────────────────── @error_handler() def cache_jsonl( files: List[str], output_dir: str, processor: BaseProcessor, *, pack_size: int = -1, pad_value: int = 1, ) -> 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 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} 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: result = processor.process(json.loads(line)) if result is not None: for key in output_keys: arrows[key].append(result[key]) except json.JSONDecodeError as e: logger.warning( f"JSON decode error in {file_path} line {line_num}: {e}. Skipping line." ) continue except Exception as e: logger.warning( f"Unexpected error processing line {line_num} in {file_path}: {e}. Skipping line." ) continue if pack_size > 0: output = {} for key in output_keys: packer = SequencePacker(pack_size, pad_value) output[key] = packer.pack(arrows[key]) else: output = arrows IOHandler.save_h5(output_dir, file_name, output) h5_path = os.path.join(output_dir, f"{file_name}.h5") output_files.append(h5_path) logger.info(f"Saved {h5_path}") return output_files