refactor: 从data 模块分离tokenizer
This commit is contained in:
+5
-15
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"],
|
||||
)
|
||||
@@ -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
|
||||
)
|
||||
Reference in New Issue
Block a user