cwenzi/neuroflow-cpp
1
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 