From bea4f8384219a2d4b5157e6a4d7771b7526c6c7e Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Tue, 15 Jul 2025 12:18:58 +0800 Subject: [PATCH] =?UTF-8?q?refactor(dataset):=20=E9=87=8D=E6=9E=84?= =?UTF-8?q?=E6=95=B0=E6=8D=AE=E5=A4=84=E7=90=86=E8=84=9A=E6=9C=AC=E5=B9=B6?= =?UTF-8?q?=E6=94=AF=E6=8C=81=E8=87=AA=E5=AE=9A=E4=B9=89=E6=95=B0=E6=8D=AE?= =?UTF-8?q?=E9=9B=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- chinese-cosmopedia.py | 14 ++++++++++++++ chinese-wiki.py | 8 -------- english-fineweb.py | 1 - tokenizer.py | 4 ++-- utils.py | 13 +++++++++---- 5 files changed, 25 insertions(+), 15 deletions(-) create mode 100644 chinese-cosmopedia.py delete mode 100644 chinese-wiki.py diff --git a/chinese-cosmopedia.py b/chinese-cosmopedia.py new file mode 100644 index 0000000..6f55162 --- /dev/null +++ b/chinese-cosmopedia.py @@ -0,0 +1,14 @@ +from utils import process_dataset + +if __name__ == "__main__": + max_chunk_num = 10 + chunk_size = 1000000 + item_size = max_chunk_num * chunk_size + + process_dataset( + dataset_name="chinese-cosmopedia", + output_subdir="chinese-wiki", + split=f"train[:{item_size}]", + max_chunk_num=max_chunk_num, + chunk_size=chunk_size, + ) \ No newline at end of file diff --git a/chinese-wiki.py b/chinese-wiki.py deleted file mode 100644 index 2e08cc1..0000000 --- a/chinese-wiki.py +++ /dev/null @@ -1,8 +0,0 @@ -from utils import process_dataset - -if __name__ == "__main__": - process_dataset( - dataset_name="Blaze7451/Wiki-zh-20250601", - output_subdir="chinese-wiki", - max_chunk_size=5, - ) \ No newline at end of file diff --git a/english-fineweb.py b/english-fineweb.py index 3333b4f..4d53dab 100644 --- a/english-fineweb.py +++ b/english-fineweb.py @@ -1,6 +1,5 @@ from utils import process_dataset - if __name__ == "__main__": process_dataset( dataset_name="HuggingFaceFW/fineweb", diff --git a/tokenizer.py b/tokenizer.py index 689e3fc..ba52a31 100644 --- a/tokenizer.py +++ b/tokenizer.py @@ -93,8 +93,8 @@ class BpeTokenizer: else: return [encoding.tokens for encoding in encodings] - def decode(self, tokens: List[int]) -> str: - return self._tokenizer.decode(tokens) + def decode(self, tokens: List[int], skip_special_tokens=True) -> str: + return self._tokenizer.decode(tokens, skip_special_tokens=skip_special_tokens) def __len__(self) -> int: return self._tokenizer.get_vocab_size() diff --git a/utils.py b/utils.py index 131b110..a2441b4 100644 --- a/utils.py +++ b/utils.py @@ -61,19 +61,24 @@ def process_dataset( dataset_name: str, output_subdir: str, dataset_config: str = None, - max_chunk_size: int = None, + max_chunk_num: int = None, + split: str = None, split_name: str = "train", column_name: str = "text", chunk_size: int = 1000000, normalization_func=comprehensive_normalization ): - dataset_dict = load_dataset(dataset_name, dataset_config) + dataset_dict = load_dataset( + data_dir=dataset_name, + data_files=dataset_config, + split=split + ) + train_dataset = dataset_dict[split_name] - total_samples = len(train_dataset) num_chunks = (total_samples // chunk_size) + 1 - lim_chunks = min(max_chunk_size, num_chunks) if max_chunk_size else num_chunks + lim_chunks = min(max_chunk_num, num_chunks) if max_chunk_num else num_chunks script_dir = os.path.dirname(os.path.abspath(__file__)) output_dir = os.path.join(script_dir, "dataset", output_subdir)