Team Ai
Modelpublic

cwenzi/neuroflow-cpp

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
1likes
prepare_training_data.py395 linesDownload Raw Back to scripts
1#!/usr/bin/env python32"""3NeuroFlow 训练数据预处理工具4 5将 D:\语料\ 中的原始语料预处理为高效的 .tok1 二进制格式。6存放于 WSL 虚拟盘 (速度 ~200 MB/s vs HDD ~125 MB/s)。7 8用法:9  python3 scripts/prepare_training_data.py \10    --corpus D:/语料 \11    --tokenizer configs/tokenizer_128k.json \12    --output /home/user/neuroflow_data \13    --max-seq-len 128 \14    --max-samples 500000015 16输出目录结构:17  output/18    train.tok1       # 训练数据19    stats.json        # 统计信息20"""21 22import argparse23import json24import os25import struct26import sys27import time28from pathlib import Path29 30# ─── BPE Tokenizer ───────────────────────────────────────────31class BPETokenizer:32    """与 C++ BPETokenizer 兼容的 Python 实现"""33 34    def __init__(self, config_path: str):35        with open(config_path, 'r', encoding='utf-8') as f:36            config = json.load(f)37 38        self.vocab = config.get('vocab', {})39        self.id_to_token = {v: k for k, v in self.vocab.items()}40        self.merges = config.get('merges', [])41        self.vocab_size = len(self.vocab)42 43        # Build merge ranks44        self.merge_ranks = {}45        for i, (a, b) in enumerate(self.merges):46            self.merge_ranks[(a, b)] = i47 48        # Special tokens49        self.pad_id = 050        self.unk_id = 151        self.bos_id = 252        self.eos_id = 353 54    def apply_bpe(self, token: str) -> str:55        """Apply BPE merges to a single token using priority queue algorithm"""56        if len(token) <= 1 or not self.merge_ranks:57            return token58 59        symbols = list(token)60        n = len(symbols)61 62        import heapq63 64        # Build linked list65        next_link = list(range(1, n)) + [None]66        prev_link = [None] + list(range(0, n - 1))67 68        pq = []69 70        def push_pair(i):71            j = next_link[i]72            if j is None:73                return74            pair = (symbols[i], symbols[j])75            rank = self.merge_ranks.get(pair)76            if rank is not None:77                heapq.heappush(pq, (rank, i))78 79        for i in range(n):80            push_pair(i)81 82        while pq:83            rank, i = heapq.heappop(pq)84            j = next_link[i]85            if j is None:86                continue87            pair = (symbols[i], symbols[j])88            if self.merge_ranks.get(pair) != rank:89                continue90 91            # Merge92            symbols[i] = symbols[i] + symbols[j]93            k = next_link[j]94            next_link[i] = k95            if k is not None:96                prev_link[k] = i97 98            if prev_link[i] is not None:99                push_pair(prev_link[i])100            push_pair(i)101 102        # Collect result103        result = []104        i = 0105        while i is not None:106            result.append(symbols[i])107            i = next_link[i]108        return ''.join(result)109 110    def encode(self, text: str, max_len: int = 128) -> list[int]:111        """Encode text to token IDs"""112        ids = [self.bos_id]113 114        i = 0115        while i < len(text) and len(ids) < max_len - 1:116            # Handle UTF-8 multi-byte117            char = text[i]118            byte_len = 1119            if ord(char) >= 0x80:120                if ord(char) < 0xE0:121                    byte_len = 2122                elif ord(char) < 0xF0:123                    byte_len = 3124                else:125                    byte_len = 4126 127            byte_seq = text[i:i + byte_len]128            bpe_result = self.apply_bpe(byte_seq)129 130            if bpe_result in self.vocab:131                ids.append(self.vocab[bpe_result])132            else:133                for c in bpe_result:134                    ids.append(self.vocab.get(c, self.unk_id))135            i += byte_len136 137        ids.append(self.eos_id)138        if len(ids) > max_len:139            ids = ids[:max_len - 1] + [self.eos_id]140        return ids141 142    def decode(self, ids: list[int]) -> str:143        """Decode token IDs to text"""144        result = []145        for id in ids:146            if id in (self.pad_id, self.unk_id, self.bos_id, self.eos_id):147                continue148            token = self.id_to_token.get(id, '')149            if token:150                result.append(token)151        return ''.join(result)152 153 154# ─── TOK1 Format Writer ─────────────────────────────────────155# TOK1 Binary Format:156#   Magic:     "TOK1" (4 bytes)157#   Version:   uint16 (2 bytes)158#   VocabSize: uint32 (4 bytes)159#   MaxSeqLen: uint32 (4 bytes)160#   NumSamples:uint32 (4 bytes)161#   For each sample:162#     SeqLen:   uint16 (2 bytes)163#     TokenIDs: uint32[SeqLen] (4 bytes each)164 165 166class TOK1Writer:167    def __init__(self, path: str, vocab_size: int, max_seq_len: int):168        self.f = open(path, 'wb')169        self.f.write(b'TOK1')170        self.f.write(struct.pack('<H', 1))  # version171        self.f.write(struct.pack('<I', vocab_size))172        self.f.write(struct.pack('<I', max_seq_len))173        self.f.write(struct.pack('<I', 0))  # num_samples placeholder174        self.count = 0175        self.max_seq_len = max_seq_len176        self._pos_count = self.f.tell()177 178    def add(self, token_ids: list[int]):179        # Truncate if needed180        if len(token_ids) > self.max_seq_len:181            token_ids = token_ids[:self.max_seq_len]182        self.f.write(struct.pack('<H', len(token_ids)))183        self.f.write(struct.pack(f'<{len(token_ids)}I', *token_ids))184        self.count += 1185 186    def close(self):187        # Update num_samples188        self.f.seek(self._pos_count - 4)  # back to num_samples field189        self.f.write(struct.pack('<I', self.count))190        self.f.close()191        return self.count192 193 194# ─── Corpus Scanner ─────────────────────────────────────────195def scan_corpus(corpus_root: str) -> list[Path]:196    """Scan corpus directory and return list of text files"""197    root = Path(corpus_root)198    if not root.exists():199        print(f"Error: corpus root not found: {corpus_root}")200        sys.exit(1)201 202    supported_exts = {'.txt', '.json', '.jsonl', '.csv', '.tsv', '.md'}203    files = []204    for ext in supported_exts:205        found = list(root.rglob(f'*{ext}'))206        files.extend(found)207        if found:208            print(f"  {ext}: {len(found)} files")209 210    # Sort for reproducibility211    files.sort()212    return files213 214 215def extract_text_from_file(filepath: Path) -> list[str]:216    """Extract text content from various file formats"""217    texts = []218    ext = filepath.suffix.lower()219 220    try:221        if ext == '.txt' or ext == '.md':222            with open(filepath, 'r', encoding='utf-8', errors='ignore') as f:223                content = f.read()224            # Split into paragraphs225            paragraphs = [p.strip() for p in content.split('\n\n') if len(p.strip()) >= 10]226            texts.extend(paragraphs)227 228        elif ext == '.json':229            with open(filepath, 'r', encoding='utf-8', errors='ignore') as f:230                content = f.read()231            # Try JSON array232            try:233                data = json.loads(content)234                if isinstance(data, list):235                    for item in data:236                        if isinstance(item, dict):237                            for key in ('text', 'content', 'title', 'question', 'answer'):238                                if key in item and isinstance(item[key], str):239                                    texts.append(item[key])240                elif isinstance(data, dict):241                    for key in ('text', 'content', 'title', 'question', 'answer'):242                        if key in data and isinstance(data[key], str):243                            texts.append(data[key])244            except json.JSONDecodeError:245                # Fallback: treat as plain text246                if len(content) >= 10:247                    texts.append(content)248 249        elif ext == '.jsonl':250            with open(filepath, 'r', encoding='utf-8', errors='ignore') as f:251                for line in f:252                    line = line.strip()253                    if not line or line.startswith('#'):254                        continue255                    try:256                        item = json.loads(line)257                        for key in ('text', 'content', 'title', 'question', 'answer'):258                            if key in item and isinstance(item[key], str):259                                texts.append(item[key])260                                break261                    except json.JSONDecodeError:262                        pass263 264        elif ext in ('.csv', '.tsv'):265            delim = '\t' if ext == '.tsv' else ','266            with open(filepath, 'r', encoding='utf-8', errors='ignore') as f:267                for line in f:268                    line = line.strip()269                    if not line or line.startswith('#'):270                        continue271                    # Take longest field as text272                    fields = line.split(delim)273                    if fields:274                        longest = max(fields, key=len)275                        if len(longest) >= 10:276                            texts.append(longest)277    except Exception as e:278        pass  # Skip problematic files279 280    return texts281 282 283# ─── Main ────────────────────────────────────────────────────284def main():285    parser = argparse.ArgumentParser(description='NeuroFlow Training Data Preparer')286    parser.add_argument('--corpus', required=True, help='Corpus root directory')287    parser.add_argument('--tokenizer', required=True, help='Tokenizer config path')288    parser.add_argument('--output', required=True, help='Output directory')289    parser.add_argument('--max-seq-len', type=int, default=128, help='Max sequence length')290    parser.add_argument('--max-samples', type=int, default=5000000, help='Max total samples')291    parser.add_argument('--min-text-len', type=int, default=10, help='Minimum text length')292    args = parser.parse_args()293 294    os.makedirs(args.output, exist_ok=True)295 296    # Load tokenizer297    print(f"\n{'='*60}")298    print(f"NeuroFlow 训练数据预处理")299    print(f"{'='*60}")300    print(f"语料根目录: {args.corpus}")301    print(f"输出目录:   {args.output}")302    print(f"最大序列:   {args.max_seq_len}")303    print(f"最大样本:   {args.max_samples:,}")304    print()305 306    print("加载分词器...")307    tok = BPETokenizer(args.tokenizer)308    print(f"  词表大小: {tok.vocab_size}")309    print(f"  Merges:   {len(tok.merges)}")310 311    # Scan corpus312    print("\n扫描语料...")313    files = scan_corpus(args.corpus)314    print(f"总计: {len(files):,} 文件")315 316    # Process files317    print(f"\n开始处理...")318    t0 = time.time()319    total_chars = 0320    total_samples = 0321    skipped = 0322 323    train_path = os.path.join(args.output, 'train.tok1')324    writer = TOK1Writer(train_path, tok.vocab_size, args.max_seq_len)325 326    for i, filepath in enumerate(files):327        texts = extract_text_from_file(filepath)328 329        for text in texts:330            if len(text) < args.min_text_len:331                skipped += 1332                continue333 334            try:335                ids = tok.encode(text, args.max_seq_len)336                if len(ids) >= 4:  # at least bos + 2 tokens + eos337                    writer.add(ids)338                    total_samples += 1339                    total_chars += len(text)340            except Exception:341                skipped += 1342                continue343 344            if total_samples >= args.max_samples:345                break346 347        if total_samples >= args.max_samples:348            break349 350        # Progress report351        if (i + 1) % 1000 == 0:352            elapsed = time.time() - t0353            rate = total_samples / elapsed if elapsed > 0 else 0354            print(f"  [{i+1}/{len(files)}] "355                  f"samples={total_samples:,} "356                  f"chars={total_chars:,} "357                  f"rate={rate:.0f} samples/s "358                  f"elapsed={elapsed:.1f}s")359 360    final_count = writer.close()361    elapsed = time.time() - t0362 363    # Stats364    file_size = os.path.getsize(train_path)365    print(f"\n{'='*60}")366    print(f"处理完成!")367    print(f"{'='*60}")368    print(f"  样本数:    {final_count:,}")369    print(f"  字符数:    {total_chars:,}")370    print(f"  跳过:      {skipped:,}")371    print(f"  文件大小:  {file_size / (1024**3):.2f} GB")372    print(f"  耗时:      {elapsed:.1f}s ({elapsed/60:.1f}min)")373    print(f"  速率:      {final_count/elapsed:.0f} samples/s")374    print(f"  输出:      {train_path}")375 376    # Save stats377    stats = {378        'num_samples': final_count,379        'total_chars': total_chars,380        'skipped': skipped,381        'file_size_bytes': file_size,382        'vocab_size': tok.vocab_size,383        'max_seq_len': args.max_seq_len,384        'corpus_root': args.corpus,385        'elapsed_seconds': elapsed,386    }387    stats_path = os.path.join(args.output, 'stats.json')388    with open(stats_path, 'w') as f:389        json.dump(stats, f, indent=2, ensure_ascii=False)390    print(f"  统计:      {stats_path}")391 392 393if __name__ == '__main__':394    main()395