refactor(dataset): 重构数据处理脚本并支持自定义数据集

This commit is contained in:
2025-07-15 12:18:58 +08:00
parent 1c7e49a631
commit bea4f83842
5 changed files with 25 additions and 15 deletions
+14
View File
@@ -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,
)
-8
View File
@@ -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,
)
-1
View File
@@ -1,6 +1,5 @@
from utils import process_dataset from utils import process_dataset
if __name__ == "__main__": if __name__ == "__main__":
process_dataset( process_dataset(
dataset_name="HuggingFaceFW/fineweb", dataset_name="HuggingFaceFW/fineweb",
+2 -2
View File
@@ -93,8 +93,8 @@ class BpeTokenizer:
else: else:
return [encoding.tokens for encoding in encodings] return [encoding.tokens for encoding in encodings]
def decode(self, tokens: List[int]) -> str: def decode(self, tokens: List[int], skip_special_tokens=True) -> str:
return self._tokenizer.decode(tokens) return self._tokenizer.decode(tokens, skip_special_tokens=skip_special_tokens)
def __len__(self) -> int: def __len__(self) -> int:
return self._tokenizer.get_vocab_size() return self._tokenizer.get_vocab_size()
+9 -4
View File
@@ -61,19 +61,24 @@ def process_dataset(
dataset_name: str, dataset_name: str,
output_subdir: str, output_subdir: str,
dataset_config: str = None, dataset_config: str = None,
max_chunk_size: int = None, max_chunk_num: int = None,
split: str = None,
split_name: str = "train", split_name: str = "train",
column_name: str = "text", column_name: str = "text",
chunk_size: int = 1000000, chunk_size: int = 1000000,
normalization_func=comprehensive_normalization 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] train_dataset = dataset_dict[split_name]
total_samples = len(train_dataset) total_samples = len(train_dataset)
num_chunks = (total_samples // chunk_size) + 1 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__)) script_dir = os.path.dirname(os.path.abspath(__file__))
output_dir = os.path.join(script_dir, "dataset", output_subdir) output_dir = os.path.join(script_dir, "dataset", output_subdir)