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
|
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
@@ -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()
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user