202 lines
6.9 KiB
Python
202 lines
6.9 KiB
Python
"""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
|