fix: 修复特殊token 问题
This commit is contained in:
+18
-13
@@ -5,24 +5,29 @@ from pipeline.packing import SequencePacker
|
||||
from pipeline.io import IOHandler, export_dataset, cache_jsonl
|
||||
from pipeline.processors import ProcessorFactory, BaseProcessor
|
||||
from pipeline.utils import setup_logging
|
||||
from pipeline.strategies import PromptStrategy, ChatMLStrategy, AlpacaStrategy, StrategyFactory
|
||||
from pipeline.strategies import (
|
||||
PromptStrategy,
|
||||
ChatMLStrategy,
|
||||
AlpacaStrategy,
|
||||
StrategyFactory,
|
||||
)
|
||||
|
||||
# Configure project-level logging
|
||||
setup_logging()
|
||||
|
||||
__all__ = [
|
||||
# Core modules
|
||||
'BpeTokenizer',
|
||||
'TextNormalizer',
|
||||
'SequencePacker',
|
||||
'IOHandler',
|
||||
'ProcessorFactory',
|
||||
'BaseProcessor',
|
||||
'export_dataset',
|
||||
'cache_jsonl',
|
||||
"BpeTokenizer",
|
||||
"TextNormalizer",
|
||||
"SequencePacker",
|
||||
"IOHandler",
|
||||
"ProcessorFactory",
|
||||
"BaseProcessor",
|
||||
"export_dataset",
|
||||
"cache_jsonl",
|
||||
# Strategy pattern
|
||||
'PromptStrategy',
|
||||
'ChatMLStrategy',
|
||||
'AlpacaStrategy',
|
||||
'StrategyFactory',
|
||||
"PromptStrategy",
|
||||
"ChatMLStrategy",
|
||||
"AlpacaStrategy",
|
||||
"StrategyFactory",
|
||||
]
|
||||
|
||||
+27
-10
@@ -1,4 +1,5 @@
|
||||
"""File, HDF5, JSONL I/O operations."""
|
||||
|
||||
import json
|
||||
import os
|
||||
import logging
|
||||
@@ -33,7 +34,9 @@ class IOHandler:
|
||||
return sorted(files)
|
||||
|
||||
@staticmethod
|
||||
def fetch_folders(root_dir: str, filter_func: Optional[Callable[[str], bool]] = None) -> List[str]:
|
||||
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:
|
||||
@@ -44,15 +47,17 @@ class IOHandler:
|
||||
|
||||
@staticmethod
|
||||
@error_handler()
|
||||
def save_h5(output_dir: str, file_name: str, tensor_group: Dict[str, List[Tensor]]) -> None:
|
||||
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:
|
||||
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())
|
||||
grp.create_dataset(f"data_{idx}", data=tensor.cpu().numpy())
|
||||
|
||||
@staticmethod
|
||||
@error_handler()
|
||||
@@ -63,7 +68,7 @@ class IOHandler:
|
||||
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:
|
||||
with h5py.File(h5_file, "r") as f:
|
||||
for key in f.keys():
|
||||
grp = f[key]
|
||||
dsets = []
|
||||
@@ -92,7 +97,9 @@ def export_dataset(
|
||||
*,
|
||||
chunk_size: int = 1_000_000,
|
||||
max_chunks: Optional[int] = None,
|
||||
process_func: Optional[Callable[[Dict[str, Any]], Union[Dict[str, Any], List[Dict[str, Any]]]]] = None,
|
||||
process_func: Optional[
|
||||
Callable[[Dict[str, Any]], Union[Dict[str, Any], List[Dict[str, Any]]]]
|
||||
] = None,
|
||||
column: str = "text",
|
||||
) -> List[str]:
|
||||
"""
|
||||
@@ -125,7 +132,11 @@ def export_dataset(
|
||||
try:
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
for example in chunk:
|
||||
processed = process_func(example) if process_func else {column: example[column]}
|
||||
processed = (
|
||||
process_func(example)
|
||||
if process_func
|
||||
else {column: example[column]}
|
||||
)
|
||||
items = processed if isinstance(processed, list) else [processed]
|
||||
for item in items:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||
@@ -172,17 +183,23 @@ def cache_jsonl(
|
||||
arrows: Dict[str, List] = {key: [] for key in output_keys}
|
||||
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
for line_num, line in enumerate(tqdm(f, desc=f"Processing {file_name}", leave=False), start=1):
|
||||
for line_num, line in enumerate(
|
||||
tqdm(f, desc=f"Processing {file_name}", leave=False), start=1
|
||||
):
|
||||
try:
|
||||
result = processor.process(json.loads(line))
|
||||
if result is not None:
|
||||
for key in output_keys:
|
||||
arrows[key].append(result[key])
|
||||
except json.JSONDecodeError as e:
|
||||
logger.warning(f"JSON decode error in {file_path} line {line_num}: {e}. Skipping line.")
|
||||
logger.warning(
|
||||
f"JSON decode error in {file_path} line {line_num}: {e}. Skipping line."
|
||||
)
|
||||
continue
|
||||
except Exception as e:
|
||||
logger.warning(f"Unexpected error processing line {line_num} in {file_path}: {e}. Skipping line.")
|
||||
logger.warning(
|
||||
f"Unexpected error processing line {line_num} in {file_path}: {e}. Skipping line."
|
||||
)
|
||||
continue
|
||||
|
||||
if pack_size > 0:
|
||||
|
||||
+3
-1
@@ -37,7 +37,9 @@ class SequencePacker:
|
||||
identical chunk boundaries. Element-level correspondence is preserved.
|
||||
"""
|
||||
|
||||
def __init__(self, pack_size: int, pad_value: int = 0, dtype: torch.dtype = torch.int32):
|
||||
def __init__(
|
||||
self, pack_size: int, pad_value: int = 0, dtype: torch.dtype = torch.int32
|
||||
):
|
||||
self.pack_size = pack_size
|
||||
self.pad_value = pad_value
|
||||
self.dtype = dtype
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
Processor classes are registered at definition time via decorators and
|
||||
can be created through :class:`ProcessorFactory`.
|
||||
"""
|
||||
|
||||
from pipeline.processors.base import BaseProcessor
|
||||
from pipeline.processors.factory import ProcessorFactory
|
||||
from pipeline.processors.pretrain import PreTrainProcessor
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""Processor base class and shared utilities."""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, List, Any, Tuple
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""DPO preference learning data processor."""
|
||||
|
||||
from typing import Dict, List, Any, Optional
|
||||
|
||||
import torch
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""Factory for creating and registering processors."""
|
||||
|
||||
from typing import Dict, List, Any, Optional, Type
|
||||
|
||||
from pipeline.processors.base import BaseProcessor
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""Pre-training data processor."""
|
||||
|
||||
from typing import Dict, List, Any
|
||||
|
||||
import torch
|
||||
@@ -18,7 +19,7 @@ class PreTrainProcessor(BaseProcessor):
|
||||
|
||||
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
|
||||
segment = input_dict["text"]
|
||||
tokens = self.tokenizer.encode(f"{segment}<eos>")
|
||||
tokens = self.tokenizer.encode(f"{segment}<|end▁of▁sentence|>")
|
||||
return {"sequence": torch.tensor(tokens, dtype=torch.int32)}
|
||||
|
||||
@property
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""Supervised fine-tuning data processor."""
|
||||
|
||||
from typing import Dict, List, Any, Optional
|
||||
|
||||
import torch
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""Strategy pattern for prompt/response format abstraction."""
|
||||
|
||||
from pipeline.strategies.base import PromptStrategy
|
||||
from pipeline.strategies.factory import StrategyFactory
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""Alpaca format strategy."""
|
||||
|
||||
from typing import List
|
||||
|
||||
from pipeline.tokenizer import BpeTokenizer
|
||||
@@ -8,14 +9,14 @@ from pipeline.strategies.factory import StrategyFactory
|
||||
|
||||
@StrategyFactory.register("alpaca")
|
||||
class AlpacaStrategy(PromptStrategy):
|
||||
"""Alpaca format: ``### Instruction: ... \\n\\n### Response: ... <eos>``"""
|
||||
"""Alpaca format:"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: BpeTokenizer,
|
||||
instruction_start: str = "### Instruction:\n",
|
||||
response_start: str = "### Response:\n",
|
||||
response_suffix: str = "\n<eos>",
|
||||
response_suffix: str = "\n<|end▁of▁sentence|>",
|
||||
):
|
||||
super().__init__(tokenizer)
|
||||
self.instruction_start = instruction_start
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""Abstract base class for prompt construction strategies."""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List
|
||||
|
||||
@@ -30,7 +31,7 @@ class PromptStrategy(ABC):
|
||||
"""Assemble query tokens into a complete prompt with format tokens.
|
||||
|
||||
The prompt includes all tokens up to (and including) the response
|
||||
start marker, e.g. ``<|im_start|>assistant\n``.
|
||||
start marker, e.g. ``<|im▁start|>assistant\n``.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""ChatML format strategy."""
|
||||
|
||||
from typing import List
|
||||
|
||||
from pipeline.tokenizer import BpeTokenizer
|
||||
@@ -8,15 +9,15 @@ from pipeline.strategies.factory import StrategyFactory
|
||||
|
||||
@StrategyFactory.register("chatml")
|
||||
class ChatMLStrategy(PromptStrategy):
|
||||
"""ChatML format: ``<|im_start|>user ... <|im_end|> <|im_start|>assistant ... <|im_end|> <eos>``"""
|
||||
"""ChatML format strategy."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: BpeTokenizer,
|
||||
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<eos>",
|
||||
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",
|
||||
):
|
||||
super().__init__(tokenizer)
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""Factory for creating and registering prompt strategies."""
|
||||
|
||||
from typing import Dict, List, Type
|
||||
|
||||
from pipeline.tokenizer import BpeTokenizer
|
||||
|
||||
+15
-6
@@ -6,16 +6,25 @@ class TextNormalizer:
|
||||
"""Text normalization."""
|
||||
|
||||
DEFAULT_REPLACEMENTS = {
|
||||
"\\[": "$$", "\\]": "$$", "\\(": "$", "\\)": "$",
|
||||
'\u2018': "'", '\u2019': "'", '\u0060': "'",
|
||||
'\u201C': '"', '\u201D': '"',
|
||||
'\u2013': '-', '\u2014': '--', '\u2212': '-',
|
||||
'\u00A0': ' ', '\u2026': '...'
|
||||
"\\[": "$$",
|
||||
"\\]": "$$",
|
||||
"\\(": "$",
|
||||
"\\)": "$",
|
||||
"\u2018": "'",
|
||||
"\u2019": "'",
|
||||
"\u0060": "'",
|
||||
"\u201c": '"',
|
||||
"\u201d": '"',
|
||||
"\u2013": "-",
|
||||
"\u2014": "--",
|
||||
"\u2212": "-",
|
||||
"\u00a0": " ",
|
||||
"\u2026": "...",
|
||||
}
|
||||
|
||||
def __init__(self, custom_rules: Optional[Dict[str, str]] = None):
|
||||
self.replacements = {**self.DEFAULT_REPLACEMENTS, **(custom_rules or {})}
|
||||
self._pattern = re.compile('|'.join(re.escape(k) for k in self.replacements))
|
||||
self._pattern = re.compile("|".join(re.escape(k) for k in self.replacements))
|
||||
|
||||
def normalize(self, text: str) -> str:
|
||||
return self._pattern.sub(lambda m: self.replacements[m.group()], text)
|
||||
|
||||
+87
-44
@@ -7,100 +7,143 @@ from typing import List, Union, Optional, Tuple, Iterator
|
||||
|
||||
class BpeTokenizer:
|
||||
def __init__(self, path: Optional[str] = None):
|
||||
self._control_tokens = ["<bos>", "<eos>", "<pad>"]
|
||||
self._special_tokens = ["<|im_start|>", "<|im_end|>"]
|
||||
|
||||
self._control_tokens = [
|
||||
"<|begin▁of▁sentence|>",
|
||||
"<|end▁of▁sentence|>",
|
||||
"<|▁pad▁|>",
|
||||
]
|
||||
self._special_tokens = ["<|im▁start|>", "<|im▁end|>"]
|
||||
|
||||
model = BPE()
|
||||
self._tokenizer = Tokenizer(model)
|
||||
self._tokenizer.normalizer = normalizers.Sequence([
|
||||
normalizers.NFC(),
|
||||
normalizers.Strip()
|
||||
])
|
||||
|
||||
self._tokenizer.pre_tokenizer = pre_tokenizers.Sequence([
|
||||
pre_tokenizers.UnicodeScripts(),
|
||||
pre_tokenizers.ByteLevel(add_prefix_space=False, use_regex=True)
|
||||
])
|
||||
|
||||
self._tokenizer.normalizer = normalizers.Sequence(
|
||||
[normalizers.NFC(), normalizers.Strip()]
|
||||
)
|
||||
|
||||
self._tokenizer.pre_tokenizer = pre_tokenizers.Sequence(
|
||||
[
|
||||
pre_tokenizers.UnicodeScripts(),
|
||||
pre_tokenizers.ByteLevel(add_prefix_space=False, use_regex=True),
|
||||
]
|
||||
)
|
||||
|
||||
self._tokenizer.decoder = decoders.ByteLevel()
|
||||
self._tokenizer.post_processor = processors.ByteLevel(trim_offsets=True)
|
||||
|
||||
|
||||
if path is not None:
|
||||
self._tokenizer = Tokenizer.from_file(path)
|
||||
|
||||
def _prepare_trainer(self, vocab_size: int, min_freq: int, reserved_token_size: int, max_token_length: int = 18) -> Tuple[BpeTrainer, int, List[str]]:
|
||||
|
||||
def _prepare_trainer(
|
||||
self,
|
||||
vocab_size: int,
|
||||
min_freq: int,
|
||||
reserved_token_size: int,
|
||||
max_token_length: int = 18,
|
||||
) -> Tuple[BpeTrainer, int, List[str]]:
|
||||
assert reserved_token_size > len(self._special_tokens)
|
||||
reserved_tokens = [f"<|reserve{i:02d}|>" for i in range(reserved_token_size - len(self._special_tokens))]
|
||||
detail_vocab_size = vocab_size - (len(reserved_tokens) + len(self._special_tokens))
|
||||
|
||||
reserved_tokens = [
|
||||
f"<|reserve{i:02d}|>"
|
||||
for i in range(reserved_token_size - len(self._special_tokens))
|
||||
]
|
||||
detail_vocab_size = vocab_size - (
|
||||
len(reserved_tokens) + len(self._special_tokens)
|
||||
)
|
||||
|
||||
alphabet = pre_tokenizers.ByteLevel.alphabet()
|
||||
min_size = len(alphabet) + len(self._control_tokens)
|
||||
assert detail_vocab_size > min_size
|
||||
|
||||
|
||||
trainer = BpeTrainer(
|
||||
vocab_size=detail_vocab_size,
|
||||
min_frequency=min_freq,
|
||||
limit_alphabet=detail_vocab_size // 6,
|
||||
max_token_length=max_token_length,
|
||||
special_tokens=self._control_tokens,
|
||||
special_tokens=self._control_tokens + self._special_tokens,
|
||||
initial_alphabet=alphabet,
|
||||
show_progress=True,
|
||||
)
|
||||
|
||||
|
||||
return trainer, detail_vocab_size, reserved_tokens
|
||||
|
||||
def train(self, files: List[str], vocab_size: int, min_freq: int, reserved_token_size: int = 100) -> None:
|
||||
def train(
|
||||
self,
|
||||
files: List[str],
|
||||
vocab_size: int,
|
||||
min_freq: int,
|
||||
reserved_token_size: int = 100,
|
||||
) -> None:
|
||||
trainer, _, reserved_tokens = self._prepare_trainer(
|
||||
vocab_size=vocab_size,
|
||||
min_freq=min_freq,
|
||||
reserved_token_size=reserved_token_size
|
||||
reserved_token_size=reserved_token_size,
|
||||
)
|
||||
self._tokenizer.train(files=files, trainer=trainer)
|
||||
self._tokenizer.add_special_tokens(self._special_tokens + reserved_tokens)
|
||||
|
||||
def train_from_iterator(self, iterator: Iterator[str], vocab_size: int, min_freq: int, reserved_token_size: int = 100) -> None:
|
||||
self._tokenizer.add_special_tokens(
|
||||
self._control_tokens + self._special_tokens + reserved_tokens
|
||||
)
|
||||
|
||||
def train_from_iterator(
|
||||
self,
|
||||
iterator: Iterator[str],
|
||||
vocab_size: int,
|
||||
min_freq: int,
|
||||
reserved_token_size: int = 100,
|
||||
) -> None:
|
||||
trainer, _, reserved_tokens = self._prepare_trainer(
|
||||
vocab_size=vocab_size,
|
||||
min_freq=min_freq,
|
||||
reserved_token_size=reserved_token_size
|
||||
reserved_token_size=reserved_token_size,
|
||||
)
|
||||
self._tokenizer.train_from_iterator(iterator=iterator, trainer=trainer)
|
||||
self._tokenizer.add_special_tokens(self._special_tokens + reserved_tokens)
|
||||
|
||||
self._tokenizer.add_special_tokens(
|
||||
self._control_tokens + self._special_tokens + reserved_tokens
|
||||
)
|
||||
|
||||
def save(self, path: str) -> None:
|
||||
self._tokenizer.save(path)
|
||||
|
||||
|
||||
def load(self, path: str) -> None:
|
||||
self._tokenizer = Tokenizer.from_file(path)
|
||||
|
||||
def encode(self, tokens: Union[str, List[str]], out_ids: bool = True, add_special_tokens: bool = False) -> Union[List[int], List[str], List[List[int]], List[List[str]]]:
|
||||
def encode(
|
||||
self,
|
||||
tokens: Union[str, List[str]],
|
||||
out_ids: bool = True,
|
||||
add_special_tokens: bool = False,
|
||||
) -> Union[List[int], List[str], List[List[int]], List[List[str]]]:
|
||||
if isinstance(tokens, str):
|
||||
encoded: Encoding = self._tokenizer.encode(tokens, add_special_tokens=add_special_tokens)
|
||||
encoded: Encoding = self._tokenizer.encode(
|
||||
tokens, add_special_tokens=add_special_tokens
|
||||
)
|
||||
return encoded.ids if out_ids else encoded.tokens
|
||||
elif isinstance(tokens, list):
|
||||
encoded_list: List[Encoding] = self._tokenizer.encode_batch(tokens, add_special_tokens=add_special_tokens)
|
||||
return [encoded.ids if out_ids else encoded.tokens for encoded in encoded_list]
|
||||
encoded_list: List[Encoding] = self._tokenizer.encode_batch(
|
||||
tokens, add_special_tokens=add_special_tokens
|
||||
)
|
||||
return [
|
||||
encoded.ids if out_ids else encoded.tokens for encoded in encoded_list
|
||||
]
|
||||
|
||||
def decode(self, tokens: List[int], skip_special_tokens: bool=True) -> str:
|
||||
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
|
||||
return self._tokenizer.decode(tokens, skip_special_tokens=skip_special_tokens)
|
||||
|
||||
|
||||
def __len__(self) -> int:
|
||||
return self._tokenizer.get_vocab_size()
|
||||
|
||||
|
||||
@property
|
||||
def stop_ids(self) -> List[int]:
|
||||
stop_token = self._control_tokens + self._special_tokens
|
||||
stop_ids = [self._tokenizer.token_to_id(token) for token in stop_token]
|
||||
return stop_ids
|
||||
|
||||
|
||||
@property
|
||||
def bos_id(self) -> int:
|
||||
return self._tokenizer.token_to_id("<bos>")
|
||||
|
||||
return self._tokenizer.token_to_id("<|begin▁of▁sentence|>")
|
||||
|
||||
@property
|
||||
def eos_id(self) -> int:
|
||||
return self._tokenizer.token_to_id("<eos>")
|
||||
|
||||
return self._tokenizer.token_to_id("<|end▁of▁sentence|>")
|
||||
|
||||
@property
|
||||
def pad_id(self) -> int:
|
||||
return self._tokenizer.token_to_id("<pad>")
|
||||
return self._tokenizer.token_to_id("<|▁pad▁|>")
|
||||
|
||||
+3
-1
@@ -30,7 +30,9 @@ def error_handler(
|
||||
if reraise:
|
||||
raise
|
||||
return None
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
@@ -58,4 +60,4 @@ def setup_logging(level: Optional[int] = None) -> None:
|
||||
root_logger.addHandler(console_handler)
|
||||
|
||||
logging.getLogger("h5py").setLevel(logging.WARNING)
|
||||
logging.getLogger("torch").setLevel(logging.WARNING)
|
||||
logging.getLogger("torch").setLevel(logging.WARNING)
|
||||
|
||||
Reference in New Issue
Block a user