style: 使用ruff 工具优化代码风格
This commit is contained in:
+7
-11
@@ -1,12 +1,12 @@
|
||||
from khaosz.data.dataset import (
|
||||
BaseDataset,
|
||||
SEQDataset,
|
||||
DPODataset,
|
||||
SFTDataset,
|
||||
BaseDataset,
|
||||
SEQDataset,
|
||||
DPODataset,
|
||||
SFTDataset,
|
||||
GRPODataset,
|
||||
MultiSegmentFetcher,
|
||||
DatasetLoader,
|
||||
DatasetFactory
|
||||
DatasetFactory,
|
||||
)
|
||||
|
||||
from khaosz.data.tokenizer import BpeTokenizer
|
||||
@@ -15,21 +15,17 @@ from khaosz.data.sampler import ResumableDistributedSampler
|
||||
__all__ = [
|
||||
# Base classes
|
||||
"BaseDataset",
|
||||
|
||||
# Dataset implementations
|
||||
"SEQDataset",
|
||||
"SFTDataset",
|
||||
"DPODataset",
|
||||
"GRPODataset",
|
||||
|
||||
# Fetchers
|
||||
"MultiSegmentFetcher",
|
||||
|
||||
# Factory (DatasetLoader is alias for backward compatibility)
|
||||
"DatasetLoader",
|
||||
"DatasetFactory",
|
||||
|
||||
# Tokenizer and sampler
|
||||
"BpeTokenizer",
|
||||
"ResumableDistributedSampler"
|
||||
]
|
||||
"ResumableDistributedSampler",
|
||||
]
|
||||
|
||||
+99
-70
@@ -12,40 +12,42 @@ from typing import Callable, List, Dict, Literal, Optional, Union
|
||||
|
||||
class BaseSegmentFetcher:
|
||||
"""Fetches data segments across multiple tensor segments.
|
||||
|
||||
|
||||
Maintains cumulative lengths for efficient range queries across
|
||||
multiple discontinuous segments.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(self, segments: List[Tensor]):
|
||||
self.segments = segments
|
||||
self.cum_lengths = []
|
||||
|
||||
|
||||
total = 0
|
||||
for seg in segments:
|
||||
total += torch.numel(seg)
|
||||
self.cum_lengths.append(total)
|
||||
|
||||
|
||||
self.total_length = total
|
||||
|
||||
|
||||
def __len__(self) -> int:
|
||||
return self.total_length
|
||||
|
||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
||||
"""Fetch data in the range [begin_idx, end_idx).
|
||||
|
||||
|
||||
Args:
|
||||
begin_idx: Starting index (inclusive)
|
||||
end_idx: Ending index (exclusive)
|
||||
|
||||
|
||||
Returns:
|
||||
Concatenated tensor of data in the specified range
|
||||
"""
|
||||
if not (0 <= begin_idx < self.total_length and 0 <= end_idx <= self.total_length):
|
||||
if not (
|
||||
0 <= begin_idx < self.total_length and 0 <= end_idx <= self.total_length
|
||||
):
|
||||
raise ValueError("begin_idx or end_idx out of bounds")
|
||||
if begin_idx >= end_idx:
|
||||
return torch.tensor([], dtype=torch.long)
|
||||
|
||||
|
||||
# Find segment boundaries for the range
|
||||
seg_start_idx = bisect.bisect_right(self.cum_lengths, begin_idx)
|
||||
seg_end_idx = bisect.bisect_left(self.cum_lengths, end_idx)
|
||||
@@ -64,43 +66,44 @@ class BaseSegmentFetcher:
|
||||
|
||||
class MultiSegmentFetcher:
|
||||
"""Manages multiple segment fetchers for different data keys.
|
||||
|
||||
|
||||
Each key corresponds to a different type of data (e.g., "sequence", "mask").
|
||||
"""
|
||||
|
||||
|
||||
def __init__(self, muti_segments: Dict):
|
||||
self.muti_keys = list(muti_segments.keys())
|
||||
self.muti_fetchers = {
|
||||
key: BaseSegmentFetcher(segments)
|
||||
for key, segments in muti_segments.items()
|
||||
key: BaseSegmentFetcher(segments) for key, segments in muti_segments.items()
|
||||
}
|
||||
|
||||
|
||||
def __len__(self) -> int:
|
||||
"""Returns the minimum length across all fetchers."""
|
||||
len_list = [len(seg) for seg in self.muti_fetchers.values()]
|
||||
return min(len_list)
|
||||
|
||||
def key_fetch(self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]) -> Dict:
|
||||
|
||||
def key_fetch(
|
||||
self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]
|
||||
) -> Dict:
|
||||
"""Fetch data for specific keys.
|
||||
|
||||
|
||||
Args:
|
||||
begin_idx: Starting index
|
||||
end_idx: Ending index
|
||||
keys: Single key or list of keys to fetch
|
||||
|
||||
|
||||
Returns:
|
||||
Dictionary of tensors if multiple keys, single tensor if one key
|
||||
"""
|
||||
fetch_dict = {}
|
||||
fetch_dict = {}
|
||||
keys = [keys] if isinstance(keys, str) else keys
|
||||
|
||||
|
||||
for key in keys:
|
||||
fetcher = self.muti_fetchers[key]
|
||||
fetch_tensor = fetcher.fetch_data(begin_idx, end_idx)
|
||||
fetch_dict[key] = fetch_tensor
|
||||
|
||||
return fetch_dict if len(keys) > 1 else fetch_dict[keys[0]]
|
||||
|
||||
|
||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Dict:
|
||||
"""Fetch all keys."""
|
||||
return self.key_fetch(begin_idx, end_idx, self.muti_keys)
|
||||
@@ -108,10 +111,10 @@ class MultiSegmentFetcher:
|
||||
|
||||
class BaseDataset(Dataset, ABC):
|
||||
"""Abstract base class for all dataset types.
|
||||
|
||||
|
||||
Implements common functionality for window-based data fetching.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__()
|
||||
self.segments = {}
|
||||
@@ -122,38 +125,38 @@ class BaseDataset(Dataset, ABC):
|
||||
|
||||
def load(self, load_path: str):
|
||||
"""Load dataset from HDF5 file.
|
||||
|
||||
|
||||
Args:
|
||||
load_path: Path to the HDF5 data file
|
||||
"""
|
||||
self.segments = load_h5(load_path)
|
||||
self.fetcher = MultiSegmentFetcher(self.segments)
|
||||
self.total_samples = len(self.fetcher)
|
||||
|
||||
|
||||
def get_index(self, index: int) -> tuple:
|
||||
"""Calculate begin and end indices for a sample.
|
||||
|
||||
|
||||
Args:
|
||||
index: Sample index
|
||||
|
||||
|
||||
Returns:
|
||||
Tuple of (begin_idx, end_idx)
|
||||
"""
|
||||
assert self.total_samples > self.window_size
|
||||
|
||||
|
||||
begin_idx = min(index * self.stride, self.total_samples - 1 - self.window_size)
|
||||
end_idx = min(begin_idx + self.window_size, self.total_samples - 1)
|
||||
|
||||
|
||||
return begin_idx, end_idx
|
||||
|
||||
|
||||
@abstractmethod
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
"""Get a single sample by index.
|
||||
|
||||
|
||||
Must be implemented by subclasses.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def __len__(self) -> int:
|
||||
assert self.total_samples is not None
|
||||
if self.total_samples <= self.window_size:
|
||||
@@ -163,48 +166,50 @@ class BaseDataset(Dataset, ABC):
|
||||
|
||||
class DatasetFactory:
|
||||
"""Factory class for creating dataset instances.
|
||||
|
||||
|
||||
Supports decorator-based registration for extensible dataset types.
|
||||
All default dataset types (seq, sft, dpo, grpo) are registered automatically
|
||||
when their classes are defined with the decorator.
|
||||
|
||||
|
||||
Example usage:
|
||||
@DatasetFactory.register("custom")
|
||||
class CustomDataset(BaseDataset):
|
||||
...
|
||||
|
||||
|
||||
dataset = DatasetFactory.create("custom", window_size, stride)
|
||||
"""
|
||||
|
||||
|
||||
SUPPORTED_TYPES = frozenset({"seq", "sft", "dpo", "grpo"})
|
||||
DATASET_MAP: Dict[str, type] = {}
|
||||
|
||||
|
||||
@classmethod
|
||||
def register(cls, name: str):
|
||||
"""Decorator to register a new dataset class.
|
||||
|
||||
|
||||
Args:
|
||||
name: Registration name for the dataset type
|
||||
|
||||
|
||||
Returns:
|
||||
Decorator function that registers the dataset class
|
||||
"""
|
||||
|
||||
def decorator(dataset_cls: type) -> type:
|
||||
if not issubclass(dataset_cls, BaseDataset):
|
||||
raise TypeError(f"{dataset_cls.__name__} must inherit from BaseDataset")
|
||||
cls.DATASET_MAP[name] = dataset_cls
|
||||
return dataset_cls
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
@classmethod
|
||||
def create(cls, train_type: str, window_size: int, stride: int) -> BaseDataset:
|
||||
"""Create a dataset instance.
|
||||
|
||||
|
||||
Args:
|
||||
train_type: Type of training ("seq", "sft", "dpo", "grpo")
|
||||
window_size: Window size for data sampling
|
||||
stride: Stride between consecutive samples
|
||||
|
||||
|
||||
Returns:
|
||||
Dataset instance
|
||||
"""
|
||||
@@ -213,36 +218,42 @@ class DatasetFactory:
|
||||
f"Unknown dataset type: '{train_type}'. "
|
||||
f"Supported types: {sorted(cls.SUPPORTED_TYPES)}"
|
||||
)
|
||||
|
||||
|
||||
if train_type not in cls.DATASET_MAP:
|
||||
raise NotImplementedError(
|
||||
f"Dataset type '{train_type}' is supported but not yet implemented."
|
||||
)
|
||||
|
||||
|
||||
dataset_cls = cls.DATASET_MAP[train_type]
|
||||
return dataset_cls(window_size, stride)
|
||||
|
||||
|
||||
@classmethod
|
||||
def load(cls, train_type: str, load_path: str, window_size: int, stride: Optional[int] = None) -> BaseDataset:
|
||||
def load(
|
||||
cls,
|
||||
train_type: str,
|
||||
load_path: str,
|
||||
window_size: int,
|
||||
stride: Optional[int] = None,
|
||||
) -> BaseDataset:
|
||||
"""Create and load a dataset in one step.
|
||||
|
||||
|
||||
Args:
|
||||
train_type: Type of training dataset
|
||||
load_path: Path to the data file
|
||||
window_size: Window size for data sampling
|
||||
stride: Stride between consecutive samples (default: same as window_size)
|
||||
|
||||
|
||||
Returns:
|
||||
Loaded dataset instance
|
||||
"""
|
||||
if stride is None:
|
||||
stride = window_size
|
||||
|
||||
|
||||
dataset = cls.create(train_type, window_size, stride)
|
||||
dataset.load(load_path)
|
||||
|
||||
|
||||
return dataset
|
||||
|
||||
|
||||
@classmethod
|
||||
def available_types(cls) -> list:
|
||||
"""Return list of registered dataset type names."""
|
||||
@@ -256,46 +267,50 @@ class DatasetFactory:
|
||||
@DatasetFactory.register("seq")
|
||||
class SEQDataset(BaseDataset):
|
||||
"""Dataset for sequential next-token prediction training."""
|
||||
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
||||
return self.fetcher.key_fetch(begin_idx, end_idx, "sequence")
|
||||
|
||||
|
||||
def __getitem__(self, index):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
|
||||
x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long)
|
||||
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long)
|
||||
|
||||
|
||||
return {"input_ids": x, "target_ids": y}
|
||||
|
||||
|
||||
@DatasetFactory.register("sft")
|
||||
class SFTDataset(BaseDataset):
|
||||
"""Dataset for supervised fine-tuning with loss masking."""
|
||||
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
||||
|
||||
|
||||
def __getitem__(self, index):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
|
||||
x = self._fetch_data(begin_idx, end_idx, "sequence").to(dtype=torch.long)
|
||||
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence").to(dtype=torch.long)
|
||||
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask").to(dtype=torch.bool)
|
||||
|
||||
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence").to(
|
||||
dtype=torch.long
|
||||
)
|
||||
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask").to(
|
||||
dtype=torch.bool
|
||||
)
|
||||
|
||||
return {"input_ids": x, "target_ids": y, "loss_mask": loss_mask}
|
||||
|
||||
|
||||
@DatasetFactory.register("dpo")
|
||||
class DPODataset(BaseDataset):
|
||||
"""Dataset for Direct Preference Optimization training."""
|
||||
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
|
||||
@@ -304,25 +319,34 @@ class DPODataset(BaseDataset):
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
|
||||
chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long)
|
||||
rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
|
||||
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(dtype=torch.bool)
|
||||
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(dtype=torch.bool)
|
||||
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(
|
||||
dtype=torch.bool
|
||||
)
|
||||
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(
|
||||
dtype=torch.bool
|
||||
)
|
||||
|
||||
return {"chosen": chosen, "rejected": rejected, "chosen_mask": chosen_mask, "rejected_mask": rejected_mask}
|
||||
return {
|
||||
"chosen": chosen,
|
||||
"rejected": rejected,
|
||||
"chosen_mask": chosen_mask,
|
||||
"rejected_mask": rejected_mask,
|
||||
}
|
||||
|
||||
|
||||
@DatasetFactory.register("grpo")
|
||||
class GRPODataset(BaseDataset):
|
||||
"""Dataset for Group Relative Policy Optimization training."""
|
||||
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
||||
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
@@ -330,8 +354,13 @@ class GRPODataset(BaseDataset):
|
||||
responses = self._fetch_data(begin_idx, end_idx, "responses")
|
||||
masks = self._fetch_data(begin_idx, end_idx, "masks")
|
||||
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
|
||||
|
||||
return {"prompts": prompts, "responses": responses, "masks": masks, "rewards": rewards}
|
||||
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"responses": responses,
|
||||
"masks": masks,
|
||||
"rewards": rewards,
|
||||
}
|
||||
|
||||
|
||||
# Backward compatibility alias
|
||||
|
||||
+25
-25
@@ -7,45 +7,45 @@ from typing import Optional
|
||||
|
||||
class ResumableDistributedSampler(Sampler[int]):
|
||||
def __init__(
|
||||
self,
|
||||
self,
|
||||
data_source: Dataset,
|
||||
start_epoch: int=0,
|
||||
start_iter: int=0,
|
||||
seed: int=42,
|
||||
drop_last: bool=False,
|
||||
shuffle: bool=True,
|
||||
process_group: Optional[dist.ProcessGroup]=None,
|
||||
start_epoch: int = 0,
|
||||
start_iter: int = 0,
|
||||
seed: int = 42,
|
||||
drop_last: bool = False,
|
||||
shuffle: bool = True,
|
||||
process_group: Optional[dist.ProcessGroup] = None,
|
||||
):
|
||||
self.epoch = start_epoch
|
||||
self.iter = start_iter
|
||||
self.seed = seed
|
||||
self.num_samples = len(data_source)
|
||||
|
||||
|
||||
if process_group is not None:
|
||||
# input process group
|
||||
self.rank = dist.get_rank(process_group)
|
||||
self.num_replicas = dist.get_world_size(process_group)
|
||||
|
||||
|
||||
elif dist.is_available() and dist.is_initialized():
|
||||
# use default process group
|
||||
process_group = dist.group.WORLD
|
||||
self.rank = dist.get_rank()
|
||||
self.num_replicas = dist.get_world_size()
|
||||
|
||||
|
||||
else:
|
||||
# single process
|
||||
self.rank = 0
|
||||
self.num_replicas = 1
|
||||
|
||||
|
||||
self.drop_last = drop_last
|
||||
self.shuffle = shuffle
|
||||
|
||||
offset = 0 if drop_last else self.num_replicas - 1
|
||||
|
||||
offset = 0 if drop_last else self.num_replicas - 1
|
||||
self.num_samples_per_replica = (self.num_samples + offset) // self.num_replicas
|
||||
self.total_size = self.num_samples_per_replica * self.num_replicas
|
||||
|
||||
|
||||
self._indices = None
|
||||
|
||||
|
||||
def _get_indices(self):
|
||||
if self.shuffle:
|
||||
generator = torch.Generator()
|
||||
@@ -53,26 +53,26 @@ class ResumableDistributedSampler(Sampler[int]):
|
||||
indices = torch.randperm(self.num_samples, generator=generator).tolist()
|
||||
else:
|
||||
indices = torch.arange(self.num_samples).tolist()
|
||||
|
||||
|
||||
if not self.drop_last and self.num_samples < self.total_size:
|
||||
padding_size = self.total_size - len(indices)
|
||||
indices += indices[:padding_size]
|
||||
|
||||
local_indices = indices[self.rank:self.total_size:self.num_replicas]
|
||||
|
||||
|
||||
local_indices = indices[self.rank : self.total_size : self.num_replicas]
|
||||
|
||||
self.iter = self.iter % self.num_samples_per_replica
|
||||
self._indices = local_indices[self.iter:]
|
||||
|
||||
self._indices = local_indices[self.iter :]
|
||||
|
||||
def __iter__(self):
|
||||
if self._indices is None:
|
||||
self._get_indices()
|
||||
|
||||
|
||||
for i in self._indices:
|
||||
self.iter += 1
|
||||
yield i
|
||||
|
||||
|
||||
self.epoch += 1
|
||||
self._indices = None
|
||||
|
||||
|
||||
def __len__(self):
|
||||
return self.num_samples_per_replica
|
||||
return self.num_samples_per_replica
|
||||
|
||||
@@ -10,24 +10,26 @@ from torch import Tensor
|
||||
from typing import Any, Dict, List
|
||||
from khaosz.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:
|
||||
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)
|
||||
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:
|
||||
with h5py.File(h5_file, "r") as f:
|
||||
for key in f.keys():
|
||||
grp = f[key]
|
||||
dsets = []
|
||||
@@ -37,7 +39,7 @@ def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
||||
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)
|
||||
@@ -60,7 +62,7 @@ class Checkpoint:
|
||||
self,
|
||||
save_dir: str,
|
||||
) -> None:
|
||||
|
||||
|
||||
save_path = Path(save_dir)
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
@@ -72,7 +74,7 @@ class Checkpoint:
|
||||
}
|
||||
with open(save_path / "meta.json", "w") as f:
|
||||
json.dump(meta, f, indent=2)
|
||||
|
||||
|
||||
st.save_file(self.state_dict, save_path / f"state_dict.safetensors")
|
||||
|
||||
@classmethod
|
||||
@@ -83,7 +85,7 @@ class Checkpoint:
|
||||
|
||||
rank = get_rank()
|
||||
save_path = Path(save_dir)
|
||||
|
||||
|
||||
meta = {}
|
||||
if rank == 0:
|
||||
with open(Path(save_dir) / "meta.json", "r") as f:
|
||||
@@ -100,4 +102,4 @@ class Checkpoint:
|
||||
state_dict=state_dict,
|
||||
epoch=meta["epoch"],
|
||||
iteration=meta["iteration"],
|
||||
)
|
||||
)
|
||||
|
||||
+61
-36
@@ -9,34 +9,46 @@ class BpeTokenizer:
|
||||
def __init__(self, path=None):
|
||||
self._control_tokens = ["<bos>", "<eos>", "<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=18) -> tuple:
|
||||
|
||||
def _prepare_trainer(
|
||||
self,
|
||||
vocab_size: int,
|
||||
min_freq: int,
|
||||
reserved_token_size: int,
|
||||
max_token_length=18,
|
||||
) -> tuple:
|
||||
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,
|
||||
@@ -46,61 +58,74 @@ class BpeTokenizer:
|
||||
initial_alphabet=alphabet,
|
||||
show_progress=True,
|
||||
)
|
||||
|
||||
|
||||
return trainer, detail_vocab_size, reserved_tokens
|
||||
|
||||
def train(self, files, vocab_size, min_freq, reserved_token_size=100):
|
||||
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, vocab_size, min_freq, reserved_token_size=100):
|
||||
|
||||
def train_from_iterator(
|
||||
self, iterator, vocab_size, min_freq, reserved_token_size=100
|
||||
):
|
||||
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)
|
||||
|
||||
|
||||
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:
|
||||
def encode(
|
||||
self,
|
||||
tokens: Union[str, List[str]],
|
||||
out_ids: bool = True,
|
||||
add_special_tokens: bool = False,
|
||||
) -> List:
|
||||
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>")
|
||||
|
||||
|
||||
@property
|
||||
def eos_id(self) -> int:
|
||||
return self._tokenizer.token_to_id("<eos>")
|
||||
|
||||
|
||||
@property
|
||||
def pad_id(self) -> int:
|
||||
return self._tokenizer.token_to_id("<pad>")
|
||||
return self._tokenizer.token_to_id("<pad>")
|
||||
|
||||
Reference in New Issue
Block a user