| |
| import os |
| from pathlib import Path |
| from typing import Optional, Dict |
|
|
| from datasets import load_dataset, concatenate_datasets |
| |
| |
| |
|
|
| from lmr.tokenizer import Tokenizer |
| 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) |
| """ |
| |
| 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.get_instance() |
|
|
| |
| raw_dataset = prepare_raw_datasets() |
| total = len(raw_dataset) |
| if total == 0: |
| raise RuntimeError("Combined raw dataset is empty.") |
|
|
| |
| n_train = int(total * split_ratios["train"]) |
| n_val = total - n_train |
|
|
| |
| if total >= 2 and n_val == 0: |
| n_val = 1 |
| n_train = total - 1 |
|
|
| |
| 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)) |
|
|
| |
| if split_token_limits is None: |
| split_token_limits = {k: None for k in splits.keys()} |
| else: |
| |
| split_token_limits = {k: split_token_limits.get(k) for k in splits.keys()} |
|
|
| |
| 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) |
|
|
| |
| 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}") |
|
|
| |
| if __name__ == "__main__": |
| class DummyConfig: |
| dataset_name = "my_merged_dataset" |
|
|
| |
| initialize_dataset(DummyConfig(), Path("./data")) |
|
|