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