refactor: 从data 模块分离tokenizer

This commit is contained in:
2026-04-04 16:12:58 +08:00
parent b531232a9b
commit bd9741dc5f
16 changed files with 108 additions and 92 deletions
+5 -15
View File
@@ -1,29 +1,19 @@
from astrai.data.dataset import (
BaseDataset,
DatasetFactory,
DPODataset,
GRPODataset,
BaseSegmentFetcher,
MultiSegmentFetcher,
SEQDataset,
SFTDataset,
)
from astrai.data.sampler import ResumableDistributedSampler
from astrai.data.tokenizer import BpeTokenizer
__all__ = [
# Base classes
"BaseDataset",
# Dataset implementations
"SEQDataset",
"SFTDataset",
"DPODataset",
"GRPODataset",
# Factory
"DatasetFactory",
# Fetchers
"BaseSegmentFetcher",
"MultiSegmentFetcher",
# Factory (DatasetFactory is alias for backward compatibility)
"DatasetFactory",
"DatasetFactory",
# Tokenizer and sampler
"BpeTokenizer",
# Sampler
"ResumableDistributedSampler",
]
+1 -1
View File
@@ -9,7 +9,7 @@ from torch import Tensor
from torch.utils.data import Dataset
from astrai.factory import BaseFactory
from astrai.data.serialization import load_h5
from astrai.serialization import load_h5
class BaseSegmentFetcher:
-106
View File
@@ -1,106 +0,0 @@
import json
import os
from pathlib import Path
from typing import Any, Dict, List
import h5py
import safetensors.torch as st
import torch
import torch.distributed as dist
from torch import Tensor
from astrai.parallel.setup import get_rank
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
os.makedirs(file_path, exist_ok=True)
full_file_path = os.path.join(file_path, f"{file_name}.h5")
with h5py.File(full_file_path, "w") as f:
for key, tensors in tensor_group.items():
grp = f.create_group(key)
for idx, tensor in enumerate(tensors):
arr = tensor.cpu().numpy()
grp.create_dataset(f"data_{idx}", data=arr)
def load_h5(file_path: str, share_memory=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
class Checkpoint:
def __init__(
self,
state_dict: Dict[str, Any],
epoch: int = 0,
iteration: int = 0,
):
self.state_dict = state_dict
self.epoch = epoch
self.iteration = iteration
def save(
self,
save_dir: str,
) -> None:
save_path = Path(save_dir)
save_path.mkdir(parents=True, exist_ok=True)
rank = get_rank()
if rank == 0:
meta = {
"epoch": self.epoch,
"iteration": self.iteration,
}
with open(save_path / "meta.json", "w") as f:
json.dump(meta, f, indent=2)
st.save_file(self.state_dict, save_path / "state_dict.safetensors")
@classmethod
def load(
cls,
save_dir: str,
) -> "Checkpoint":
rank = get_rank()
save_path = Path(save_dir)
meta = {}
if rank == 0:
with open(Path(save_dir) / "meta.json", "r") as f:
meta = json.load(f)
if dist.is_initialized():
meta_list = [meta]
dist.broadcast_object_list(meta_list, src=0)
meta = meta_list[0]
state_dict = st.load_file(save_path / "state_dict.safetensors")
return cls(
state_dict=state_dict,
epoch=meta["epoch"],
iteration=meta["iteration"],
)
-212
View File
@@ -1,212 +0,0 @@
from abc import ABC, abstractmethod
from typing import List, Union
from tokenizers import Tokenizer, decoders, normalizers, pre_tokenizers, processors
from tokenizers.models import BPE
from tokenizers.trainers import BpeTrainer as BpeTrainerImpl
class BaseTokenizer(ABC):
@abstractmethod
def _init_tokenizer(self):
pass
@abstractmethod
def save(self, path):
pass
@abstractmethod
def load(self, path):
pass
@abstractmethod
def encode(
self,
tokens: Union[str, List[str]],
out_ids: bool = True,
add_special_tokens: bool = False,
) -> List:
pass
@abstractmethod
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
pass
@abstractmethod
def __len__(self) -> int:
pass
@property
@abstractmethod
def stop_ids(self) -> List[int]:
pass
@property
@abstractmethod
def bos_id(self) -> int:
pass
@property
@abstractmethod
def eos_id(self) -> int:
pass
@property
@abstractmethod
def pad_id(self) -> int:
pass
class BaseTrainer(ABC):
def __init__(self, tokenizer: BaseTokenizer):
self.tokenizer = tokenizer
@abstractmethod
def train(self, files, vocab_size, min_freq, **kwargs):
pass
@abstractmethod
def train_from_iterator(self, iterator, vocab_size, min_freq, **kwargs):
pass
class BpeTokenizer(BaseTokenizer):
def __init__(
self,
control_tokens: List[str] = None,
special_tokens: List[str] = None,
path=None,
):
self._control_tokens = control_tokens or [
"<begin▁of▁sentence>",
"<end▁of▁sentence>",
"<|▁pad▁|>",
]
self._special_tokens = special_tokens or [
"<im▁start>",
"<im▁end>",
]
self._tokenizer = None
self._init_tokenizer()
if path is not None:
self.load(path)
def _init_tokenizer(self):
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.decoder = decoders.ByteLevel()
self._tokenizer.post_processor = processors.ByteLevel(trim_offsets=True)
def save(self, path):
self._tokenizer.save(path)
def load(self, path):
self._tokenizer = Tokenizer.from_file(path)
def encode(
self,
tokens: Union[str, List[str]],
out_ids: bool = True,
add_special_tokens: bool = False,
) -> List:
if isinstance(tokens, str):
encoded = self._tokenizer.encode(
tokens, add_special_tokens=add_special_tokens
)
return encoded.ids if out_ids else encoded.tokens
else:
encoded_list = 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:
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
return [self._tokenizer.token_to_id(tok) for tok in stop_token]
@property
def bos_id(self) -> int:
return self._tokenizer.token_to_id(self._control_tokens[0])
@property
def eos_id(self) -> int:
return self._tokenizer.token_to_id(self._control_tokens[1])
@property
def pad_id(self) -> int:
return self._tokenizer.token_to_id(self._control_tokens[2])
class BpeTrainer(BaseTrainer):
def __init__(self, tokenizer: BaseTokenizer):
super().__init__(tokenizer)
def _prepare_trainer(
self,
vocab_size: int,
min_freq: int,
reserved_token_size: int,
max_token_length=18,
):
assert reserved_token_size > len(self.tokenizer._special_tokens)
reserved_tokens = [
f"<|reserve{i:02d}|>"
for i in range(reserved_token_size - len(self.tokenizer._special_tokens))
]
detail_vocab_size = vocab_size - (
len(reserved_tokens) + len(self.tokenizer._special_tokens)
)
alphabet = pre_tokenizers.ByteLevel.alphabet()
min_size = len(alphabet) + len(self.tokenizer._control_tokens)
assert detail_vocab_size > min_size
trainer = BpeTrainerImpl(
vocab_size=detail_vocab_size,
min_frequency=min_freq,
limit_alphabet=detail_vocab_size // 6,
max_token_length=max_token_length,
special_tokens=self.tokenizer._control_tokens,
initial_alphabet=alphabet,
show_progress=True,
)
return trainer, reserved_tokens
def train(self, files, vocab_size, min_freq, reserved_token_size=100, **kwargs):
trainer, reserved_tokens = self._prepare_trainer(
vocab_size, min_freq, reserved_token_size, **kwargs
)
self.tokenizer._tokenizer.train(files=files, trainer=trainer)
self.tokenizer._tokenizer.add_special_tokens(
self.tokenizer._special_tokens + reserved_tokens
)
def train_from_iterator(
self, iterator, vocab_size, min_freq, reserved_token_size=100, **kwargs
):
trainer, reserved_tokens = self._prepare_trainer(
vocab_size, min_freq, reserved_token_size, **kwargs
)
self.tokenizer._tokenizer.train_from_iterator(
iterator=iterator, trainer=trainer
)
self.tokenizer._tokenizer.add_special_tokens(
self.tokenizer._special_tokens + reserved_tokens
)