refactor(dataset): 重构数据集处理逻辑

This commit is contained in:
2025-07-14 08:45:01 +08:00
parent d4df706bc5
commit 7591a231df
4 changed files with 57 additions and 80 deletions
+5 -26
View File
@@ -1,28 +1,7 @@
from datasets import load_dataset from utils import process_dataset
import json
import os
if __name__ == "__main__": if __name__ == "__main__":
dataset_dict = load_dataset("shjwudp/chinese-c4") process_dataset(
train_dataset = dataset_dict["train"] dataset_name="shjwudp/chinese-c4",
output_subdir="chinese-c4"
chunk_size = 1000000 )
total_samples = len(train_dataset)
num_chunks = (total_samples // chunk_size) + 1
script_dir = os.path.dirname(os.path.abspath(__file__))
output_dir = os.path(script_dir, "dataset", "chinese-c4")
os.makedirs(output_dir, exist_ok=True)
for i in range(num_chunks):
start_idx = i * chunk_size
end_idx = min((i + 1) * chunk_size, total_samples)
chunk = train_dataset.select(range(start_idx, end_idx))
output_path = f"{output_dir}/chinese-c4_text_chunk_{i}.jsonl"
with open(output_path, "w", encoding="utf-8") as f:
for example in chunk:
# 每行写入一个 {"text": "xxx"} 对象
json_line = {"text": example["text"]}
f.write(json.dumps(json_line, ensure_ascii=False) + "\n")
print(f"Saved text chunk {i} to {output_path}")
+5 -26
View File
@@ -1,28 +1,7 @@
from datasets import load_dataset from utils import process_dataset
import json
import os
if __name__ == "__main__": if __name__ == "__main__":
dataset_dict = load_dataset("Blaze7451/Wiki-zh-20250601") process_dataset(
train_dataset = dataset_dict["train"] dataset_name="Blaze7451/Wiki-zh-20250601",
chunk_size = 1000000 output_subdir="chinese-wiki"
total_samples = len(train_dataset) )
num_chunks = (total_samples // chunk_size) + 1
script_dir = os.path.dirname(os.path.abspath(__file__))
output_dir = os.path.join(script_dir, "dataset", "chinese-wiki")
os.makedirs(output_dir, exist_ok=True)
for i in range(num_chunks):
start_idx = i * chunk_size
end_idx = min((i + 1) * chunk_size, total_samples)
chunk = train_dataset.select(range(start_idx, end_idx))
output_path = f"{output_dir}/chinese-wiki_text_chunk_{i}.jsonl"
with open(output_path, "w", encoding="utf-8") as f:
for example in chunk:
json_line = {"text": example["text"]}
f.write(json.dumps(json_line, ensure_ascii=False) + "\n")
print(f"Saved text chunk {i} to {output_path}")
+8 -28
View File
@@ -1,6 +1,4 @@
from datasets import load_dataset from utils import process_dataset
import json
import os
import re import re
def comprehensive_normalization(text): def comprehensive_normalization(text):
@@ -14,29 +12,11 @@ def comprehensive_normalization(text):
pattern = re.compile('|'.join(re.escape(k) for k in replacements)) pattern = re.compile('|'.join(re.escape(k) for k in replacements))
return pattern.sub(lambda m: replacements[m.group()], text) return pattern.sub(lambda m: replacements[m.group()], text)
if __name__ == "__main__": if __name__ == "__main__":
dataset_dict = load_dataset("HuggingFaceFW/fineweb","sample-10BT") process_dataset(
train_dataset = dataset_dict["train"] dataset_name="HuggingFaceFW/fineweb",
output_subdir="english-fineweb",
chunk_size = 1000000 dataset_config="sample-10BT",
total_samples = len(train_dataset) normalization_func=comprehensive_normalization
num_chunks = (total_samples // chunk_size) + 1 )
script_dir = os.path.dirname(os.path.abspath(__file__))
output_dir = os.path(script_dir, "dataset", "english-fineweb")
os.makedirs(output_dir, exist_ok=True)
for i in range(num_chunks):
if i == 10:
break
start_idx = i * chunk_size
end_idx = min((i + 1) * chunk_size, total_samples)
chunk = train_dataset.select(range(start_idx, end_idx))
output_path = f"{output_dir}/english-fineweb_text_chunk_{i}.jsonl"
with open(output_path, "w", encoding="utf-8") as f:
for example in chunk:
json_line = {"text": comprehensive_normalization(example["text"])}
f.write(json.dumps(json_line, ensure_ascii=False) + "\n")
print(f"Saved text chunk {i} to {output_path}")
+39
View File
@@ -0,0 +1,39 @@
from datasets import load_dataset
import json
import os
def process_dataset(
dataset_name: str,
output_subdir: str,
dataset_config: str = None,
split_name: str = "train",
chunk_size: int = 1000000,
normalization_func=None
):
dataset_dict = load_dataset(dataset_name, dataset_config) if dataset_config else load_dataset(dataset_name)
train_dataset = dataset_dict[split_name]
total_samples = len(train_dataset)
num_chunks = (total_samples // chunk_size) + 1
script_dir = os.path.dirname(os.path.abspath(__file__))
output_dir = os.path.join(script_dir, "dataset", output_subdir)
os.makedirs(output_dir, exist_ok=True)
for i in range(num_chunks):
start_idx = i * chunk_size
end_idx = min((i + 1) * chunk_size, total_samples)
chunk = train_dataset.select(range(start_idx, end_idx))
output_path = os.path.join(output_dir, f"{output_subdir}_text_chunk_{i}.jsonl")
with open(output_path, "w", encoding="utf-8") as f:
for example in chunk:
text = example["text"]
if normalization_func:
text = normalization_func(text)
json_line = {"text": text}
f.write(json.dumps(json_line, ensure_ascii=False) + "\n")
print(f"Saved text chunk {i} to {output_path}")