File size: 4,929 Bytes
3b2d368 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 | # dataset_prep_simple.py
import os
from pathlib import Path
from typing import Optional, Dict
from datasets import load_dataset, concatenate_datasets
# 不再需要 dotenv / huggingface_hub login
# from dotenv import load_dotenv
# from huggingface_hub import login
from lmr.tokenizer import Tokenizer
from .disk_dataset import DiskDataset # 如果是独立脚本改为 from disk_dataset import DiskDataset
def prepare_raw_datasets() -> "datasets.Dataset":
"""
加载并合并 bookcorpus 与 english wikipedia(2022-03-01),
仅保留 'text' 列,返回合并后的单一 Dataset。
"""
bookcorpus = load_dataset("bookcorpus", split="train")
wiki = load_dataset("wikipedia", "20220301.en", split="train")
def keep_only_text(ds, name):
if "text" not in ds.column_names:
raise ValueError(f"Dataset {name} doesn't have a 'text' column: columns={ds.column_names}")
cols_to_remove = [c for c in ds.column_names if c != "text"]
return ds.remove_columns(cols_to_remove) if cols_to_remove else ds
bookcorpus = keep_only_text(bookcorpus, "bookcorpus")
wiki = keep_only_text(wiki, "wikipedia")
raw = concatenate_datasets([bookcorpus, wiki])
return raw
def initialize_dataset(
dataset_config,
dataset_dir: Path,
split_ratios: Optional[Dict[str, float]] = None,
split_token_limits: Optional[Dict[str, Optional[int]]] = None,
):
"""
dataset_config: 提供 dataset_name(属性或 dict key)
dataset_dir: 输出目录(Path 或 字符串)
split_ratios: 可选,比例字典,例如 {'train':0.99,'validation':0.01}
若为 None,将使用默认 {'train':0.99,'validation':0.01}
split_token_limits: 可选,按 split 指定 token_limit(如果不需要可传 None)
"""
# 默认比例:99% train,1% validation
if split_ratios is None:
split_ratios = {"train": 0.99, "validation": 0.01}
# 校验比例和顺序
if not ("train" in split_ratios and "validation" in split_ratios):
raise ValueError("split_ratios must contain 'train' and 'validation' keys.")
total_ratio = sum(split_ratios.values())
if abs(total_ratio - 1.0) > 1e-6:
# 归一化(宽容处理)
split_ratios = {k: v / total_ratio for k, v in split_ratios.items()}
# tokenizer 单例
tokenizer = Tokenizer.get_instance()
# 读取合并数据集
raw_dataset = prepare_raw_datasets()
total = len(raw_dataset)
if total == 0:
raise RuntimeError("Combined raw dataset is empty.")
# 计算切分索引(确保每个 split 至少一个样本,若样本量极小则按可用样本分配)
n_train = int(total * split_ratios["train"])
n_val = total - n_train
# 保证 validation 至少 1(若 total>=2),否则把一个 sample 划给 validation
if total >= 2 and n_val == 0:
n_val = 1
n_train = total - 1
# 如果 total==1,就全部给 train(没有 validation)
indices_train = range(0, n_train)
indices_val = range(n_train, n_train + n_val)
splits = {
"train": raw_dataset.select(list(indices_train)),
}
if n_val > 0:
splits["validation"] = raw_dataset.select(list(indices_val))
# 默认 token limits 全为 None(如果用户传入则使用)
if split_token_limits is None:
split_token_limits = {k: None for k in splits.keys()}
else:
# 只保留我们实际有的 splits,其他忽略
split_token_limits = {k: split_token_limits.get(k) for k in splits.keys()}
# dataset_config 支持对象或 dict
if hasattr(dataset_config, "dataset_name"):
dataset_name = dataset_config.dataset_name
elif isinstance(dataset_config, dict) and "dataset_name" in dataset_config:
dataset_name = dataset_config["dataset_name"]
else:
raise ValueError("dataset_config must provide dataset_name (attribute or dict key).")
# 输出目录
dataset_dir = Path(dataset_dir)
output_dir = dataset_dir / dataset_name
output_dir.mkdir(parents=True, exist_ok=True)
# 生成每个 split 的 bin 文件
for split_name, ds in splits.items():
token_limit = split_token_limits.get(split_name)
out_path = output_dir / f"{split_name}.bin"
metadata_path = output_dir / f"metadata_{split_name}.json"
print(f"[{dataset_name}] Generating split '{split_name}': {len(ds)} samples -> {out_path} (token_limit={token_limit})")
DiskDataset.generate_bin(ds, tokenizer, out_path, token_limit=token_limit, metadata_path=metadata_path)
print(f"Done. Outputs under: {output_dir}")
# ----------------- usage example -----------------
if __name__ == "__main__":
class DummyConfig:
dataset_name = "my_merged_dataset"
# 不传 split_token_limits,使用默认按大小分配
initialize_dataset(DummyConfig(), Path("./data"))
|