refactor(dataset): 重构数据处理脚本并支持自定义数据集
This commit is contained in:
@@ -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,
|
||||
)
|
||||
@@ -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,6 +1,5 @@
|
||||
from utils import process_dataset
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
process_dataset(
|
||||
dataset_name="HuggingFaceFW/fineweb",
|
||||
|
||||
+2
-2
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
train_dataset = dataset_dict[split_name]
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user