merge remote main
This commit is contained in:
+12
-2
@@ -28,7 +28,13 @@ Usage::
|
||||
from pipeline.pipeline import Pipeline, PipelineConfig, Stage, TransformStage
|
||||
from pipeline.tokenize import AutoTokenizer, ChatTemplate, train_bpe_tokenizer
|
||||
from pipeline.text import TextNormalizer
|
||||
from pipeline.packing import SequencePacker
|
||||
from pipeline.packing import (
|
||||
GreedyPacker,
|
||||
FfDPacker,
|
||||
BfdPacker,
|
||||
BasePacker,
|
||||
pack_tensors,
|
||||
)
|
||||
|
||||
# I/O module
|
||||
from pipeline.io import FileScanner, HDF5Handler, export_dataset, cache_jsonl
|
||||
@@ -70,7 +76,11 @@ __all__ = [
|
||||
"train_bpe_tokenizer",
|
||||
# Text processing
|
||||
"TextNormalizer",
|
||||
"SequencePacker",
|
||||
"GreedyPacker",
|
||||
"FfDPacker",
|
||||
"BfdPacker",
|
||||
"BasePacker",
|
||||
"pack_tensors",
|
||||
# I/O
|
||||
"FileScanner",
|
||||
"HDF5Handler",
|
||||
|
||||
+11
-1
@@ -4,16 +4,26 @@ 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
|
||||
- cache_jsonl: JSONL to HDF5/binary tokenization and caching
|
||||
- dedup_jsonl: MinHash+LSH deduplication for pretraining text
|
||||
- writers: BaseWriter / H5Writer / BinWriter / TextWriter (Strategy + Factory)
|
||||
"""
|
||||
|
||||
from pipeline.io.file_scanner import FileScanner
|
||||
from pipeline.io.hdf5_handler import HDF5Handler
|
||||
from pipeline.io.export import export_dataset, cache_jsonl
|
||||
from pipeline.io.dedup import dedup_jsonl
|
||||
from pipeline.io.writers import BaseWriter, H5Writer, BinWriter, TextWriter, create_writer
|
||||
|
||||
__all__ = [
|
||||
"FileScanner",
|
||||
"HDF5Handler",
|
||||
"export_dataset",
|
||||
"cache_jsonl",
|
||||
"dedup_jsonl",
|
||||
"BaseWriter",
|
||||
"H5Writer",
|
||||
"BinWriter",
|
||||
"TextWriter",
|
||||
"create_writer",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
"""MinHash + LSH deduplication for pretraining text data."""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Iterator, List, Set, Tuple
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from pipeline.io.writers import TextWriter
|
||||
from pipeline.utils import error_handler
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _tokenize(text: str, ngram: int = 3) -> Set[str]:
|
||||
return {text[i : i + ngram] for i in range(len(text) - ngram + 1)}
|
||||
|
||||
|
||||
def _iter_docs(input_dir: Path) -> Iterator[Tuple[str, dict]]:
|
||||
for fpath in sorted(input_dir.glob("*.jsonl")):
|
||||
with open(fpath, encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
record = json.loads(line)
|
||||
text = record.get("text", "")
|
||||
if text:
|
||||
yield text, record
|
||||
|
||||
|
||||
def _write_h5(records: List[dict], output_dir: str, chunk_idx: int):
|
||||
import h5py
|
||||
|
||||
fname = os.path.join(output_dir, f"chunk_{chunk_idx}.h5")
|
||||
texts = [rec.get("text", "") for rec in records]
|
||||
with h5py.File(fname, "w") as f:
|
||||
dt = h5py.special_dtype(vlen=str)
|
||||
ds = f.create_dataset("text", (len(texts),), dtype=dt)
|
||||
for i, t in enumerate(texts):
|
||||
ds[i] = t
|
||||
|
||||
|
||||
def _write_bin(records: List[dict], output_dir: Path, chunk_idx: int):
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
texts = [rec.get("text", "") + "\n" for rec in records]
|
||||
|
||||
meta = {"chunk": chunk_idx, "count": len(texts), "format": "text", "encoding": "utf-8"}
|
||||
meta_path = output_dir / "meta.json"
|
||||
existing = json.loads(meta_path.read_text()) if meta_path.exists() else {}
|
||||
existing[str(chunk_idx)] = meta
|
||||
meta_path.write_text(json.dumps(existing, indent=2))
|
||||
|
||||
(output_dir / f"text_{chunk_idx}.bin").write_bytes("".join(texts).encode("utf-8"))
|
||||
|
||||
|
||||
_WRITERS = {
|
||||
"jsonl": TextWriter,
|
||||
"h5": lambda: None, # handled inline below
|
||||
"bin": lambda: None,
|
||||
}
|
||||
|
||||
|
||||
@error_handler()
|
||||
def dedup_jsonl(
|
||||
input_dir: str,
|
||||
output_dir: str,
|
||||
*,
|
||||
threshold: float = 0.8,
|
||||
num_perm: int = 128,
|
||||
ngram: int = 3,
|
||||
output_format: str = "jsonl",
|
||||
chunk_size: int = 1_000_000,
|
||||
) -> Tuple[int, int]:
|
||||
"""Deduplicate JSONL text files using MinHash + LSH.
|
||||
|
||||
Args:
|
||||
input_dir: Directory with source ``*.jsonl`` files.
|
||||
output_dir: Directory for deduplicated output.
|
||||
threshold: Jaccard similarity threshold (0–1).
|
||||
num_perm: Number of MinHash permutations.
|
||||
ngram: Character n-gram size.
|
||||
output_format: ``"jsonl"``, ``"h5"``, or ``"bin"``.
|
||||
chunk_size: Records per output chunk file.
|
||||
|
||||
Returns:
|
||||
``(kept, removed)`` counts.
|
||||
"""
|
||||
from datasketch import MinHash, MinHashLSH
|
||||
|
||||
input_path = Path(input_dir)
|
||||
output_path = Path(output_dir)
|
||||
output_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
logger.info(
|
||||
f"Deduplicating {input_dir} -> {output_dir} "
|
||||
f"(threshold={threshold}, perm={num_perm}, fmt={output_format})"
|
||||
)
|
||||
|
||||
lsh = MinHashLSH(threshold=threshold, num_perm=num_perm)
|
||||
|
||||
kept = 0
|
||||
removed = 0
|
||||
|
||||
dup_doc_ids: Set[int] = set()
|
||||
for doc_id, (text, _record) in enumerate(tqdm(_iter_docs(input_path), desc="indexing", unit="docs")):
|
||||
shingles = _tokenize(text, ngram=ngram)
|
||||
if len(shingles) < ngram * 2:
|
||||
dup_doc_ids.add(doc_id)
|
||||
continue
|
||||
|
||||
m = MinHash(num_perm=num_perm)
|
||||
for s in shingles:
|
||||
m.update(s.encode("utf-8"))
|
||||
|
||||
if lsh.query(m):
|
||||
dup_doc_ids.add(doc_id)
|
||||
else:
|
||||
lsh.insert(doc_id, m)
|
||||
|
||||
logger.info(f"Found {len(dup_doc_ids)} duplicates, writing deduplicated data")
|
||||
|
||||
buffer: List[dict] = []
|
||||
chunk_idx = 0
|
||||
writer = TextWriter(chunk_size) if output_format == "jsonl" else None
|
||||
|
||||
for doc_id, (_text, record) in enumerate(tqdm(_iter_docs(input_path), desc="writing", unit="docs")):
|
||||
if doc_id in dup_doc_ids:
|
||||
removed += 1
|
||||
continue
|
||||
|
||||
kept += 1
|
||||
buffer.append(record)
|
||||
|
||||
if len(buffer) >= chunk_size:
|
||||
_flush_chunk(buffer, output_path, chunk_idx, output_format, writer)
|
||||
chunk_idx += 1
|
||||
buffer = []
|
||||
|
||||
if buffer:
|
||||
_flush_chunk(buffer, output_path, chunk_idx, output_format, writer)
|
||||
|
||||
if writer:
|
||||
writer.flush(output_path)
|
||||
|
||||
logger.info(f"Done. kept={kept}, removed={removed}")
|
||||
return kept, removed
|
||||
|
||||
|
||||
def _flush_chunk(
|
||||
records: List[dict],
|
||||
output_dir: Path,
|
||||
chunk_idx: int,
|
||||
output_format: str,
|
||||
writer=None,
|
||||
):
|
||||
if output_format == "jsonl":
|
||||
for rec in records:
|
||||
writer.write_record(rec, output_dir)
|
||||
elif output_format == "h5":
|
||||
_write_h5(records, str(output_dir), chunk_idx)
|
||||
elif output_format == "bin":
|
||||
_write_bin(records, output_dir, chunk_idx)
|
||||
else:
|
||||
raise ValueError(f"Unknown output format: {output_format}")
|
||||
+132
-37
@@ -4,14 +4,18 @@ import json
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from datasets import Dataset
|
||||
from torch import Tensor
|
||||
from tqdm import tqdm
|
||||
|
||||
from pipeline.io.file_scanner import FileScanner
|
||||
from pipeline.io.hdf5_handler import HDF5Handler
|
||||
from pipeline.io.writers import create_writer, BaseWriter
|
||||
from pipeline.processors import BaseProcessor
|
||||
from pipeline.packing import pack_tensors
|
||||
from pipeline.packing import pack_tensors, BasePacker
|
||||
from pipeline.utils import error_handler
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -75,6 +79,32 @@ def export_dataset(
|
||||
return output_files
|
||||
|
||||
|
||||
def merge_tensors(
|
||||
tensors: List[Tensor],
|
||||
group_size: int,
|
||||
) -> List[Tensor]:
|
||||
"""Merge a list of tensors into fewer larger tensors.
|
||||
|
||||
Concatenates every group_size consecutive tensors into one merged
|
||||
tensor. This reduces the number of shm blocks when loading.
|
||||
|
||||
Args:
|
||||
tensors: List of 1D tensors.
|
||||
group_size: Number of tensors to merge into each group.
|
||||
|
||||
Returns:
|
||||
List of merged tensors.
|
||||
"""
|
||||
if not tensors:
|
||||
return []
|
||||
|
||||
merged: List[Tensor] = []
|
||||
for i in range(0, len(tensors), group_size):
|
||||
merged.append(torch.cat(tensors[i : i + group_size]))
|
||||
|
||||
return merged
|
||||
|
||||
|
||||
@error_handler()
|
||||
def cache_jsonl(
|
||||
files: List[str],
|
||||
@@ -83,41 +113,85 @@ def cache_jsonl(
|
||||
*,
|
||||
pack_size: int = -1,
|
||||
pad_value: int = 0,
|
||||
batch_size: int = 256,
|
||||
group_size: int = 1_000,
|
||||
pack_algo: Optional[str] = None,
|
||||
output_format: str = "h5",
|
||||
batch_size: int = 1000,
|
||||
) -> List[str]:
|
||||
"""Tokenize JSONL files and pack them into HDF5 storage.
|
||||
"""Tokenize JSONL files and save as HDF5 or binary.
|
||||
|
||||
BFD packs in group_size-bounded batches to avoid O(N²), then all
|
||||
packed chunks are merged and saved as one file per input file.
|
||||
|
||||
Args:
|
||||
files: List of JSONL file paths.
|
||||
output_dir: H5 output directory.
|
||||
output_dir: 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.
|
||||
group_size: BFD batch granularity (token count threshold for each
|
||||
packing batch) and merge granularity, <=0 means no merging.
|
||||
pack_algo: Packing algorithm: 'bfd' (default), 'ffd',
|
||||
'greedy'. Only used when pack_size > 0.
|
||||
output_format: ``"h5"`` or ``"bin"``.
|
||||
batch_size: Number of lines to batch-process together for parallel
|
||||
tokenization via encode_batch (default: 1000).
|
||||
|
||||
Returns:
|
||||
List of generated H5 file paths.
|
||||
List of generated file paths.
|
||||
"""
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
output_files: List[str] = []
|
||||
output_keys = processor.output_keys
|
||||
|
||||
dtypes = (
|
||||
dict(processor.schema.output_fields)
|
||||
if processor.schema is not None
|
||||
else None
|
||||
)
|
||||
pad_values = {k: (0 if k == "position_ids" else (False if k.endswith("_mask") else pad_value)) for k in output_keys}
|
||||
|
||||
target_tokens = group_size * pack_size if group_size > 0 and pack_size > 0 else 0
|
||||
|
||||
for file_path in files:
|
||||
file_name = Path(file_path).stem
|
||||
|
||||
arrows: Dict[str, List] = {key: [] for key in output_keys}
|
||||
all_packed: Dict[str, List[Tensor]] = {key: [] for key in output_keys}
|
||||
arrows_batch: Dict[str, List] = {key: [] for key in output_keys}
|
||||
batch_tokens: int = 0
|
||||
|
||||
def append_batch(batch):
|
||||
items = [item for _, item in batch]
|
||||
buf: List[Tuple[int, str]] = []
|
||||
|
||||
def flush_buf():
|
||||
nonlocal batch_tokens
|
||||
if not buf:
|
||||
return
|
||||
samples = []
|
||||
for line_num, line in buf:
|
||||
try:
|
||||
samples.append((line_num, json.loads(line)))
|
||||
except json.JSONDecodeError as e:
|
||||
logger.warning(
|
||||
f"JSON decode error in {file_path} line {line_num}: "
|
||||
f"{e}. Skipping line."
|
||||
)
|
||||
buf.clear()
|
||||
if not samples:
|
||||
return
|
||||
items = [item for _, item in samples]
|
||||
try:
|
||||
results = processor.process_batch(items)
|
||||
results = (
|
||||
processor.process_batch(items)
|
||||
if hasattr(processor, "process_batch")
|
||||
else [processor.process(s) for s in 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:
|
||||
for line_num, item in samples:
|
||||
try:
|
||||
results.append(processor.process(item))
|
||||
except Exception as e:
|
||||
@@ -126,42 +200,63 @@ def cache_jsonl(
|
||||
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])
|
||||
arrows_batch[key].append(result[key])
|
||||
if target_tokens > 0:
|
||||
batch_tokens += int(result[output_keys[0]].shape[0])
|
||||
|
||||
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)
|
||||
buf.append((line_num, line))
|
||||
if len(buf) >= batch_size:
|
||||
flush_buf()
|
||||
if target_tokens > 0 and batch_tokens >= target_tokens:
|
||||
packed = pack_tensors(
|
||||
arrows_batch,
|
||||
pack_size,
|
||||
pad_value,
|
||||
dtypes,
|
||||
pad_values=pad_values,
|
||||
algo=pack_algo,
|
||||
)
|
||||
for key in output_keys:
|
||||
all_packed[key].extend(packed[key])
|
||||
arrows_batch[key] = []
|
||||
batch_tokens = 0
|
||||
|
||||
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)
|
||||
flush_buf()
|
||||
|
||||
if arrows_batch[output_keys[0]]:
|
||||
if pack_size > 0:
|
||||
packed = pack_tensors(arrows_batch, pack_size, pad_value, dtypes, pad_values=pad_values, algo=pack_algo)
|
||||
for key in output_keys:
|
||||
all_packed[key].extend(packed[key])
|
||||
else:
|
||||
for key in output_keys:
|
||||
all_packed[key].extend(arrows_batch[key])
|
||||
|
||||
if not all_packed[output_keys[0]]:
|
||||
logger.warning(f"No valid samples in {file_path}, skipping")
|
||||
continue
|
||||
|
||||
if pack_size <= 0:
|
||||
output = all_packed
|
||||
elif group_size > 0 and all_packed[output_keys[0]]:
|
||||
output = {
|
||||
key: merge_tensors(tensors, group_size)
|
||||
for key, tensors in all_packed.items()
|
||||
}
|
||||
else:
|
||||
output = arrows
|
||||
output = all_packed
|
||||
|
||||
h5_path = HDF5Handler.save(output_dir, file_name, output)
|
||||
output_files.append(h5_path)
|
||||
logger.info(f"Saved {h5_path}")
|
||||
writer: BaseWriter = create_writer(output_format)
|
||||
saved = writer.save(output_dir, file_name, output)
|
||||
output_files.append(saved)
|
||||
logger.info(f"Saved {saved}")
|
||||
|
||||
return output_files
|
||||
|
||||
@@ -0,0 +1,107 @@
|
||||
"""Storage backends for tensor / text output (Strategy + Factory).
|
||||
|
||||
Each backend implements a common ``save()`` interface so callers use
|
||||
polymorphism instead of ``if fmt == "h5" ... elif fmt == "bin" ...``.
|
||||
|
||||
Supports:
|
||||
- **H5Writer**: HDF5 format (via HDF5Handler)
|
||||
- **BinWriter**: binary format – meta.json + {key}.bin (memmap-compatible)
|
||||
- **TextWriter**: raw JSONL text (for dedup output)
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
from pathlib import Path
|
||||
from typing import Dict, List
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class BaseWriter(ABC):
|
||||
"""Abstract writer – call ``save(dir, name, data)`` without caring
|
||||
about the underlying format."""
|
||||
|
||||
@abstractmethod
|
||||
def save(self, output_dir: str, file_name: str, data: Dict[str, List[Tensor]]) -> str:
|
||||
...
|
||||
|
||||
|
||||
class H5Writer(BaseWriter):
|
||||
def save(self, output_dir: str, file_name: str, data: Dict[str, List[Tensor]]) -> str:
|
||||
from pipeline.io.hdf5_handler import HDF5Handler
|
||||
return HDF5Handler.save(output_dir, file_name, data)
|
||||
|
||||
|
||||
class BinWriter(BaseWriter):
|
||||
def save(self, output_dir: str, file_name: str, data: Dict[str, List[Tensor]]) -> str:
|
||||
import numpy as np
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
sub_dir = os.path.join(output_dir, file_name)
|
||||
os.makedirs(sub_dir, exist_ok=True)
|
||||
|
||||
meta: Dict[str, Dict] = {}
|
||||
for key, tensors in data.items():
|
||||
cat = torch.cat(tensors, dim=0)
|
||||
meta[key] = {"shape": list(cat.shape), "dtype": str(cat.dtype).split(".")[-1]}
|
||||
np.asarray(cat.cpu().numpy()).tofile(os.path.join(sub_dir, f"{key}.bin"))
|
||||
|
||||
with open(os.path.join(sub_dir, "meta.json"), "w") as f:
|
||||
json.dump(meta, f, indent=2)
|
||||
|
||||
return sub_dir
|
||||
|
||||
|
||||
class TextWriter(BaseWriter):
|
||||
"""Write raw text records as JSONL (used by dedup output)."""
|
||||
|
||||
def __init__(self, chunk_size: int = 1_000_000):
|
||||
self._chunk_size = chunk_size
|
||||
self._buffer: List[dict] = []
|
||||
self._chunk_idx = 0
|
||||
|
||||
def save(self, output_dir: str, file_name: str, data: Dict[str, List[Tensor]]) -> str:
|
||||
raise NotImplementedError("TextWriter.save_one is for tensor data; use write_record()")
|
||||
|
||||
def write_record(self, record: dict, output_dir: Path):
|
||||
self._buffer.append(record)
|
||||
if len(self._buffer) >= self._chunk_size:
|
||||
self._flush(output_dir)
|
||||
|
||||
def flush(self, output_dir: Path):
|
||||
if self._buffer:
|
||||
self._flush(output_dir)
|
||||
|
||||
def _flush(self, output_dir: Path):
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
fpath = output_dir / f"chunk_{self._chunk_idx}.jsonl"
|
||||
with open(fpath, "w", encoding="utf-8") as f:
|
||||
for rec in self._buffer:
|
||||
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
|
||||
self._chunk_idx += 1
|
||||
self._buffer = []
|
||||
|
||||
|
||||
_WRITER_REGISTRY: Dict[str, type] = {}
|
||||
|
||||
|
||||
def register_writer(name: str):
|
||||
def decorator(cls):
|
||||
_WRITER_REGISTRY[name] = cls
|
||||
return cls
|
||||
return decorator
|
||||
|
||||
|
||||
def create_writer(name: str, **kwargs) -> BaseWriter:
|
||||
cls = _WRITER_REGISTRY.get(name)
|
||||
if cls is None:
|
||||
raise ValueError(f"Unknown writer: {name}. Available: {list(_WRITER_REGISTRY)}")
|
||||
return cls(**kwargs)
|
||||
|
||||
|
||||
# Register built-in writers
|
||||
register_writer("h5")(H5Writer)
|
||||
register_writer("bin")(BinWriter)
|
||||
register_writer("jsonl")(TextWriter)
|
||||
@@ -1,146 +0,0 @@
|
||||
import logging
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from pipeline.utils import error_handler
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SequencePacker:
|
||||
"""
|
||||
Stream-concatenation packer for LLM training sequences.
|
||||
|
||||
Algorithm (streaming concat):
|
||||
|
||||
Input: sequences = [A(len=3), B(len=5), C(len=2)], pack_size = 6
|
||||
|
||||
1. Validate & Normalize
|
||||
- Check 1D dimension, unify dtype, warn on overlong sequences
|
||||
- Result: [A, B, C]
|
||||
|
||||
2. Stream into buffer, slice off full chunks
|
||||
- buffer += A(3) -> [a1 a2 a3], pos=3
|
||||
- buffer += B(5) -> [a1 a2 a3 b1 b2 b3 b4 b5], pos=8
|
||||
pos >= 6 -> flush [a1 a2 a3 b1 b2 b3], buffer=[b4 b5], pos=2
|
||||
- buffer += C(2) -> [b4 b5 c1 c2], pos=4
|
||||
loop ends -> flush tail [b4 b5 c1 c2 PAD PAD]
|
||||
|
||||
Output: [[a1 a2 a3 b1 b2 b3], [b4 b5 c1 c2 PAD PAD]]
|
||||
|
||||
Samples may be split across chunks — this is intentional and standard
|
||||
practice in LLM training (TRL, Megatron-LM, etc.).
|
||||
|
||||
Cross-group consistency:
|
||||
Different tensor groups (e.g. input_ids, loss_masks) packed with
|
||||
separate packer instances on samples with matching lengths produce
|
||||
identical chunk boundaries. Element-level correspondence is preserved.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
self.pack_size = pack_size
|
||||
self.pad_value = pad_value
|
||||
self.dtype = dtype
|
||||
self._buffer: List = []
|
||||
self._pos: int = 0
|
||||
self._packages: List[Tensor] = []
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Reset packer state for instance reuse."""
|
||||
self._buffer = []
|
||||
self._pos = 0
|
||||
self._packages = []
|
||||
|
||||
@error_handler()
|
||||
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||
"""
|
||||
Pack sequences via streaming concatenation into fixed-size chunks.
|
||||
|
||||
Sequences are concatenated in order and sliced at pack_size boundaries.
|
||||
The final chunk is padded with pad_value.
|
||||
|
||||
When dtype is not set at init, it is inferred from the first input tensor.
|
||||
|
||||
Args:
|
||||
sequences: List of 1D input tensors.
|
||||
|
||||
Returns:
|
||||
List of packed tensors, each with length equal to pack_size.
|
||||
"""
|
||||
if not sequences:
|
||||
return []
|
||||
|
||||
# --- auto-infer dtype from first sequence ---
|
||||
if self.dtype is None:
|
||||
self.dtype = sequences[0].dtype
|
||||
|
||||
# --- validate & normalize ---
|
||||
normalized: List[Tensor] = []
|
||||
for i, seq in enumerate(sequences):
|
||||
if seq.dim() != 1:
|
||||
raise ValueError(
|
||||
f"Expected 1D tensor at index {i}, got {seq.dim()}D tensor with shape {seq.shape}"
|
||||
)
|
||||
if seq.dtype != self.dtype:
|
||||
seq = seq.to(self.dtype)
|
||||
normalized.append(seq)
|
||||
|
||||
# --- stream into buffer, slice off full chunks ---
|
||||
self._buffer = []
|
||||
self._packages = []
|
||||
pack_size = self.pack_size
|
||||
buf = self._buffer
|
||||
|
||||
for seq in normalized:
|
||||
buf.extend(seq.tolist())
|
||||
while len(buf) >= pack_size:
|
||||
self._packages.append(torch.tensor(buf[:pack_size], dtype=self.dtype))
|
||||
buf = buf[pack_size:]
|
||||
|
||||
# flush tail with padding
|
||||
if buf:
|
||||
padded = buf + [self.pad_value] * (pack_size - len(buf))
|
||||
self._packages.append(torch.tensor(padded, dtype=self.dtype))
|
||||
|
||||
self._pos = len(buf)
|
||||
return self._packages
|
||||
|
||||
|
||||
def pack_tensors(
|
||||
tensors: Dict[str, List[Tensor]],
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtypes: Optional[Dict[str, torch.dtype]] = None,
|
||||
) -> Dict[str, List[Tensor]]:
|
||||
"""
|
||||
Pack multiple named tensor groups in parallel.
|
||||
|
||||
Each group is packed independently with its own SequencePacker instance.
|
||||
When dtypes is provided, packers use the declared dtype per key;
|
||||
otherwise dtype is auto-inferred from the first tensor in each group.
|
||||
|
||||
Args:
|
||||
tensors: Dict mapping key names to lists of 1D tensors.
|
||||
pack_size: Fixed chunk length.
|
||||
pad_value: Padding value for non-bool tensors.
|
||||
dtypes: Optional per-key dtype declarations.
|
||||
|
||||
Returns:
|
||||
Dict mapping key names to lists of packed tensors.
|
||||
"""
|
||||
if dtypes is None:
|
||||
dtypes = {}
|
||||
|
||||
output: Dict[str, List[Tensor]] = {}
|
||||
for key, seqs in tensors.items():
|
||||
dtype = dtypes.get(key)
|
||||
packer = SequencePacker(pack_size, pad_value, dtype=dtype)
|
||||
output[key] = packer.pack(seqs)
|
||||
return output
|
||||
@@ -0,0 +1,84 @@
|
||||
"""Sequence packing algorithms for LLM training data.
|
||||
|
||||
Available packers:
|
||||
- BfdPacker: Best-Fit Decreasing, samples never split (default)
|
||||
- FfDPacker: First-Fit Decreasing, samples never split
|
||||
- GreedyPacker: First-fit in input order, samples never split
|
||||
"""
|
||||
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
import torch
|
||||
|
||||
from pipeline.packing.base import BasePacker
|
||||
from pipeline.packing.binpack import GreedyPacker, FfDPacker, BfdPacker
|
||||
|
||||
|
||||
def pack_tensors(
|
||||
tensors: Dict[str, List[torch.Tensor]],
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtypes: Optional[Dict[str, torch.dtype]] = None,
|
||||
pad_values: Optional[Dict[str, Union[int, bool]]] = None,
|
||||
algo: Optional[Union[str, BasePacker]] = None,
|
||||
) -> Dict[str, List[torch.Tensor]]:
|
||||
"""Pack multiple named tensor groups in parallel.
|
||||
|
||||
Each group is packed independently with its own packer instance.
|
||||
|
||||
Args:
|
||||
tensors: Dict mapping key names to lists of 1D tensors.
|
||||
pack_size: Fixed chunk length.
|
||||
pad_value: Default padding value, used for keys not in pad_values.
|
||||
dtypes: Optional per-key dtype declarations.
|
||||
pad_values: Optional per-key padding values (e.g. pad_token_id for
|
||||
'sequence', False for 'loss_mask', 0 for 'position_ids').
|
||||
algo: Packing algorithm to use. Can be 'bfd' (default),
|
||||
'ffd', 'greedy', or a BasePacker instance.
|
||||
|
||||
Returns:
|
||||
Dict mapping key names to lists of packed tensors.
|
||||
"""
|
||||
if dtypes is None:
|
||||
dtypes = {}
|
||||
if pad_values is None:
|
||||
pad_values = {}
|
||||
|
||||
output: Dict[str, List[torch.Tensor]] = {}
|
||||
for key, seqs in tensors.items():
|
||||
key_pad = pad_values.get(key, pad_value)
|
||||
actual_packer = _resolve_algo(algo, pack_size, key_pad)
|
||||
dtype = dtypes.get(key)
|
||||
if dtype is not None:
|
||||
actual_packer.dtype = dtype
|
||||
output[key] = actual_packer.pack(seqs)
|
||||
return output
|
||||
|
||||
|
||||
def _resolve_algo(
|
||||
algo: Optional[Union[str, BasePacker]],
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool],
|
||||
) -> BasePacker:
|
||||
if algo is None or algo == "bfd":
|
||||
return BfdPacker(pack_size, pad_value)
|
||||
if isinstance(algo, BasePacker):
|
||||
cls = type(algo)
|
||||
return cls(pack_size, pad_value)
|
||||
if algo == "ffd":
|
||||
return FfDPacker(pack_size, pad_value)
|
||||
if algo == "greedy":
|
||||
return GreedyPacker(pack_size, pad_value)
|
||||
raise ValueError(
|
||||
f"Unknown packing algorithm: {algo}. "
|
||||
f"Choose from: bfd, ffd, greedy"
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"BasePacker",
|
||||
"BfdPacker",
|
||||
"FfDPacker",
|
||||
"GreedyPacker",
|
||||
"pack_tensors",
|
||||
]
|
||||
@@ -0,0 +1,49 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class BasePacker(ABC):
|
||||
"""Abstract base class for sequence packing algorithms.
|
||||
|
||||
All packers must implement pack() and reset().
|
||||
pack() takes a list of 1D tensors and returns a list of packed fixed-size tensors.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
self.pack_size = pack_size
|
||||
self.pad_value = pad_value
|
||||
self.dtype = dtype
|
||||
|
||||
@abstractmethod
|
||||
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||
"""Pack sequences into fixed-size chunks."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def reset(self) -> None:
|
||||
"""Reset packer state for instance reuse."""
|
||||
...
|
||||
|
||||
def _validate_and_normalize(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||
"""Validate 1D tensors and unify dtype."""
|
||||
if self.dtype is None and sequences:
|
||||
self.dtype = sequences[0].dtype
|
||||
|
||||
normalized: List[Tensor] = []
|
||||
for i, seq in enumerate(sequences):
|
||||
if seq.dim() != 1:
|
||||
raise ValueError(
|
||||
f"Expected 1D tensor at index {i}, got {seq.dim()}D tensor with shape {seq.shape}"
|
||||
)
|
||||
if seq.dtype != self.dtype:
|
||||
seq = seq.to(self.dtype)
|
||||
normalized.append(seq)
|
||||
return normalized
|
||||
@@ -0,0 +1,174 @@
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from pipeline.packing.base import BasePacker
|
||||
from pipeline.utils import error_handler
|
||||
|
||||
|
||||
def _truncate(tokens: List, max_len: int) -> List:
|
||||
return tokens[:max_len]
|
||||
|
||||
|
||||
def _pad_bin(bin_list: List, target_len: int, pad_value: Union[int, bool], dtype: torch.dtype) -> Tensor:
|
||||
bin_list.extend([pad_value] * (target_len - len(bin_list)))
|
||||
return torch.tensor(bin_list, dtype=dtype)
|
||||
|
||||
|
||||
class GreedyPacker(BasePacker):
|
||||
"""Greedy first-fit packer (no sorting).
|
||||
|
||||
Sequences are packed in input order into the first bin with enough space.
|
||||
Overlong sequences (> pack_size) are truncated to pack_size.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__(pack_size, pad_value, dtype)
|
||||
self._bins: List[List] = []
|
||||
|
||||
def reset(self) -> None:
|
||||
self._bins = []
|
||||
|
||||
@error_handler()
|
||||
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||
if not sequences:
|
||||
return []
|
||||
|
||||
normalized = self._validate_and_normalize(sequences)
|
||||
self._bins = []
|
||||
pack_size = self.pack_size
|
||||
pad_value = self.pad_value
|
||||
|
||||
for seq in normalized:
|
||||
seq_len = int(seq.shape[0])
|
||||
if seq_len > pack_size:
|
||||
self._bins.append(_truncate(seq.tolist(), pack_size))
|
||||
continue
|
||||
placed = False
|
||||
for bin_list in self._bins:
|
||||
if len(bin_list) + seq_len <= pack_size:
|
||||
bin_list.extend(seq.tolist())
|
||||
placed = True
|
||||
break
|
||||
if not placed:
|
||||
self._bins.append(list(seq.tolist()))
|
||||
|
||||
packages: List[Tensor] = []
|
||||
for bin_list in self._bins:
|
||||
packages.append(_pad_bin(bin_list, pack_size, pad_value, self.dtype))
|
||||
|
||||
return packages
|
||||
|
||||
|
||||
class FfDPacker(BasePacker):
|
||||
"""First-Fit Decreasing (FFD) bin-packing packer.
|
||||
|
||||
Sequences are sorted by descending length, then packed into the first
|
||||
bin with enough space. Overlong sequences are truncated to pack_size.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__(pack_size, pad_value, dtype)
|
||||
self._bins: List[List] = []
|
||||
|
||||
def reset(self) -> None:
|
||||
self._bins = []
|
||||
|
||||
@error_handler()
|
||||
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||
if not sequences:
|
||||
return []
|
||||
|
||||
normalized = self._validate_and_normalize(sequences)
|
||||
self._bins = []
|
||||
pack_size = self.pack_size
|
||||
pad_value = self.pad_value
|
||||
|
||||
indexed = [(int(s.shape[0]), s) for s in normalized]
|
||||
indexed.sort(key=lambda x: x[0], reverse=True)
|
||||
|
||||
for seq_len, seq in indexed:
|
||||
if seq_len > pack_size:
|
||||
self._bins.append(_truncate(seq.tolist(), pack_size))
|
||||
continue
|
||||
placed = False
|
||||
for bin_list in self._bins:
|
||||
if len(bin_list) + seq_len <= pack_size:
|
||||
bin_list.extend(seq.tolist())
|
||||
placed = True
|
||||
break
|
||||
if not placed:
|
||||
self._bins.append(list(seq.tolist()))
|
||||
|
||||
packages: List[Tensor] = []
|
||||
for bin_list in self._bins:
|
||||
packages.append(_pad_bin(bin_list, pack_size, pad_value, self.dtype))
|
||||
|
||||
return packages
|
||||
|
||||
|
||||
class BfdPacker(BasePacker):
|
||||
"""Best-Fit Decreasing (BFD) bin-packing packer.
|
||||
|
||||
Sequences are sorted by descending length, then packed into the bin
|
||||
that minimizes remaining space (tightest fit).
|
||||
Overlong sequences are truncated to pack_size.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__(pack_size, pad_value, dtype)
|
||||
self._bins: List[List] = []
|
||||
|
||||
def reset(self) -> None:
|
||||
self._bins = []
|
||||
|
||||
@error_handler()
|
||||
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||
if not sequences:
|
||||
return []
|
||||
|
||||
normalized = self._validate_and_normalize(sequences)
|
||||
self._bins = []
|
||||
pack_size = self.pack_size
|
||||
pad_value = self.pad_value
|
||||
|
||||
indexed = [(int(s.shape[0]), s) for s in normalized]
|
||||
indexed.sort(key=lambda x: x[0], reverse=True)
|
||||
|
||||
for seq_len, seq in indexed:
|
||||
if seq_len > pack_size:
|
||||
self._bins.append(_truncate(seq.tolist(), pack_size))
|
||||
continue
|
||||
best_idx = -1
|
||||
best_remain = pack_size + 1
|
||||
for i, bin_list in enumerate(self._bins):
|
||||
remain = pack_size - len(bin_list)
|
||||
if seq_len <= remain < best_remain:
|
||||
best_remain = remain
|
||||
best_idx = i
|
||||
if best_idx >= 0:
|
||||
self._bins[best_idx].extend(seq.tolist())
|
||||
else:
|
||||
self._bins.append(list(seq.tolist()))
|
||||
|
||||
packages: List[Tensor] = []
|
||||
for bin_list in self._bins:
|
||||
packages.append(_pad_bin(bin_list, pack_size, pad_value, self.dtype))
|
||||
|
||||
return packages
|
||||
@@ -88,6 +88,21 @@ class BaseProcessor(ABC):
|
||||
"""Return list of output tensor key names."""
|
||||
pass
|
||||
|
||||
def process_batch(self, input_dicts: List[Dict[str, Any]]) -> List[Dict[str, Tensor]]:
|
||||
"""Process a batch of input samples.
|
||||
|
||||
Default implementation calls process() for each sample.
|
||||
Subclasses should override for efficient batch processing
|
||||
(e.g., using tokenizer.encode_batch).
|
||||
|
||||
Args:
|
||||
input_dicts: List of input dictionaries.
|
||||
|
||||
Returns:
|
||||
List of output dictionaries mapping output key names to tensors.
|
||||
"""
|
||||
return [self.process(d) for d in input_dicts]
|
||||
|
||||
def validate_input(self, input_dict: Dict[str, Any]) -> None:
|
||||
"""Validate input against schema before processing.
|
||||
|
||||
|
||||
@@ -44,11 +44,11 @@ class PreTrainProcessor(BaseProcessor):
|
||||
return {"sequence": torch.tensor(tokens, dtype=torch.int32)}
|
||||
|
||||
def process_batch(self, input_dicts: List[Dict[str, Any]]) -> List[Dict[str, Tensor]]:
|
||||
texts = [f"{item['text']}{self._eos_token}" for item in input_dicts]
|
||||
encoded = self.tokenizer.encode(texts)
|
||||
texts = [f"{d['text']}{self._eos_token}" for d in input_dicts]
|
||||
batch_tokens = self.tokenizer.encode(texts)
|
||||
return [
|
||||
{"sequence": torch.tensor(tokens, dtype=torch.int32)}
|
||||
for tokens in encoded
|
||||
for tokens in batch_tokens
|
||||
]
|
||||
|
||||
@property
|
||||
|
||||
+72
-134
@@ -7,7 +7,7 @@ from torch import Tensor
|
||||
|
||||
from pipeline.tokenize import AutoTokenizer
|
||||
from pipeline.strategies import PromptStrategy, ChatMLStrategy
|
||||
from pipeline.processors.base import BaseProcessor, ProcessorSchema, encode_with_mask
|
||||
from pipeline.processors.base import BaseProcessor, ProcessorSchema
|
||||
from pipeline.processors.factory import ProcessorFactory
|
||||
|
||||
|
||||
@@ -15,15 +15,15 @@ from pipeline.processors.factory import ProcessorFactory
|
||||
class SFTProcessor(BaseProcessor):
|
||||
"""Supervised fine-tuning data processor.
|
||||
|
||||
Supports two input formats:
|
||||
Input formats:
|
||||
1. messages (recommended):
|
||||
``{"messages": [{"role": "user", "content": "..."},
|
||||
{"role": "assistant", "content": "..."}]}``
|
||||
Multi-turn and system prompts are supported.
|
||||
The tokenizer's ``apply_chat_template`` is used for rendering.
|
||||
Multi-turn and system prompts are supported. Each assistant
|
||||
turn gets ``loss_mask = 1``; all other roles get 0.
|
||||
2. legacy query/response:
|
||||
``{"query": "...", "response": "..."}``
|
||||
Falls back to the configured PromptStrategy (ChatML by default).
|
||||
Internally converted to messages.
|
||||
|
||||
Output schema:
|
||||
- sequence: int32 tensor - Combined token IDs (prompt + response)
|
||||
@@ -63,39 +63,23 @@ class SFTProcessor(BaseProcessor):
|
||||
if "messages" in input_dict:
|
||||
return self._process_messages(input_dict["messages"])
|
||||
if "query" in input_dict and "response" in input_dict:
|
||||
return self._process_legacy(input_dict)
|
||||
return self._process_messages([
|
||||
{"role": "user", "content": input_dict["query"]},
|
||||
{"role": "assistant", "content": input_dict["response"]},
|
||||
])
|
||||
raise KeyError(
|
||||
"Input must contain 'messages' or 'query'/'response' pair"
|
||||
)
|
||||
|
||||
def process_batch(self, input_dicts: List[Dict[str, Any]]) -> List[Dict[str, Tensor]]:
|
||||
results: List[Optional[Dict[str, Tensor]]] = [None] * len(input_dicts)
|
||||
message_indices = [i for i, item in enumerate(input_dicts) if "messages" in item]
|
||||
legacy_indices = [
|
||||
i
|
||||
for i, item in enumerate(input_dicts)
|
||||
if "messages" not in item and "query" in item and "response" in item
|
||||
]
|
||||
if len(message_indices) + len(legacy_indices) != len(input_dicts):
|
||||
raise KeyError("Input must contain 'messages' or 'query'/'response' pair")
|
||||
|
||||
if message_indices:
|
||||
items = [input_dicts[i] for i in message_indices]
|
||||
batch_results = self._process_messages_batch(
|
||||
[item["messages"] for item in items]
|
||||
)
|
||||
for index, result in zip(message_indices, batch_results):
|
||||
results[index] = result
|
||||
|
||||
if legacy_indices:
|
||||
items = [input_dicts[i] for i in legacy_indices]
|
||||
batch_results = self._process_legacy_batch(items)
|
||||
for index, result in zip(legacy_indices, batch_results):
|
||||
results[index] = result
|
||||
|
||||
if any(result is None for result in results):
|
||||
raise RuntimeError("Batch processing did not produce all results")
|
||||
return results
|
||||
def _extract_messages(self, input_dict: Dict[str, Any]) -> Optional[List[Dict[str, str]]]:
|
||||
if "messages" in input_dict:
|
||||
return input_dict["messages"]
|
||||
if "query" in input_dict and "response" in input_dict:
|
||||
return [
|
||||
{"role": "user", "content": input_dict["query"]},
|
||||
{"role": "assistant", "content": input_dict["response"]},
|
||||
]
|
||||
return None
|
||||
|
||||
def _process_messages(self, messages: List[Dict[str, str]]) -> Dict[str, Tensor]:
|
||||
if not messages:
|
||||
@@ -103,118 +87,72 @@ class SFTProcessor(BaseProcessor):
|
||||
if messages[-1]["role"] != "assistant":
|
||||
raise ValueError("Last message must have role 'assistant'")
|
||||
|
||||
last_asst_idx = max(
|
||||
i for i, m in enumerate(messages) if m["role"] == "assistant"
|
||||
)
|
||||
strategy = self.strategy or ChatMLStrategy(self.tokenizer)
|
||||
|
||||
full_text = self.tokenizer.apply_chat_template(
|
||||
messages, tokenize=False, add_generation_prompt=False
|
||||
)
|
||||
full_ids = self.tokenizer.encode(full_text, add_special_tokens=False)
|
||||
prompt, resp = strategy.format_messages(messages)
|
||||
|
||||
prompt_text = self.tokenizer.apply_chat_template(
|
||||
messages[:last_asst_idx],
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
)
|
||||
prompt_ids = self.tokenizer.encode(prompt_text, add_special_tokens=False)
|
||||
|
||||
resp_ids = full_ids[len(prompt_ids) :]
|
||||
if not resp_ids:
|
||||
raise ValueError("Empty assistant response")
|
||||
|
||||
tokens, loss_mask = encode_with_mask(prompt_ids, list(resp_ids))
|
||||
|
||||
if self.max_seq_len and len(tokens) > self.max_seq_len:
|
||||
tokens = tokens[: self.max_seq_len]
|
||||
sequence = torch.tensor(prompt + resp, dtype=torch.int32)
|
||||
loss_mask = torch.zeros(len(sequence), dtype=torch.bool)
|
||||
loss_mask[len(prompt) :] = True
|
||||
if self.max_seq_len and len(sequence) > self.max_seq_len:
|
||||
sequence = sequence[: self.max_seq_len]
|
||||
loss_mask = loss_mask[: self.max_seq_len]
|
||||
|
||||
position_ids = torch.arange(len(tokens), dtype=torch.int32)
|
||||
position_ids = torch.arange(len(sequence), dtype=torch.int32)
|
||||
return {
|
||||
"sequence": tokens,
|
||||
"sequence": sequence,
|
||||
"loss_mask": loss_mask,
|
||||
"position_ids": position_ids,
|
||||
}
|
||||
|
||||
def _process_messages_batch(
|
||||
self, conversations: List[List[Dict[str, str]]]
|
||||
) -> List[Dict[str, Tensor]]:
|
||||
for messages in conversations:
|
||||
if not messages:
|
||||
raise ValueError("Messages list is empty")
|
||||
if messages[-1]["role"] != "assistant":
|
||||
raise ValueError("Last message must have role 'assistant'")
|
||||
def process_batch(self, input_dicts: List[Dict[str, Any]]) -> List[Optional[Dict[str, Tensor]]]:
|
||||
strategy = self.strategy or ChatMLStrategy(self.tokenizer)
|
||||
|
||||
assistant_indices = [
|
||||
max(i for i, message in enumerate(messages) if message["role"] == "assistant")
|
||||
for messages in conversations
|
||||
]
|
||||
full_texts = [
|
||||
self.tokenizer.apply_chat_template(
|
||||
messages, tokenize=False, add_generation_prompt=False
|
||||
)
|
||||
for messages in conversations
|
||||
]
|
||||
prompt_texts = [
|
||||
self.tokenizer.apply_chat_template(
|
||||
messages[:assistant_idx], tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
for messages, assistant_idx in zip(conversations, assistant_indices)
|
||||
]
|
||||
full_ids_batch = self.tokenizer.encode(full_texts, add_special_tokens=False)
|
||||
prompt_ids_batch = self.tokenizer.encode(prompt_texts, add_special_tokens=False)
|
||||
prompts_text: List[str] = []
|
||||
fulls_text: List[str] = []
|
||||
indices: List[int] = []
|
||||
results: List[Optional[Dict[str, Tensor]]] = [None] * len(input_dicts)
|
||||
|
||||
results = []
|
||||
for full_ids, prompt_ids in zip(full_ids_batch, prompt_ids_batch):
|
||||
resp_ids = full_ids[len(prompt_ids) :]
|
||||
if not resp_ids:
|
||||
raise ValueError("Empty assistant response")
|
||||
tokens, loss_mask = encode_with_mask(prompt_ids, list(resp_ids))
|
||||
if self.max_seq_len and len(tokens) > self.max_seq_len:
|
||||
tokens = tokens[: self.max_seq_len]
|
||||
for i, d in enumerate(input_dicts):
|
||||
try:
|
||||
messages = self._extract_messages(d)
|
||||
if not messages or messages[-1]["role"] != "assistant":
|
||||
continue
|
||||
last_asst = max(j for j, m in enumerate(messages) if m["role"] == "assistant")
|
||||
prompt_text = self.tokenizer.apply_chat_template(
|
||||
messages[:last_asst], add_generation_prompt=True, tokenize=False
|
||||
)
|
||||
full_text = self.tokenizer.apply_chat_template(
|
||||
messages[: last_asst + 1], add_generation_prompt=False, tokenize=False
|
||||
)
|
||||
prompts_text.append(prompt_text)
|
||||
fulls_text.append(full_text)
|
||||
indices.append(i)
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
if not prompts_text:
|
||||
return results
|
||||
|
||||
prompt_tokens_list = self.tokenizer.encode(prompts_text)
|
||||
full_tokens_list = self.tokenizer.encode(fulls_text)
|
||||
|
||||
for j, idx in enumerate(indices):
|
||||
prompt_tokens = prompt_tokens_list[j]
|
||||
full_tokens = full_tokens_list[j]
|
||||
resp_tokens = full_tokens[len(prompt_tokens):]
|
||||
sequence = torch.tensor(prompt_tokens + resp_tokens, dtype=torch.int32)
|
||||
loss_mask = torch.zeros(len(sequence), dtype=torch.bool)
|
||||
loss_mask[len(prompt_tokens):] = True
|
||||
if self.max_seq_len and len(sequence) > self.max_seq_len:
|
||||
sequence = sequence[: self.max_seq_len]
|
||||
loss_mask = loss_mask[: self.max_seq_len]
|
||||
results.append(
|
||||
{
|
||||
"sequence": tokens,
|
||||
"loss_mask": loss_mask,
|
||||
"position_ids": torch.arange(len(tokens), dtype=torch.int32),
|
||||
}
|
||||
)
|
||||
return results
|
||||
position_ids = torch.arange(len(sequence), dtype=torch.int32)
|
||||
results[idx] = {
|
||||
"sequence": sequence,
|
||||
"loss_mask": loss_mask,
|
||||
"position_ids": position_ids,
|
||||
}
|
||||
|
||||
def _process_legacy(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
|
||||
strategy = self.strategy or ChatMLStrategy(self.tokenizer)
|
||||
|
||||
query_tokens = self.tokenizer.encode(input_dict["query"])
|
||||
response_tokens = self.tokenizer.encode(input_dict["response"])
|
||||
|
||||
prompt = strategy.assemble_prompt(query_tokens)
|
||||
response = strategy.assemble_response(response_tokens)
|
||||
|
||||
tokens, loss_mask = encode_with_mask(prompt, response)
|
||||
position_ids = torch.arange(len(tokens), dtype=torch.int32)
|
||||
return {"sequence": tokens, "loss_mask": loss_mask, "position_ids": position_ids}
|
||||
|
||||
def _process_legacy_batch(
|
||||
self, input_dicts: List[Dict[str, Any]]
|
||||
) -> List[Dict[str, Tensor]]:
|
||||
strategy = self.strategy or ChatMLStrategy(self.tokenizer)
|
||||
query_batch = self.tokenizer.encode([item["query"] for item in input_dicts])
|
||||
response_batch = self.tokenizer.encode(
|
||||
[item["response"] for item in input_dicts]
|
||||
)
|
||||
results = []
|
||||
for query_tokens, response_tokens in zip(query_batch, response_batch):
|
||||
prompt = strategy.assemble_prompt(query_tokens)
|
||||
response = strategy.assemble_response(response_tokens)
|
||||
tokens, loss_mask = encode_with_mask(prompt, response)
|
||||
results.append(
|
||||
{
|
||||
"sequence": tokens,
|
||||
"loss_mask": loss_mask,
|
||||
"position_ids": torch.arange(len(tokens), dtype=torch.int32),
|
||||
}
|
||||
)
|
||||
return results
|
||||
|
||||
@property
|
||||
|
||||
@@ -1,42 +1,91 @@
|
||||
"""ChatML format strategy."""
|
||||
|
||||
from typing import List
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
from pipeline.tokenize import AutoTokenizer
|
||||
from pipeline.strategies.base import PromptStrategy
|
||||
from pipeline.strategies.factory import StrategyFactory
|
||||
|
||||
DEFAULT_CHATML_TEMPLATE = (
|
||||
"{% for message in messages %}"
|
||||
"{% if message['role'] == 'system' %}"
|
||||
"{{ '<|im_start|>system\n' + message['content'] + '<|im_end|>\n' }}"
|
||||
"{% elif message['role'] == 'user' %}"
|
||||
"{{ '<|im_start|>user\n' + message['content'] + '<|im_end|>\n' }}"
|
||||
"{% elif message['role'] == 'assistant' %}"
|
||||
"{{ '<|im_start|>assistant\n' + message['content'] + '<|im_end|>\n' }}"
|
||||
"{% endif %}"
|
||||
"{% endfor %}"
|
||||
"{% if add_generation_prompt %}"
|
||||
"{{ '<|im_start|>assistant\n' }}"
|
||||
"{% endif %}"
|
||||
)
|
||||
|
||||
|
||||
@StrategyFactory.register("chatml")
|
||||
class ChatMLStrategy(PromptStrategy):
|
||||
"""ChatML format strategy."""
|
||||
"""ChatML format strategy.
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: AutoTokenizer,
|
||||
user_start: str = "<|im▁start|>user\n",
|
||||
user_end: str = "<|im▁end|>\n",
|
||||
assistant_start: str = "<|im▁start|>assistant\n",
|
||||
assistant_end: str = "<|im▁end|>\n",
|
||||
):
|
||||
Renders messages using the tokenizer's jinja chat_template from
|
||||
``tokenizer_config.json``. Falls back to DEFAULT_CHATML_TEMPLATE
|
||||
when no template is configured.
|
||||
|
||||
The strategy does **not** hard-code any special tokens – all
|
||||
formatting is driven by the jinja template.
|
||||
"""
|
||||
|
||||
def __init__(self, tokenizer: AutoTokenizer):
|
||||
super().__init__(tokenizer)
|
||||
|
||||
self._user_start_ids = self._encode_format(user_start)
|
||||
self._user_end_ids = self._encode_format(user_end)
|
||||
self._assistant_start_ids = self._encode_format(assistant_start)
|
||||
self._assistant_end_ids = self._encode_format(assistant_end)
|
||||
if tokenizer._chat_template is None:
|
||||
tokenizer.set_chat_template(DEFAULT_CHATML_TEMPLATE)
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return "chatml"
|
||||
|
||||
def format_messages(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
) -> Tuple[List[int], List[int]]:
|
||||
"""Render a single-turn messages conversation.
|
||||
|
||||
Returns ``(prompt_tokens, response_tokens)`` where
|
||||
*prompt_tokens* contains everything up to (and including) the
|
||||
last assistant start marker, and *response_tokens* is the
|
||||
assistant content plus the closing markers.
|
||||
"""
|
||||
last_asst = max(
|
||||
i for i, m in enumerate(messages) if m["role"] == "assistant"
|
||||
)
|
||||
|
||||
prompt = self.tokenizer.apply_chat_template(
|
||||
messages[:last_asst],
|
||||
add_generation_prompt=True,
|
||||
tokenize=True,
|
||||
)
|
||||
full = self.tokenizer.apply_chat_template(
|
||||
messages[: last_asst + 1],
|
||||
add_generation_prompt=False,
|
||||
tokenize=True,
|
||||
)
|
||||
return prompt, full[len(prompt) :]
|
||||
|
||||
def assemble_prompt(self, query_tokens: List[int]) -> List[int]:
|
||||
return (
|
||||
self._user_start_ids
|
||||
+ query_tokens
|
||||
+ self._user_end_ids
|
||||
+ self._assistant_start_ids
|
||||
text = self.tokenizer.decode(query_tokens)
|
||||
return self.tokenizer.apply_chat_template(
|
||||
[{"role": "user", "content": text}],
|
||||
add_generation_prompt=True,
|
||||
tokenize=True,
|
||||
)
|
||||
|
||||
def assemble_response(self, response_tokens: List[int]) -> List[int]:
|
||||
return response_tokens + self._assistant_end_ids
|
||||
text = self.tokenizer.decode(response_tokens)
|
||||
full = self.tokenizer.apply_chat_template(
|
||||
[{"role": "assistant", "content": text}],
|
||||
add_generation_prompt=False,
|
||||
tokenize=True,
|
||||
)
|
||||
opening = self.tokenizer.apply_chat_template(
|
||||
[], add_generation_prompt=True, tokenize=True
|
||||
)
|
||||
return full[len(opening) :]
|
||||
|
||||
@@ -244,22 +244,18 @@ class AutoTokenizer:
|
||||
"Tokenizer not initialized. Load or create a tokenizer first."
|
||||
)
|
||||
|
||||
if isinstance(tokens, str):
|
||||
encoded = self._tokenizer.encode(
|
||||
tokens,
|
||||
is_pretokenized=is_pretokenized,
|
||||
add_special_tokens=add_special_tokens,
|
||||
)
|
||||
return encoded.ids if out_ids else encoded.tokens
|
||||
else:
|
||||
encoded_list = self._tokenizer.encode_batch(
|
||||
tokens,
|
||||
is_pretokenized=is_pretokenized,
|
||||
add_special_tokens=add_special_tokens,
|
||||
)
|
||||
return [
|
||||
encoded.ids if out_ids else encoded.tokens for encoded in encoded_list
|
||||
]
|
||||
single = isinstance(tokens, str)
|
||||
if single:
|
||||
tokens = [tokens]
|
||||
encoded_list = self._tokenizer.encode_batch(
|
||||
tokens,
|
||||
is_pretokenized=is_pretokenized,
|
||||
add_special_tokens=add_special_tokens,
|
||||
)
|
||||
result = [
|
||||
encoded.ids if out_ids else encoded.tokens for encoded in encoded_list
|
||||
]
|
||||
return result[0] if single else result
|
||||
|
||||
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
|
||||
"""Decode token IDs to text."""
|
||||
@@ -270,6 +266,12 @@ class AutoTokenizer:
|
||||
|
||||
return self._tokenizer.decode(tokens, skip_special_tokens=skip_special_tokens)
|
||||
|
||||
def token_to_id(self, token: str) -> Optional[int]:
|
||||
"""Convert a token string to its integer ID."""
|
||||
if self._tokenizer is None:
|
||||
raise RuntimeError("Tokenizer not initialized.")
|
||||
return self._tokenizer.token_to_id(token)
|
||||
|
||||
def __len__(self) -> int:
|
||||
if self._tokenizer is None:
|
||||
return 0
|
||||
|
||||
Reference in New Issue
Block a user