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"))