refactor: 重构流水线架构,添加Pipeline抽象并拆分IOHandler
This commit is contained in:
+69
-11
@@ -1,10 +1,50 @@
|
|||||||
import logging
|
"""DataPipeline: A flexible data processing pipeline for LLM training.
|
||||||
|
|
||||||
|
Architecture:
|
||||||
|
- Pipeline: Composable stage-based processing
|
||||||
|
- Processors: Data transformation (pretrain, sft, dpo)
|
||||||
|
- Strategies: Prompt format abstraction (ChatML, Alpaca)
|
||||||
|
- I/O: File scanning and HDF5 storage
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
from pipeline import Pipeline, ProcessorFactory, FileScanner, HDF5Handler
|
||||||
|
from pipeline.pipeline import TransformStage
|
||||||
|
from pipeline.io import export_dataset, cache_jsonl
|
||||||
|
|
||||||
|
# Create pipeline
|
||||||
|
pipeline = Pipeline()
|
||||||
|
pipeline.add_stages(
|
||||||
|
TransformStage("normalize", normalizer.normalize),
|
||||||
|
TransformStage("tokenize", tokenizer.encode),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Process data
|
||||||
|
results = pipeline.run(texts)
|
||||||
|
HDF5Handler.save("./output", "data", {"tokens": results})
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Core modules
|
||||||
|
from pipeline.pipeline import Pipeline, PipelineConfig, Stage, TransformStage
|
||||||
from pipeline.tokenize import AutoTokenizer, ChatTemplate, train_bpe_tokenizer
|
from pipeline.tokenize import AutoTokenizer, ChatTemplate, train_bpe_tokenizer
|
||||||
from pipeline.text import TextNormalizer
|
from pipeline.text import TextNormalizer
|
||||||
from pipeline.packing import SequencePacker
|
from pipeline.packing import SequencePacker
|
||||||
from pipeline.io import IOHandler, export_dataset, cache_jsonl
|
|
||||||
from pipeline.processors import ProcessorFactory, BaseProcessor
|
# I/O module
|
||||||
from pipeline.utils import setup_logging
|
from pipeline.io import FileScanner, HDF5Handler, export_dataset, cache_jsonl
|
||||||
|
|
||||||
|
# Processors
|
||||||
|
from pipeline.processors import (
|
||||||
|
ProcessorFactory,
|
||||||
|
BaseProcessor,
|
||||||
|
ProcessorSchema,
|
||||||
|
ProcessorConfig,
|
||||||
|
PreTrainProcessor,
|
||||||
|
SFTProcessor,
|
||||||
|
DPOProcessor,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Strategies
|
||||||
from pipeline.strategies import (
|
from pipeline.strategies import (
|
||||||
PromptStrategy,
|
PromptStrategy,
|
||||||
ChatMLStrategy,
|
ChatMLStrategy,
|
||||||
@@ -12,25 +52,43 @@ from pipeline.strategies import (
|
|||||||
StrategyFactory,
|
StrategyFactory,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Configure project-level logging
|
# Utilities (lazy initialization)
|
||||||
setup_logging()
|
from pipeline import utils
|
||||||
|
|
||||||
|
# Expose setup_logging for explicit use
|
||||||
|
setup_logging = utils.setup_logging
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
# Pipeline
|
||||||
|
"Pipeline",
|
||||||
|
"PipelineConfig",
|
||||||
|
"Stage",
|
||||||
|
"TransformStage",
|
||||||
# Tokenizer
|
# Tokenizer
|
||||||
"AutoTokenizer",
|
"AutoTokenizer",
|
||||||
"ChatTemplate",
|
"ChatTemplate",
|
||||||
"train_bpe_tokenizer",
|
"train_bpe_tokenizer",
|
||||||
# Core modules
|
# Text processing
|
||||||
"TextNormalizer",
|
"TextNormalizer",
|
||||||
"SequencePacker",
|
"SequencePacker",
|
||||||
"IOHandler",
|
# I/O
|
||||||
"ProcessorFactory",
|
"FileScanner",
|
||||||
"BaseProcessor",
|
"HDF5Handler",
|
||||||
"export_dataset",
|
"export_dataset",
|
||||||
"cache_jsonl",
|
"cache_jsonl",
|
||||||
# Strategy pattern
|
# Processors
|
||||||
|
"ProcessorFactory",
|
||||||
|
"BaseProcessor",
|
||||||
|
"ProcessorSchema",
|
||||||
|
"ProcessorConfig",
|
||||||
|
"PreTrainProcessor",
|
||||||
|
"SFTProcessor",
|
||||||
|
"DPOProcessor",
|
||||||
|
# Strategies
|
||||||
"PromptStrategy",
|
"PromptStrategy",
|
||||||
"ChatMLStrategy",
|
"ChatMLStrategy",
|
||||||
"AlpacaStrategy",
|
"AlpacaStrategy",
|
||||||
"StrategyFactory",
|
"StrategyFactory",
|
||||||
|
# Utils
|
||||||
|
"setup_logging",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -0,0 +1,19 @@
|
|||||||
|
"""I/O module for file operations, HDF5 storage, and dataset export.
|
||||||
|
|
||||||
|
This module provides:
|
||||||
|
- FileScanner: File and directory scanning utilities
|
||||||
|
- HDF5Handler: Tensor data persistence
|
||||||
|
- export_dataset: HuggingFace Dataset to JSONL export
|
||||||
|
- cache_jsonl: JSONL to HDF5 tokenization and caching
|
||||||
|
"""
|
||||||
|
|
||||||
|
from pipeline.io.file_scanner import FileScanner
|
||||||
|
from pipeline.io.hdf5_handler import HDF5Handler
|
||||||
|
from pipeline.io.export import export_dataset, cache_jsonl
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"FileScanner",
|
||||||
|
"HDF5Handler",
|
||||||
|
"export_dataset",
|
||||||
|
"cache_jsonl",
|
||||||
|
]
|
||||||
@@ -1,94 +1,23 @@
|
|||||||
"""File, HDF5, JSONL I/O operations."""
|
"""Dataset export and caching utilities."""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import os
|
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, List, Optional, Callable, Union, Any
|
from typing import Any, Callable, Dict, List, Optional, Union
|
||||||
|
|
||||||
import h5py
|
|
||||||
import torch
|
|
||||||
from torch import Tensor
|
|
||||||
from tqdm import tqdm
|
|
||||||
from datasets import Dataset
|
from datasets import Dataset
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
from pipeline.utils import error_handler
|
from pipeline.io.file_scanner import FileScanner
|
||||||
|
from pipeline.io.hdf5_handler import HDF5Handler
|
||||||
from pipeline.processors import BaseProcessor
|
from pipeline.processors import BaseProcessor
|
||||||
from pipeline.packing import SequencePacker
|
from pipeline.packing import SequencePacker
|
||||||
|
from pipeline.utils import error_handler
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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()
|
@error_handler()
|
||||||
def export_dataset(
|
def export_dataset(
|
||||||
dataset: Dataset,
|
dataset: Dataset,
|
||||||
@@ -102,20 +31,19 @@ def export_dataset(
|
|||||||
] = None,
|
] = None,
|
||||||
column: str = "text",
|
column: str = "text",
|
||||||
) -> List[str]:
|
) -> List[str]:
|
||||||
"""
|
"""Export HuggingFace Dataset to JSONL files in chunks.
|
||||||
Export HuggingFace Dataset to JSONL files in chunks.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
dataset: HuggingFace Dataset object
|
dataset: HuggingFace Dataset object.
|
||||||
output_dir: Output directory
|
output_dir: Output directory.
|
||||||
output_prefix: Output file name prefix, e.g., "chinese-c4-pretrain"
|
output_prefix: Output file name prefix, e.g., "chinese-c4-pretrain".
|
||||||
chunk_size: Maximum number of samples per file
|
chunk_size: Maximum number of samples per file.
|
||||||
max_chunks: Maximum number of chunks to process (for debugging)
|
max_chunks: Maximum number of chunks to process (for debugging).
|
||||||
process_func: Single sample transformation function (dict) -> dict | list[dict]
|
process_func: Single sample transformation function (dict) -> dict | list[dict].
|
||||||
column: Default text column name (only used when process_func is None)
|
column: Default text column name (only used when process_func is None).
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
List of generated file paths
|
List of generated file paths.
|
||||||
"""
|
"""
|
||||||
os.makedirs(output_dir, exist_ok=True)
|
os.makedirs(output_dir, exist_ok=True)
|
||||||
total = len(dataset)
|
total = len(dataset)
|
||||||
@@ -148,9 +76,6 @@ def export_dataset(
|
|||||||
return output_files
|
return output_files
|
||||||
|
|
||||||
|
|
||||||
# ── Stage 2: Tokenize JSONL and cache to HDF5 ────────────────────────────
|
|
||||||
|
|
||||||
|
|
||||||
@error_handler()
|
@error_handler()
|
||||||
def cache_jsonl(
|
def cache_jsonl(
|
||||||
files: List[str],
|
files: List[str],
|
||||||
@@ -160,18 +85,17 @@ def cache_jsonl(
|
|||||||
pack_size: int = -1,
|
pack_size: int = -1,
|
||||||
pad_value: int = 1,
|
pad_value: int = 1,
|
||||||
) -> List[str]:
|
) -> List[str]:
|
||||||
"""
|
"""Tokenize JSONL files and pack them into HDF5 storage.
|
||||||
Tokenize JSONL files and pack them into HDF5 storage.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
files: List of JSONL file paths
|
files: List of JSONL file paths.
|
||||||
output_dir: H5 output directory
|
output_dir: H5 output directory.
|
||||||
processor: Initialized Processor instance
|
processor: Initialized Processor instance.
|
||||||
pack_size: Packing length, <=0 means no packing
|
pack_size: Packing length, <=0 means no packing.
|
||||||
pad_value: Padding value
|
pad_value: Padding value.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
List of generated H5 file paths
|
List of generated H5 file paths.
|
||||||
"""
|
"""
|
||||||
os.makedirs(output_dir, exist_ok=True)
|
os.makedirs(output_dir, exist_ok=True)
|
||||||
output_files: List[str] = []
|
output_files: List[str] = []
|
||||||
@@ -210,8 +134,7 @@ def cache_jsonl(
|
|||||||
else:
|
else:
|
||||||
output = arrows
|
output = arrows
|
||||||
|
|
||||||
IOHandler.save_h5(output_dir, file_name, output)
|
h5_path = HDF5Handler.save(output_dir, file_name, output)
|
||||||
h5_path = os.path.join(output_dir, f"{file_name}.h5")
|
|
||||||
output_files.append(h5_path)
|
output_files.append(h5_path)
|
||||||
logger.info(f"Saved {h5_path}")
|
logger.info(f"Saved {h5_path}")
|
||||||
|
|
||||||
@@ -0,0 +1,98 @@
|
|||||||
|
"""File system scanning utilities."""
|
||||||
|
|
||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Callable, List, Optional
|
||||||
|
|
||||||
|
|
||||||
|
class FileScanner:
|
||||||
|
"""Utility for scanning and filtering files in a directory tree.
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
scanner = FileScanner()
|
||||||
|
files = scanner.scan("./data", suffix=".jsonl")
|
||||||
|
folders = scanner.scan_folders("./data", filter_func=is_valid_dir)
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def scan(
|
||||||
|
directory: str,
|
||||||
|
suffix: Optional[str] = None,
|
||||||
|
recursive: bool = True,
|
||||||
|
relative: bool = False,
|
||||||
|
) -> List[str]:
|
||||||
|
"""Scan directory for files matching criteria.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
directory: Root directory to scan.
|
||||||
|
suffix: Filter files by extension (e.g., ".jsonl").
|
||||||
|
recursive: Whether to search subdirectories.
|
||||||
|
relative: Return relative paths if True.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Sorted list of file paths.
|
||||||
|
"""
|
||||||
|
root_path = Path(directory)
|
||||||
|
if not root_path.exists():
|
||||||
|
return []
|
||||||
|
|
||||||
|
if recursive:
|
||||||
|
pattern = "**/*" + suffix if suffix else "**/*"
|
||||||
|
files = [str(p) for p in root_path.glob(pattern) if p.is_file()]
|
||||||
|
else:
|
||||||
|
pattern = "*" + suffix if suffix else "*"
|
||||||
|
files = [str(p) for p in root_path.glob(pattern) if p.is_file()]
|
||||||
|
|
||||||
|
if relative:
|
||||||
|
files = [str(Path(f).relative_to(root_path)) for f in files]
|
||||||
|
|
||||||
|
return sorted(files)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def scan_folders(
|
||||||
|
directory: str,
|
||||||
|
filter_func: Optional[Callable[[str], bool]] = None,
|
||||||
|
recursive: bool = True,
|
||||||
|
) -> List[str]:
|
||||||
|
"""Scan directory for subdirectories.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
directory: Root directory to scan.
|
||||||
|
filter_func: Optional predicate to filter directories.
|
||||||
|
recursive: Whether to search subdirectories.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Sorted list of directory paths.
|
||||||
|
"""
|
||||||
|
root_path = Path(directory)
|
||||||
|
if not root_path.exists():
|
||||||
|
return []
|
||||||
|
|
||||||
|
if recursive:
|
||||||
|
folders = [str(p) for p in root_path.rglob("*") if p.is_dir()]
|
||||||
|
else:
|
||||||
|
folders = [str(p) for p in root_path.glob("*") if p.is_dir()]
|
||||||
|
|
||||||
|
if filter_func:
|
||||||
|
folders = [f for f in folders if filter_func(f)]
|
||||||
|
|
||||||
|
return sorted(folders)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def group_by_extension(files: List[str]) -> dict[str, List[str]]:
|
||||||
|
"""Group files by their extension.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
files: List of file paths.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary mapping extension to list of files.
|
||||||
|
"""
|
||||||
|
groups: dict[str, List[str]] = {}
|
||||||
|
for f in files:
|
||||||
|
ext = Path(f).suffix
|
||||||
|
if ext not in groups:
|
||||||
|
groups[ext] = []
|
||||||
|
groups[ext].append(f)
|
||||||
|
return groups
|
||||||
@@ -0,0 +1,133 @@
|
|||||||
|
"""HDF5 storage operations for tensor data."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
|
import h5py
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from pipeline.utils import error_handler
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class HDF5Handler:
|
||||||
|
"""Handler for reading and writing tensor data to HDF5 files.
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
handler = HDF5Handler()
|
||||||
|
handler.save(output_dir, "data", {"input_ids": [tensor1, tensor2]})
|
||||||
|
|
||||||
|
loaded = handler.load("./output/data.h5")
|
||||||
|
for tensor in loaded["input_ids"]:
|
||||||
|
print(tensor.shape)
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
@error_handler()
|
||||||
|
def save(
|
||||||
|
output_dir: str,
|
||||||
|
file_name: str,
|
||||||
|
tensor_group: Dict[str, List[Tensor]],
|
||||||
|
extension: str = ".h5",
|
||||||
|
) -> str:
|
||||||
|
"""Save tensor groups to HDF5 file.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
output_dir: Output directory path.
|
||||||
|
file_name: Base name for the output file (without extension).
|
||||||
|
tensor_group: Dictionary mapping group names to tensor lists.
|
||||||
|
extension: File extension (default: ".h5").
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Path to the saved file.
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
os.makedirs(output_dir, exist_ok=True)
|
||||||
|
full_path = os.path.join(output_dir, f"{file_name}{extension}")
|
||||||
|
|
||||||
|
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())
|
||||||
|
|
||||||
|
logger.info(f"Saved HDF5 file: {full_path}")
|
||||||
|
return full_path
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
@error_handler()
|
||||||
|
def load(
|
||||||
|
file_path: str,
|
||||||
|
share_memory: bool = True,
|
||||||
|
device: Optional[torch.device] = None,
|
||||||
|
) -> Dict[str, List[Tensor]]:
|
||||||
|
"""Load tensor groups from HDF5 file.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_path: Path to HDF5 file or directory containing HDF5 files.
|
||||||
|
share_memory: Whether to use shared memory for tensors.
|
||||||
|
device: Target device for tensors (default: CPU).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary mapping group names to tensor lists.
|
||||||
|
"""
|
||||||
|
root_path = Path(file_path)
|
||||||
|
h5_files = []
|
||||||
|
|
||||||
|
if root_path.is_file():
|
||||||
|
h5_files = [root_path]
|
||||||
|
else:
|
||||||
|
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
||||||
|
|
||||||
|
if not h5_files:
|
||||||
|
logger.warning(f"No HDF5 files found at: {file_path}")
|
||||||
|
return {}
|
||||||
|
|
||||||
|
tensor_group: Dict[str, List[Tensor]] = {}
|
||||||
|
|
||||||
|
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 device is not None:
|
||||||
|
tensor = tensor.to(device)
|
||||||
|
elif share_memory:
|
||||||
|
tensor = tensor.share_memory_()
|
||||||
|
|
||||||
|
dsets.append(tensor)
|
||||||
|
|
||||||
|
if tensor_group.get(key) is None:
|
||||||
|
tensor_group[key] = []
|
||||||
|
tensor_group[key].extend(dsets)
|
||||||
|
|
||||||
|
logger.info(f"Loaded HDF5: {len(tensor_group)} groups, "
|
||||||
|
f"{sum(len(v) for v in tensor_group.values())} total tensors")
|
||||||
|
|
||||||
|
return tensor_group
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_metadata(file_path: str) -> Dict[str, int]:
|
||||||
|
"""Get metadata about an HDF5 file without loading full data.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_path: Path to HDF5 file.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with group names and tensor counts.
|
||||||
|
"""
|
||||||
|
metadata: Dict[str, int] = {}
|
||||||
|
|
||||||
|
with h5py.File(file_path, "r") as f:
|
||||||
|
for key in f.keys():
|
||||||
|
metadata[key] = len(f[key].keys())
|
||||||
|
|
||||||
|
return metadata
|
||||||
@@ -0,0 +1,196 @@
|
|||||||
|
"""Pipeline abstraction for composable data processing stages."""
|
||||||
|
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any, Dict, Iterator, List, Optional, TypeVar, Generic
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
T = TypeVar("T")
|
||||||
|
R = TypeVar("R")
|
||||||
|
|
||||||
|
|
||||||
|
class Stage(ABC, Generic[T, R]):
|
||||||
|
"""Abstract base class for pipeline stages.
|
||||||
|
|
||||||
|
A Stage represents a single processing step that transforms input data
|
||||||
|
and can be composed with other stages to form a pipeline.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@property
|
||||||
|
@abstractmethod
|
||||||
|
def name(self) -> str:
|
||||||
|
"""Return the stage name for logging/debugging."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def process(self, input_data: T) -> R:
|
||||||
|
"""Process input data and return transformed output."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
return f"{self.__class__.__name__}(name='{self.name}')"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class PipelineConfig:
|
||||||
|
"""Configuration for pipeline execution.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
name: Pipeline name for identification.
|
||||||
|
enable_logging: Enable per-stage logging.
|
||||||
|
continue_on_error: Continue processing if a stage fails.
|
||||||
|
error_threshold: Maximum errors before stopping (-1 for unlimited).
|
||||||
|
"""
|
||||||
|
|
||||||
|
name: str = "pipeline"
|
||||||
|
enable_logging: bool = True
|
||||||
|
continue_on_error: bool = True
|
||||||
|
error_threshold: int = -1
|
||||||
|
|
||||||
|
|
||||||
|
class Pipeline(Generic[T]):
|
||||||
|
"""Composable pipeline for sequential data processing.
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
pipeline = Pipeline(config=PipelineConfig(name="data-prep"))
|
||||||
|
pipeline.add_stage(TextNormalizationStage(normalizer))
|
||||||
|
pipeline.add_stage(TokenizationStage(tokenizer))
|
||||||
|
pipeline.add_stage(PackingStage(packer))
|
||||||
|
|
||||||
|
results = pipeline.run(input_data)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, config: Optional[PipelineConfig] = None):
|
||||||
|
self.config = config or PipelineConfig()
|
||||||
|
self._stages: List[Stage] = []
|
||||||
|
self._error_count: int = 0
|
||||||
|
|
||||||
|
def add_stage(self, stage: Stage) -> "Pipeline":
|
||||||
|
"""Add a stage to the pipeline (fluent interface)."""
|
||||||
|
self._stages.append(stage)
|
||||||
|
return self
|
||||||
|
|
||||||
|
def add_stages(self, *stages: Stage) -> "Pipeline":
|
||||||
|
"""Add multiple stages at once."""
|
||||||
|
self._stages.extend(stages)
|
||||||
|
return self
|
||||||
|
|
||||||
|
def run(self, input_data: List[T]) -> List[Any]:
|
||||||
|
"""Run all stages sequentially on input data.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
input_data: List of input items to process.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of processed results.
|
||||||
|
"""
|
||||||
|
if self.config.enable_logging:
|
||||||
|
logger.info(f"Starting pipeline '{self.config.name}' with {len(self._stages)} stages")
|
||||||
|
|
||||||
|
results = input_data
|
||||||
|
for stage in self._stages:
|
||||||
|
if self.config.enable_logging:
|
||||||
|
logger.info(f"Running stage: {stage.name}")
|
||||||
|
|
||||||
|
new_results = []
|
||||||
|
for item in results:
|
||||||
|
if self._should_stop():
|
||||||
|
break
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = stage.process(item)
|
||||||
|
if result is not None:
|
||||||
|
new_results.append(result)
|
||||||
|
except Exception as e:
|
||||||
|
self._handle_error(stage, item, e)
|
||||||
|
|
||||||
|
results = new_results
|
||||||
|
|
||||||
|
if self.config.enable_logging:
|
||||||
|
logger.info(f"Pipeline '{self.config.name}' completed: {len(results)} items")
|
||||||
|
|
||||||
|
return results
|
||||||
|
|
||||||
|
def run_stream(self, input_data: List[T]) -> Iterator[Any]:
|
||||||
|
"""Run pipeline as a generator for memory-efficient processing.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
input_data: List of input items to process.
|
||||||
|
|
||||||
|
Yields:
|
||||||
|
Processed results one at a time.
|
||||||
|
"""
|
||||||
|
for stage in self._stages:
|
||||||
|
if self.config.enable_logging:
|
||||||
|
logger.info(f"Running stage: {stage.name}")
|
||||||
|
|
||||||
|
for item in input_data:
|
||||||
|
if self._should_stop():
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = stage.process(item)
|
||||||
|
if result is not None:
|
||||||
|
yield result
|
||||||
|
except Exception as e:
|
||||||
|
self._handle_error(stage, item, e)
|
||||||
|
|
||||||
|
def _should_stop(self) -> bool:
|
||||||
|
"""Check if pipeline should stop processing."""
|
||||||
|
if self.config.error_threshold < 0:
|
||||||
|
return False
|
||||||
|
return self._error_count >= self.config.error_threshold
|
||||||
|
|
||||||
|
def _handle_error(self, stage: Stage, item: Any, error: Exception) -> None:
|
||||||
|
"""Handle processing error."""
|
||||||
|
self._error_count += 1
|
||||||
|
error_msg = f"Error in stage '{stage.name}': {error}"
|
||||||
|
if self.config.continue_on_error:
|
||||||
|
logger.warning(error_msg)
|
||||||
|
else:
|
||||||
|
raise RuntimeError(error_msg) from error
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
stage_names = [s.name for s in self._stages]
|
||||||
|
return f"Pipeline(name='{self.config.name}', stages={stage_names})"
|
||||||
|
|
||||||
|
|
||||||
|
# ── Common Stage Implementations ──────────────────────────────────────────────
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TransformStage(Stage):
|
||||||
|
"""Stage that applies a transformation function.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
transform: Callable that transforms input to output.
|
||||||
|
name: Stage name.
|
||||||
|
"""
|
||||||
|
|
||||||
|
transform: callable
|
||||||
|
name: str
|
||||||
|
_name: str = field(init=False, repr=False, compare=False, hash=False)
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
self._name = self.name
|
||||||
|
|
||||||
|
def process(self, input_data: T) -> R:
|
||||||
|
return self.transform(input_data)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class FilterStage(Stage):
|
||||||
|
"""Stage that filters items based on a predicate.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
predicate: Callable that returns True to keep item.
|
||||||
|
name: Stage name.
|
||||||
|
"""
|
||||||
|
|
||||||
|
predicate: callable
|
||||||
|
name: str
|
||||||
|
|
||||||
|
def process(self, input_data: T) -> Optional[T]:
|
||||||
|
return input_data if self.predicate(input_data) else None
|
||||||
@@ -4,15 +4,18 @@ Processor classes are registered at definition time via decorators and
|
|||||||
can be created through :class:`ProcessorFactory`.
|
can be created through :class:`ProcessorFactory`.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from pipeline.processors.base import BaseProcessor
|
from pipeline.processors.base import BaseProcessor, ProcessorSchema, encode_with_mask
|
||||||
from pipeline.processors.factory import ProcessorFactory
|
from pipeline.processors.factory import ProcessorFactory, ProcessorConfig
|
||||||
from pipeline.processors.pretrain import PreTrainProcessor
|
from pipeline.processors.pretrain import PreTrainProcessor
|
||||||
from pipeline.processors.sft import SFTProcessor
|
from pipeline.processors.sft import SFTProcessor
|
||||||
from pipeline.processors.dpo import DPOProcessor
|
from pipeline.processors.dpo import DPOProcessor
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"BaseProcessor",
|
"BaseProcessor",
|
||||||
|
"ProcessorSchema",
|
||||||
|
"ProcessorConfig",
|
||||||
"ProcessorFactory",
|
"ProcessorFactory",
|
||||||
|
"encode_with_mask",
|
||||||
"PreTrainProcessor",
|
"PreTrainProcessor",
|
||||||
"SFTProcessor",
|
"SFTProcessor",
|
||||||
"DPOProcessor",
|
"DPOProcessor",
|
||||||
|
|||||||
+106
-12
@@ -1,32 +1,126 @@
|
|||||||
"""Processor base class and shared utilities."""
|
"""Processor base class and shared utilities."""
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Dict, List, Any, Tuple
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, Dict, List, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
def _encode_with_mask(
|
@dataclass(frozen=True)
|
||||||
prompt_tokens: List[int],
|
class ProcessorSchema:
|
||||||
response_tokens: List[int],
|
"""Schema definition for processor input/output contracts.
|
||||||
) -> Tuple[Tensor, Tensor]:
|
|
||||||
"""Concatenate token lists and build loss mask (prompt=False, response=True)."""
|
Attributes:
|
||||||
q_len = len(prompt_tokens)
|
input_fields: Required input field names and their types.
|
||||||
combined = torch.tensor(prompt_tokens + response_tokens, dtype=torch.int32)
|
output_fields: Output field names and their tensor dtypes.
|
||||||
mask = torch.zeros(q_len + len(response_tokens), dtype=torch.bool)
|
"""
|
||||||
mask[q_len:] = True
|
|
||||||
return combined, mask
|
input_fields: Dict[str, type]
|
||||||
|
output_fields: Dict[str, torch.dtype]
|
||||||
|
|
||||||
|
|
||||||
class BaseProcessor(ABC):
|
class BaseProcessor(ABC):
|
||||||
"""Abstract base class for processors."""
|
"""Abstract base class for data processors.
|
||||||
|
|
||||||
|
Processors transform raw data (e.g., text, JSON) into tokenized tensors
|
||||||
|
suitable for model training.
|
||||||
|
|
||||||
|
Subclasses must implement:
|
||||||
|
- process(): Transform a single input sample
|
||||||
|
- output_keys: Declare output tensor names
|
||||||
|
- schema: Define input/output contracts (optional but recommended)
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
class MyProcessor(BaseProcessor):
|
||||||
|
@property
|
||||||
|
def schema(self) -> ProcessorSchema:
|
||||||
|
return ProcessorSchema(
|
||||||
|
input_fields={"text": str},
|
||||||
|
output_fields={"tokens": torch.int32}
|
||||||
|
)
|
||||||
|
|
||||||
|
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
|
||||||
|
tokens = self.tokenizer.encode(input_dict["text"])
|
||||||
|
return {"tokens": torch.tensor(tokens, dtype=torch.int32)}
|
||||||
|
|
||||||
|
@property
|
||||||
|
def output_keys(self) -> List[str]:
|
||||||
|
return ["tokens"]
|
||||||
|
"""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def schema(self) -> ProcessorSchema:
|
||||||
|
"""Return the input/output schema for this processor.
|
||||||
|
|
||||||
|
Override this property to define explicit contracts.
|
||||||
|
Default returns None, meaning schema is not defined.
|
||||||
|
"""
|
||||||
|
return None
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
|
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
|
||||||
|
"""Process a single input sample.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
input_dict: Dictionary containing input fields as defined by schema.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary mapping output key names to tensors.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
KeyError: If required input fields are missing.
|
||||||
|
ValueError: If input data is invalid.
|
||||||
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def output_keys(self) -> List[str]:
|
def output_keys(self) -> List[str]:
|
||||||
|
"""Return list of output tensor key names."""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def validate_input(self, input_dict: Dict[str, Any]) -> None:
|
||||||
|
"""Validate input against schema before processing.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
input_dict: Input dictionary to validate.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
KeyError: If required fields are missing.
|
||||||
|
TypeError: If field types don't match schema.
|
||||||
|
"""
|
||||||
|
schema = self.schema
|
||||||
|
if schema is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
for field_name, expected_type in schema.input_fields.items():
|
||||||
|
if field_name not in input_dict:
|
||||||
|
raise KeyError(f"Missing required input field: '{field_name}'")
|
||||||
|
if not isinstance(input_dict[field_name], expected_type):
|
||||||
|
raise TypeError(
|
||||||
|
f"Field '{field_name}' expected type {expected_type.__name__}, "
|
||||||
|
f"got {type(input_dict[field_name]).__name__}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def encode_with_mask(
|
||||||
|
prompt_tokens: List[int],
|
||||||
|
response_tokens: List[int],
|
||||||
|
) -> Tuple[Tensor, Tensor]:
|
||||||
|
"""Concatenate token lists and build loss mask (prompt=False, response=True).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt_tokens: Token IDs for the prompt/question.
|
||||||
|
response_tokens: Token IDs for the response/answer.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (combined_tokens, loss_mask) where loss_mask is True for response tokens.
|
||||||
|
"""
|
||||||
|
q_len = len(prompt_tokens)
|
||||||
|
combined = torch.tensor(prompt_tokens + response_tokens, dtype=torch.int32)
|
||||||
|
mask = torch.zeros(q_len + len(response_tokens), dtype=torch.bool)
|
||||||
|
mask[q_len:] = True
|
||||||
|
return combined, mask
|
||||||
|
|||||||
@@ -1,21 +1,32 @@
|
|||||||
"""DPO preference learning data processor."""
|
"""DPO preference learning data processor."""
|
||||||
|
|
||||||
from typing import Dict, List, Any, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from pipeline.tokenize import AutoTokenizer
|
from pipeline.tokenize import AutoTokenizer
|
||||||
from pipeline.strategies import PromptStrategy, ChatMLStrategy
|
from pipeline.strategies import PromptStrategy, ChatMLStrategy
|
||||||
from pipeline.processors.base import BaseProcessor, _encode_with_mask
|
from pipeline.processors.base import BaseProcessor, ProcessorSchema, encode_with_mask
|
||||||
from pipeline.processors.factory import ProcessorFactory
|
from pipeline.processors.factory import ProcessorFactory
|
||||||
|
|
||||||
|
|
||||||
@ProcessorFactory.register("dpo")
|
@ProcessorFactory.register("dpo")
|
||||||
class DPOProcessor(BaseProcessor):
|
class DPOProcessor(BaseProcessor):
|
||||||
"""DPO preference learning data processor.
|
"""DPO (Direct Preference Optimization) data processor.
|
||||||
|
|
||||||
Supports custom prompt strategy via constructor parameter.
|
Processes query, chosen, and rejected responses for preference learning.
|
||||||
|
|
||||||
|
Input schema:
|
||||||
|
- query: str - User query/prompt
|
||||||
|
- chosen: str - Preferred assistant response
|
||||||
|
- rejected: str - Dispreferred assistant response
|
||||||
|
|
||||||
|
Output schema:
|
||||||
|
- chosen: int32 tensor - Token IDs for preferred response
|
||||||
|
- chosen_mask: bool tensor - True for response tokens
|
||||||
|
- rejected: int32 tensor - Token IDs for dispreferred response
|
||||||
|
- rejected_mask: bool tensor - True for response tokens
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -26,6 +37,22 @@ class DPOProcessor(BaseProcessor):
|
|||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
self.strategy = strategy or ChatMLStrategy(tokenizer)
|
self.strategy = strategy or ChatMLStrategy(tokenizer)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def schema(self) -> ProcessorSchema:
|
||||||
|
return ProcessorSchema(
|
||||||
|
input_fields={
|
||||||
|
"query": str,
|
||||||
|
"chosen": str,
|
||||||
|
"rejected": str,
|
||||||
|
},
|
||||||
|
output_fields={
|
||||||
|
"chosen": torch.int32,
|
||||||
|
"chosen_mask": torch.bool,
|
||||||
|
"rejected": torch.int32,
|
||||||
|
"rejected_mask": torch.bool,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
|
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
|
||||||
query_tokens = self.tokenizer.encode(input_dict["query"])
|
query_tokens = self.tokenizer.encode(input_dict["query"])
|
||||||
chosen_tokens = self.tokenizer.encode(input_dict["chosen"])
|
chosen_tokens = self.tokenizer.encode(input_dict["chosen"])
|
||||||
@@ -33,10 +60,10 @@ class DPOProcessor(BaseProcessor):
|
|||||||
|
|
||||||
prompt = self.strategy.assemble_prompt(query_tokens)
|
prompt = self.strategy.assemble_prompt(query_tokens)
|
||||||
|
|
||||||
chosen_t, chosen_m = _encode_with_mask(
|
chosen_t, chosen_m = encode_with_mask(
|
||||||
prompt, self.strategy.assemble_response(chosen_tokens)
|
prompt, self.strategy.assemble_response(chosen_tokens)
|
||||||
)
|
)
|
||||||
rejected_t, rejected_m = _encode_with_mask(
|
rejected_t, rejected_m = encode_with_mask(
|
||||||
prompt, self.strategy.assemble_response(rejected_tokens)
|
prompt, self.strategy.assemble_response(rejected_tokens)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,12 +1,32 @@
|
|||||||
"""Factory for creating and registering processors."""
|
"""Factory for creating and registering processors with unified interface."""
|
||||||
|
|
||||||
from typing import Dict, List, Any, Optional, Type
|
from dataclasses import dataclass
|
||||||
|
from typing import Dict, List, Optional, Type, Union
|
||||||
|
|
||||||
from pipeline.processors.base import BaseProcessor
|
from pipeline.processors.base import BaseProcessor
|
||||||
from pipeline.tokenize import AutoTokenizer
|
from pipeline.tokenize import AutoTokenizer
|
||||||
from pipeline.strategies import PromptStrategy, StrategyFactory
|
from pipeline.strategies import PromptStrategy, StrategyFactory
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ProcessorConfig:
|
||||||
|
"""Configuration for creating a processor.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
processor_type: Type name for the processor ("pt", "sft", "dpo").
|
||||||
|
tokenizer: Tokenizer instance (required).
|
||||||
|
strategy_name: Name of the strategy to use (optional).
|
||||||
|
strategy: Pre-created strategy instance (optional).
|
||||||
|
strategy_kwargs: Additional arguments for strategy creation.
|
||||||
|
"""
|
||||||
|
|
||||||
|
processor_type: str
|
||||||
|
tokenizer: AutoTokenizer
|
||||||
|
strategy_name: Optional[str] = None
|
||||||
|
strategy: Optional[PromptStrategy] = None
|
||||||
|
strategy_kwargs: Optional[Dict] = None
|
||||||
|
|
||||||
|
|
||||||
class ProcessorFactory:
|
class ProcessorFactory:
|
||||||
"""Registry and factory for BaseProcessor implementations.
|
"""Registry and factory for BaseProcessor implementations.
|
||||||
|
|
||||||
@@ -18,7 +38,16 @@ class ProcessorFactory:
|
|||||||
class CustomProcessor(BaseProcessor):
|
class CustomProcessor(BaseProcessor):
|
||||||
...
|
...
|
||||||
|
|
||||||
processor = ProcessorFactory.create(optimizer, "custom", **kwargs)
|
# Using config object (recommended)
|
||||||
|
config = ProcessorConfig(
|
||||||
|
processor_type="sft",
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
strategy_name="alpaca"
|
||||||
|
)
|
||||||
|
processor = ProcessorFactory.create_from_config(config)
|
||||||
|
|
||||||
|
# Using direct arguments
|
||||||
|
processor = ProcessorFactory.create("pt", tokenizer)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
PROCESSOR_MAP: Dict[str, Type[BaseProcessor]] = {}
|
PROCESSOR_MAP: Dict[str, Type[BaseProcessor]] = {}
|
||||||
@@ -46,10 +75,10 @@ class ProcessorFactory:
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def create(cls, processor_type: str, tokenizer: AutoTokenizer) -> BaseProcessor:
|
def create(cls, processor_type: str, tokenizer: AutoTokenizer) -> BaseProcessor:
|
||||||
"""Create a processor by type name (uses default ChatMLStrategy for SFT/DPO).
|
"""Create a processor by type name.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
processor_type: Registered processor name (e.g. ``"pt"``, ``"sft"``, ``"dpo"``).
|
processor_type: Registered processor name (e.g. "pt", "sft", "dpo").
|
||||||
tokenizer: Tokenizer instance.
|
tokenizer: Tokenizer instance.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -72,14 +101,12 @@ class ProcessorFactory:
|
|||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
strategy: PromptStrategy,
|
strategy: PromptStrategy,
|
||||||
) -> BaseProcessor:
|
) -> BaseProcessor:
|
||||||
"""Create a processor with a custom strategy.
|
"""Create a processor with a pre-configured strategy.
|
||||||
|
|
||||||
Only SFT and DPO processors accept a strategy; PreTrain ignores it.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
processor_type: Registered processor name.
|
processor_type: Registered processor name.
|
||||||
tokenizer: Tokenizer instance.
|
tokenizer: Tokenizer instance.
|
||||||
strategy: Prompt strategy instance.
|
strategy: Pre-created strategy instance.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Processor instance configured with strategy.
|
Processor instance configured with strategy.
|
||||||
@@ -108,8 +135,8 @@ class ProcessorFactory:
|
|||||||
Args:
|
Args:
|
||||||
processor_type: Registered processor name.
|
processor_type: Registered processor name.
|
||||||
tokenizer: Tokenizer instance.
|
tokenizer: Tokenizer instance.
|
||||||
strategy_name: Registered strategy name (``"chatml"``, ``"alpaca"``, etc.).
|
strategy_name: Registered strategy name ("chatml", "alpaca", etc.).
|
||||||
**strategy_kwargs: Forwarded to the strategy constructor.
|
**strategy_kwargs: Additional arguments forwarded to strategy constructor.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Processor instance.
|
Processor instance.
|
||||||
@@ -117,6 +144,44 @@ class ProcessorFactory:
|
|||||||
strategy = StrategyFactory.create(strategy_name, tokenizer, **strategy_kwargs)
|
strategy = StrategyFactory.create(strategy_name, tokenizer, **strategy_kwargs)
|
||||||
return cls.create_with_strategy(processor_type, tokenizer, strategy)
|
return cls.create_with_strategy(processor_type, tokenizer, strategy)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create_from_config(cls, config: ProcessorConfig) -> BaseProcessor:
|
||||||
|
"""Create a processor from a configuration object (unified interface).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config: ProcessorConfig with all creation parameters.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Processor instance.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If processor_type is not registered or strategy is invalid.
|
||||||
|
"""
|
||||||
|
if config.processor_type not in cls.PROCESSOR_MAP:
|
||||||
|
raise ValueError(
|
||||||
|
f"Unknown processor type: '{config.processor_type}'. "
|
||||||
|
f"Supported types: {sorted(cls.PROCESSOR_MAP.keys())}"
|
||||||
|
)
|
||||||
|
|
||||||
|
tokenizer = config.tokenizer
|
||||||
|
strategy_kwargs = config.strategy_kwargs or {}
|
||||||
|
|
||||||
|
# Determine strategy to use
|
||||||
|
strategy: Optional[PromptStrategy] = None
|
||||||
|
if config.strategy is not None:
|
||||||
|
strategy = config.strategy
|
||||||
|
elif config.strategy_name is not None:
|
||||||
|
strategy = StrategyFactory.create(
|
||||||
|
config.strategy_name, tokenizer, **strategy_kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
# Create processor
|
||||||
|
if strategy is not None:
|
||||||
|
return cls.create_with_strategy(
|
||||||
|
config.processor_type, tokenizer, strategy
|
||||||
|
)
|
||||||
|
return cls.create(config.processor_type, tokenizer)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def available_types(cls) -> List[str]:
|
def available_types(cls) -> List[str]:
|
||||||
"""Return list of registered processor type names."""
|
"""Return list of registered processor type names."""
|
||||||
|
|||||||
@@ -1,25 +1,46 @@
|
|||||||
"""Pre-training data processor."""
|
"""Pre-training data processor."""
|
||||||
|
|
||||||
from typing import Dict, List, Any
|
from typing import Any, Dict, List
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from pipeline.tokenize import AutoTokenizer
|
from pipeline.tokenize import AutoTokenizer
|
||||||
from pipeline.processors.base import BaseProcessor
|
from pipeline.processors.base import BaseProcessor, ProcessorSchema
|
||||||
from pipeline.processors.factory import ProcessorFactory
|
from pipeline.processors.factory import ProcessorFactory
|
||||||
|
|
||||||
|
|
||||||
@ProcessorFactory.register("pt")
|
@ProcessorFactory.register("pt")
|
||||||
class PreTrainProcessor(BaseProcessor):
|
class PreTrainProcessor(BaseProcessor):
|
||||||
"""Pre-training data processor."""
|
"""Pre-training data processor.
|
||||||
|
|
||||||
def __init__(self, tokenizer: AutoTokenizer):
|
Processes raw text into tokenized sequences with EOS tokens.
|
||||||
|
|
||||||
|
Input schema:
|
||||||
|
- text: str - Raw text string to tokenize
|
||||||
|
|
||||||
|
Output schema:
|
||||||
|
- sequence: int32 tensor - Token IDs with EOS appended
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
tokenizer: AutoTokenizer,
|
||||||
|
eos_token: str = "<|end▁of▁sentence|>",
|
||||||
|
):
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
|
self._eos_token = eos_token
|
||||||
|
|
||||||
|
@property
|
||||||
|
def schema(self) -> ProcessorSchema:
|
||||||
|
return ProcessorSchema(
|
||||||
|
input_fields={"text": str},
|
||||||
|
output_fields={"sequence": torch.int32},
|
||||||
|
)
|
||||||
|
|
||||||
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
|
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
|
||||||
segment = input_dict["text"]
|
segment = input_dict["text"]
|
||||||
tokens = self.tokenizer.encode(f"{segment}<|end▁of▁sentence|>")
|
tokens = self.tokenizer.encode(f"{segment}{self._eos_token}")
|
||||||
return {"sequence": torch.tensor(tokens, dtype=torch.int32)}
|
return {"sequence": torch.tensor(tokens, dtype=torch.int32)}
|
||||||
|
|
||||||
@property
|
@property
|
||||||
|
|||||||
@@ -1,13 +1,13 @@
|
|||||||
"""Supervised fine-tuning data processor."""
|
"""Supervised fine-tuning data processor."""
|
||||||
|
|
||||||
from typing import Dict, List, Any, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from pipeline.tokenize import AutoTokenizer
|
from pipeline.tokenize import AutoTokenizer
|
||||||
from pipeline.strategies import PromptStrategy, ChatMLStrategy
|
from pipeline.strategies import PromptStrategy, ChatMLStrategy
|
||||||
from pipeline.processors.base import BaseProcessor, _encode_with_mask
|
from pipeline.processors.base import BaseProcessor, ProcessorSchema, encode_with_mask
|
||||||
from pipeline.processors.factory import ProcessorFactory
|
from pipeline.processors.factory import ProcessorFactory
|
||||||
|
|
||||||
|
|
||||||
@@ -15,7 +15,15 @@ from pipeline.processors.factory import ProcessorFactory
|
|||||||
class SFTProcessor(BaseProcessor):
|
class SFTProcessor(BaseProcessor):
|
||||||
"""Supervised fine-tuning data processor.
|
"""Supervised fine-tuning data processor.
|
||||||
|
|
||||||
Supports custom prompt strategy via constructor parameter.
|
Processes query-response pairs into tokenized sequences with loss masks.
|
||||||
|
|
||||||
|
Input schema:
|
||||||
|
- query: str - User query/prompt
|
||||||
|
- response: str - Assistant response
|
||||||
|
|
||||||
|
Output schema:
|
||||||
|
- sequence: int32 tensor - Combined token IDs (query + response)
|
||||||
|
- loss_mask: bool tensor - True for response tokens (compute loss)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -26,6 +34,16 @@ class SFTProcessor(BaseProcessor):
|
|||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
self.strategy = strategy or ChatMLStrategy(tokenizer)
|
self.strategy = strategy or ChatMLStrategy(tokenizer)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def schema(self) -> ProcessorSchema:
|
||||||
|
return ProcessorSchema(
|
||||||
|
input_fields={"query": str, "response": str},
|
||||||
|
output_fields={
|
||||||
|
"sequence": torch.int32,
|
||||||
|
"loss_mask": torch.bool,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
|
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
|
||||||
query_tokens = self.tokenizer.encode(input_dict["query"])
|
query_tokens = self.tokenizer.encode(input_dict["query"])
|
||||||
response_tokens = self.tokenizer.encode(input_dict["response"])
|
response_tokens = self.tokenizer.encode(input_dict["response"])
|
||||||
@@ -33,7 +51,7 @@ class SFTProcessor(BaseProcessor):
|
|||||||
prompt = self.strategy.assemble_prompt(query_tokens)
|
prompt = self.strategy.assemble_prompt(query_tokens)
|
||||||
response = self.strategy.assemble_response(response_tokens)
|
response = self.strategy.assemble_response(response_tokens)
|
||||||
|
|
||||||
tokens, loss_mask = _encode_with_mask(prompt, response)
|
tokens, loss_mask = encode_with_mask(prompt, response)
|
||||||
return {"sequence": tokens, "loss_mask": loss_mask}
|
return {"sequence": tokens, "loss_mask": loss_mask}
|
||||||
|
|
||||||
@property
|
@property
|
||||||
|
|||||||
+26
-9
@@ -11,8 +11,14 @@ Usage:
|
|||||||
import argparse
|
import argparse
|
||||||
import os
|
import os
|
||||||
|
|
||||||
from pipeline import AutoTokenizer, ProcessorFactory, cache_jsonl
|
from pipeline import (
|
||||||
from pipeline.io import IOHandler
|
AutoTokenizer,
|
||||||
|
ProcessorFactory,
|
||||||
|
ProcessorConfig,
|
||||||
|
FileScanner,
|
||||||
|
cache_jsonl,
|
||||||
|
setup_logging,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
@@ -47,9 +53,19 @@ def main():
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--pad-value", type=int, default=1, help="Padding value (default: 1)"
|
"--pad-value", type=int, default=1, help="Padding value (default: 1)"
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--log-level",
|
||||||
|
default="INFO",
|
||||||
|
choices=["DEBUG", "INFO", "WARNING", "ERROR"],
|
||||||
|
help="Logging level (default: INFO)",
|
||||||
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
jsonl_files = IOHandler.fetch_files(args.input_dir, suffix=".jsonl")
|
# Initialize logging explicitly (not automatic anymore)
|
||||||
|
import logging
|
||||||
|
setup_logging(getattr(logging, args.log_level))
|
||||||
|
|
||||||
|
jsonl_files = FileScanner.scan(args.input_dir, suffix=".jsonl")
|
||||||
if not jsonl_files:
|
if not jsonl_files:
|
||||||
print(f"[ERROR] No JSONL files found in {args.input_dir}")
|
print(f"[ERROR] No JSONL files found in {args.input_dir}")
|
||||||
return
|
return
|
||||||
@@ -64,12 +80,13 @@ def main():
|
|||||||
tokenizer = AutoTokenizer(args.tokenizer)
|
tokenizer = AutoTokenizer(args.tokenizer)
|
||||||
print(f"Tokenizer loaded: vocab_size={len(tokenizer)}")
|
print(f"Tokenizer loaded: vocab_size={len(tokenizer)}")
|
||||||
|
|
||||||
if args.strategy:
|
# Use unified config interface
|
||||||
processor = ProcessorFactory.create_with_strategy_name(
|
config = ProcessorConfig(
|
||||||
args.type, tokenizer, args.strategy
|
processor_type=args.type,
|
||||||
)
|
tokenizer=tokenizer,
|
||||||
else:
|
strategy_name=args.strategy,
|
||||||
processor = ProcessorFactory.create(args.type, tokenizer)
|
)
|
||||||
|
processor = ProcessorFactory.create_from_config(config)
|
||||||
|
|
||||||
print(f"Processor: {args.type} ({processor.__class__.__name__})")
|
print(f"Processor: {args.type} ({processor.__class__.__name__})")
|
||||||
print(f"Output keys: {processor.output_keys}")
|
print(f"Output keys: {processor.output_keys}")
|
||||||
|
|||||||
+53
-23
@@ -7,11 +7,11 @@ import torch
|
|||||||
import h5py
|
import h5py
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from pipeline.io import IOHandler
|
from pipeline.io import FileScanner, HDF5Handler
|
||||||
|
|
||||||
|
|
||||||
class TestIOHandler:
|
class TestFileScanner:
|
||||||
def test_fetch_files_in_directory(self):
|
def test_scan_files_in_directory(self):
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
Path(tmpdir, "file1.txt").touch()
|
Path(tmpdir, "file1.txt").touch()
|
||||||
Path(tmpdir, "file2.txt").touch()
|
Path(tmpdir, "file2.txt").touch()
|
||||||
@@ -19,42 +19,61 @@ class TestIOHandler:
|
|||||||
os.makedirs(subdir)
|
os.makedirs(subdir)
|
||||||
Path(subdir, "file3.txt").touch()
|
Path(subdir, "file3.txt").touch()
|
||||||
|
|
||||||
files = IOHandler.fetch_files(tmpdir)
|
files = FileScanner.scan(tmpdir)
|
||||||
assert len(files) == 3
|
assert len(files) == 3
|
||||||
|
|
||||||
def test_fetch_files_empty_directory(self):
|
def test_scan_empty_directory(self):
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
assert IOHandler.fetch_files(tmpdir) == []
|
assert FileScanner.scan(tmpdir) == []
|
||||||
|
|
||||||
def test_fetch_folders_in_directory(self):
|
def test_scan_with_suffix_filter(self):
|
||||||
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
|
Path(tmpdir, "file1.txt").touch()
|
||||||
|
Path(tmpdir, "file2.json").touch()
|
||||||
|
|
||||||
|
txt_files = FileScanner.scan(tmpdir, suffix=".txt")
|
||||||
|
assert len(txt_files) == 1
|
||||||
|
assert txt_files[0].endswith(".txt")
|
||||||
|
|
||||||
|
def test_scan_folders_in_directory(self):
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
os.makedirs(os.path.join(tmpdir, "folder1"))
|
os.makedirs(os.path.join(tmpdir, "folder1"))
|
||||||
os.makedirs(os.path.join(tmpdir, "folder2"))
|
os.makedirs(os.path.join(tmpdir, "folder2"))
|
||||||
os.makedirs(os.path.join(tmpdir, "folder1", "nested"))
|
os.makedirs(os.path.join(tmpdir, "folder1", "nested"))
|
||||||
|
|
||||||
folders = IOHandler.fetch_folders(tmpdir)
|
folders = FileScanner.scan_folders(tmpdir)
|
||||||
assert len(folders) == 3
|
assert len(folders) == 3
|
||||||
|
|
||||||
def test_fetch_folders_with_filter(self):
|
def test_scan_folders_with_filter(self):
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
os.makedirs(os.path.join(tmpdir, "folder1"))
|
os.makedirs(os.path.join(tmpdir, "folder1"))
|
||||||
os.makedirs(os.path.join(tmpdir, "folder2"))
|
os.makedirs(os.path.join(tmpdir, "folder2"))
|
||||||
|
|
||||||
folders = IOHandler.fetch_folders(
|
folders = FileScanner.scan_folders(
|
||||||
tmpdir, filter_func=lambda x: "folder1" in x
|
tmpdir, filter_func=lambda x: "folder1" in x
|
||||||
)
|
)
|
||||||
assert len(folders) == 1
|
assert len(folders) == 1
|
||||||
|
|
||||||
def test_save_and_load_h5(self):
|
def test_group_by_extension(self):
|
||||||
|
files = ["/path/file1.txt", "/path/file2.txt", "/path/file3.json"]
|
||||||
|
groups = FileScanner.group_by_extension(files)
|
||||||
|
assert ".txt" in groups
|
||||||
|
assert ".json" in groups
|
||||||
|
assert len(groups[".txt"]) == 2
|
||||||
|
assert len(groups[".json"]) == 1
|
||||||
|
|
||||||
|
|
||||||
|
class TestHDF5Handler:
|
||||||
|
def test_save_and_load(self):
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
tensor_group = {
|
tensor_group = {
|
||||||
"sequence": [torch.tensor([1, 2, 3], dtype=torch.int32)],
|
"sequence": [torch.tensor([1, 2, 3], dtype=torch.int32)],
|
||||||
"labels": [torch.tensor([4, 5], dtype=torch.int32)],
|
"labels": [torch.tensor([4, 5], dtype=torch.int32)],
|
||||||
}
|
}
|
||||||
IOHandler.save_h5(tmpdir, "test", tensor_group)
|
h5_path = HDF5Handler.save(tmpdir, "test", tensor_group)
|
||||||
|
|
||||||
assert os.path.exists(os.path.join(tmpdir, "test.h5"))
|
assert os.path.exists(h5_path)
|
||||||
loaded = IOHandler.load_h5(tmpdir, share_memory=False)
|
loaded = HDF5Handler.load(h5_path, share_memory=False)
|
||||||
assert "sequence" in loaded
|
assert "sequence" in loaded
|
||||||
assert "labels" in loaded
|
assert "labels" in loaded
|
||||||
assert torch.equal(
|
assert torch.equal(
|
||||||
@@ -64,14 +83,14 @@ class TestIOHandler:
|
|||||||
loaded["labels"][0], torch.tensor([4, 5], dtype=torch.int32)
|
loaded["labels"][0], torch.tensor([4, 5], dtype=torch.int32)
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_save_h5_creates_directory(self):
|
def test_save_creates_directory(self):
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
output_dir = os.path.join(tmpdir, "nested", "output")
|
output_dir = os.path.join(tmpdir, "nested", "output")
|
||||||
IOHandler.save_h5(output_dir, "test", {"data": [torch.tensor([1, 2, 3])]})
|
HDF5Handler.save(output_dir, "test", {"data": [torch.tensor([1, 2, 3])]})
|
||||||
assert os.path.exists(output_dir)
|
assert os.path.exists(output_dir)
|
||||||
assert os.path.exists(os.path.join(output_dir, "test.h5"))
|
assert os.path.exists(os.path.join(output_dir, "test.h5"))
|
||||||
|
|
||||||
def test_load_h5_multiple_files(self):
|
def test_load_directory_with_multiple_files(self):
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
for i, data in enumerate([[1, 2, 3], [4, 5, 6]]):
|
for i, data in enumerate([[1, 2, 3], [4, 5, 6]]):
|
||||||
h5_path = os.path.join(tmpdir, f"file{i}.h5")
|
h5_path = os.path.join(tmpdir, f"file{i}.h5")
|
||||||
@@ -79,10 +98,10 @@ class TestIOHandler:
|
|||||||
grp = f.create_group("data")
|
grp = f.create_group("data")
|
||||||
grp.create_dataset("data_0", data=data)
|
grp.create_dataset("data_0", data=data)
|
||||||
|
|
||||||
loaded = IOHandler.load_h5(tmpdir, share_memory=False)
|
loaded = HDF5Handler.load(tmpdir, share_memory=False)
|
||||||
assert len(loaded["data"]) == 2
|
assert len(loaded["data"]) == 2
|
||||||
|
|
||||||
def test_load_h5_with_rglob(self):
|
def test_load_directory_with_nested_files(self):
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
subdir = os.path.join(tmpdir, "subdir")
|
subdir = os.path.join(tmpdir, "subdir")
|
||||||
os.makedirs(subdir)
|
os.makedirs(subdir)
|
||||||
@@ -91,11 +110,11 @@ class TestIOHandler:
|
|||||||
grp = f.create_group("test")
|
grp = f.create_group("test")
|
||||||
grp.create_dataset("data_0", data=[1, 2])
|
grp.create_dataset("data_0", data=[1, 2])
|
||||||
|
|
||||||
loaded = IOHandler.load_h5(tmpdir, share_memory=False)
|
loaded = HDF5Handler.load(tmpdir, share_memory=False)
|
||||||
assert "test" in loaded
|
assert "test" in loaded
|
||||||
assert len(loaded["test"]) == 1
|
assert len(loaded["test"]) == 1
|
||||||
|
|
||||||
def test_save_h5_multiple_tensors_per_key(self):
|
def test_save_multiple_tensors_per_key(self):
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
tensor_group = {
|
tensor_group = {
|
||||||
"batch": [
|
"batch": [
|
||||||
@@ -104,6 +123,17 @@ class TestIOHandler:
|
|||||||
torch.tensor([6]),
|
torch.tensor([6]),
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
IOHandler.save_h5(tmpdir, "multi", tensor_group)
|
HDF5Handler.save(tmpdir, "multi", tensor_group)
|
||||||
loaded = IOHandler.load_h5(tmpdir, share_memory=False)
|
loaded = HDF5Handler.load(tmpdir, share_memory=False)
|
||||||
assert len(loaded["batch"]) == 3
|
assert len(loaded["batch"]) == 3
|
||||||
|
|
||||||
|
def test_get_metadata(self):
|
||||||
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
|
tensor_group = {
|
||||||
|
"data": [torch.tensor([1, 2, 3]) for _ in range(5)],
|
||||||
|
}
|
||||||
|
HDF5Handler.save(tmpdir, "meta", tensor_group)
|
||||||
|
|
||||||
|
h5_path = os.path.join(tmpdir, "meta.h5")
|
||||||
|
metadata = HDF5Handler.get_metadata(h5_path)
|
||||||
|
assert metadata["data"] == 5
|
||||||
|
|||||||
Reference in New Issue
Block a user