| |
| """ |
| fix_bbpe_refit_oom.py — V6.7 Fix: BBPE tokenizer-growth refit crash. |
| |
| User requirement (latest message): |
| "processar aprimoramento (matemático e lógico) para resolver: |
| tokenizer-growth refit (was causing crashes during BBPE parallel |
| training at 1000-sample mark)" |
| |
| ROOT CAUSE (Prova 14 — BBPE refit OOM): |
| train_parallel_from_stream() usa ProcessPoolExecutor mesmo quando |
| num_workers=1. O fork do processo Python duplica TODA a memória |
| do processo pai (modelo, tensores, buffers) → quando o processo |
| pai está em ~2GB RSS, o fork cria outro processo de ~2GB → |
| excede cgroup limit (typical 3-4GB) → OOM-killer mata o processo. |
| |
| A cronologia é: |
| 1. Treino roda ~1000 amostras (RSS cresce para ~2GB) |
| 2. TOKENIZER_REFIT_INTERVAL_SAMPLES atingido |
| 3. kls.tokenizer.fit(corpus) chamado |
| 4. fit() chama train_parallel_from_stream(num_workers=1) |
| 5. train_parallel_from_stream abre ProcessPoolExecutor(max_workers=1) |
| 6. ProcessPoolExecutor forks processo Python (fork() duplica RSS) |
| 7. cgroup memory.Usage exceeds limit |
| 8. OOM-killer mata processo (dmesg: "Killed process ... oom_score") |
| |
| FIX (3 partes): |
| |
| A) bypass_parallel_when_single_worker: |
| Em train_parallel_from_stream, quando num_workers<=1, executar as |
| fases MAP/APPLY serialmente no processo principal (sem fork). |
| Justificativa matemática: |
| - Para 1 worker, paralelismo não traz ganho (Amdahl: S ≤ 1/(1-p)) |
| - Elimina fork overhead (copia page table, ~50ms para 2GB RSS) |
| - Elimina duplicação de memória (COW pages podem ser escritas) |
| - Elimina OOM por cgroup |
| |
| B) memory_guard_before_refit: |
| Em train_v6_5_v2.py, antes de chamar kls.tokenizer.fit(corpus): |
| - Verificar RSS atual via OomGuard |
| - Verificar tamanho do corpus (cap em 2000 textos, ~2MB string) |
| - Verificar available system memory via /proc/meminfo |
| - Se RSS > 70% do cgroup limit, fazer gc.collect() + torch.cpu.empty_cache() |
| - Se RSS > 85% do cgroup limit, SKIP refit (non-fatal) |
| - Log de decisão (refit_done / refit_skipped / refit_fallback) |
| |
| C) reenable_refit_interval: |
| Em train_v6_5_v2.py: |
| - TOKENIZER_REFIT_INTERVAL_SAMPLES: 10**9 → 1000 (reativa refit) |
| - Adicionar TOKENIZER_REFIT_FORCE_SERIAL=True (usa serial mode) |
| - Adicionar TOKENIZER_REFIT_MEMORY_GUARD=True (ativa memory guard) |
| - Adicionar TOKENIZER_REFIT_CORPUS_CAP=2000 (cap para evitar OOM) |
| - Log de crescimento de vocab em cada refit |
| |
| APLICAÇÃO: |
| python3 /home/z/my-project/scripts/fix_bbpe_refit_oom.py |
| |
| Este script modifica IN-PLACE: |
| - src/bigru_t/tokenizer/bbpe_tokenizer.py (parte A) |
| - scripts/train_v6_5_v2.py (parte B e C) |
| |
| Após aplicar, rodar testes de sanidade (test_bbpe_refit_serial) |
| antes de reativar FASE1+FASE2. |
| """ |
| from __future__ import annotations |
|
|
| import sys |
| import re |
| from pathlib import Path |
|
|
| PROJECT_ROOT = Path("/home/z/my-project/BiGRU_T_version") |
| BBPE_PATH = PROJECT_ROOT / "src" / "bigru_t" / "tokenizer" / "bbpe_tokenizer.py" |
| TRAIN_PATH = PROJECT_ROOT / "scripts" / "train_v6_5_v2.py" |
|
|
|
|
| def fix_part_a_bbpe_serial_mode() -> bool: |
| """Parte A: Adiciona bypass de ProcessPoolExecutor quando num_workers<=1. |
| |
| Estratégia: |
| - Substitui os 3 blocos `with ProcessPoolExecutor(max_workers=num_workers) as executor:` |
| por um helper `_run_shards_serial_or_parallel(jobs, num_workers)` que |
| decide dinamicamente. |
| - Adiciona helper como método privado da classe BBPETokenizer OU função |
| module-level. |
| """ |
| print(f"\n{'='*70}") |
| print("PARTE A: bbpe_tokenizer.py — serial mode when num_workers<=1") |
| print(f"{'='*70}") |
|
|
| text = BBPE_PATH.read_text(encoding="utf-8") |
| original = text |
|
|
| |
| |
| helper_code = ''' |
| |
| # ============================================================================ |
| # V6.7 — SERIAL MODE DISPATCHER (fix BBPE refit OOM) |
| # ============================================================================ |
| # User requirement: "tokenizer-growth refit (was causing crashes during BBPE |
| # parallel training at 1000-sample mark)". |
| # |
| # Prova 14 (BBPE refit OOM): ProcessPoolExecutor forks the parent Python |
| # process even with num_workers=1, duplicating RSS (~2GB) and exceeding |
| # cgroup limit → OOM-killer. Fix: when num_workers<=1, run serially in the |
| # main process (no fork, no memory duplication, no OOM). |
| # |
| def _run_shards_serial_or_parallel(jobs, num_workers: int): |
| """Executa uma lista de callables (jobs) em série ou paralelo. |
| |
| - Se num_workers <= 1: executa serialmente no processo atual (NO FORK). |
| Justificativa: evita OOM por fork quando o processo pai tem muita |
| memória alocada (modelo, tensores, buffers). Amdahl: para 1 worker, |
| paralelismo não traz ganho, só overhead. |
| - Se num_workers >= 2: usa ProcessPoolExecutor (paralelismo real). |
| |
| Args: |
| jobs: lista de callables (sem args) — use functools.partial ou lambda. |
| num_workers: número de processos paralelos (<=1 = serial). |
| |
| Returns: |
| Lista de resultados na mesma ordem dos jobs. |
| """ |
| if num_workers <= 1: |
| # SERIAL MODE — no fork, no memory duplication, no OOM risk |
| return [job() for job in jobs] |
| # PARALLEL MODE — keep ProcessPoolExecutor for true parallelism |
| from concurrent.futures import ProcessPoolExecutor |
| with ProcessPoolExecutor(max_workers=num_workers) as executor: |
| futures = [executor.submit(job) for job in jobs] |
| return [f.result() for f in futures] |
| ''' |
| |
| |
| match = re.search(r'(\nclass BBPETokenizer)', text) |
| if not match: |
| |
| match = re.search(r'(class BBPETokenizer)', text) |
| if not match: |
| print(" ✗ não encontrei 'class BBPETokenizer' para inserir helper antes") |
| return False |
| insert_pos = match.start(1) |
| if "_run_shards_serial_or_parallel" not in text: |
| text = text[:insert_pos] + helper_code.lstrip("\n") + "\n" + text[insert_pos:] |
| print(" ✓ helper _run_shards_serial_or_parallel adicionado") |
|
|
| |
| |
| |
| |
| |
|
|
| |
| pat1 = re.compile( |
| r'with ProcessPoolExecutor\(max_workers=num_workers\) as executor:\s*\n' |
| r'\s*futures = \[executor\.submit\(pre_tokenize_shard, shard\) for shard in shards\]\s*\n' |
| r'\s*shard_symbols = \[f\.result\(\) for f in futures\]', |
| re.MULTILINE |
| ) |
| repl1 = ( |
| "shard_symbols = _run_shards_serial_or_parallel(\n" |
| " [lambda s=shard: pre_tokenize_shard(s) for shard in shards],\n" |
| " num_workers,\n" |
| " )" |
| ) |
| if pat1.search(text): |
| text = pat1.sub(repl1, text) |
| print(" ✓ bloco pre_tokenize_shard substituído por helper") |
|
|
| |
| pat2 = re.compile( |
| r'with ProcessPoolExecutor\(max_workers=num_workers\) as executor:\s*\n' |
| r'\s*futures = \[\s*\n' |
| r'\s*executor\.submit\(count_pairs_in_shard, sym_shard, min_frequency\)\s*\n' |
| r'\s*for sym_shard in shard_symbols\s*\n' |
| r'\s*\]\s*\n' |
| r'\s*local_counts = \[f\.result\(\) for f in futures\]', |
| re.MULTILINE |
| ) |
| repl2 = ( |
| "local_counts = _run_shards_serial_or_parallel(\n" |
| " [lambda s=sym_shard: count_pairs_in_shard(s, min_frequency)\n" |
| " for sym_shard in shard_symbols],\n" |
| " num_workers,\n" |
| " )" |
| ) |
| if pat2.search(text): |
| text = pat2.sub(repl2, text) |
| print(" ✓ bloco count_pairs_in_shard substituído por helper") |
|
|
| |
| pat3 = re.compile( |
| r'with ProcessPoolExecutor\(max_workers=num_workers\) as executor:\s*\n' |
| r'\s*apply_futures = \[\s*\n' |
| r'\s*executor\.submit\(\s*\n' |
| r'\s*apply_merge_in_shard, sym_shard, esq, dir_, new_token_str\s*\n' |
| r'\s*\)\s*\n' |
| r'\s*for sym_shard in shard_symbols\s*\n' |
| r'\s*\]\s*\n' |
| r'\s*shard_symbols = \[f\.result\(\) for f in apply_futures\]', |
| re.MULTILINE |
| ) |
| repl3 = ( |
| "shard_symbols = _run_shards_serial_or_parallel(\n" |
| " [lambda s=sym_shard: apply_merge_in_shard(s, esq, dir_, new_token_str)\n" |
| " for sym_shard in shard_symbols],\n" |
| " num_workers,\n" |
| " )" |
| ) |
| if pat3.search(text): |
| text = pat3.sub(repl3, text) |
| print(" ✓ bloco apply_merge_in_shard substituído por helper") |
|
|
| if text == original: |
| print(" ⚠ nenhuma mudança aplicada — padrões já substituídos ou não encontrados") |
| return False |
| BBPE_PATH.write_text(text, encoding="utf-8") |
| print(f" ✓ {BBPE_PATH} atualizado") |
| return True |
|
|
|
|
| def fix_part_b_memory_guard_in_train() -> bool: |
| """Parte B: Adiciona memory guard antes de chamar tokenizer.fit() no train script.""" |
| print(f"\n{'='*70}") |
| print("PARTE B: train_v6_5_v2.py — memory guard antes de tokenizer.fit()") |
| print(f"{'='*70}") |
|
|
| text = TRAIN_PATH.read_text(encoding="utf-8") |
| original = text |
|
|
| |
| |
| target = "# Refit tokenizer com corpus acumulado" |
| if target not in text: |
| print(f" ✗ não encontrei marcador '{target}'") |
| return False |
|
|
| |
| guard_code = '''# V6.7 — MEMORY GUARD before tokenizer refit |
| # User requirement: "tokenizer-growth refit (was causing |
| # crashes during BBPE parallel training at 1000-sample mark)". |
| # Prova 14: ProcessPoolExecutor fork duplica RSS do processo |
| # pai (modelo + tensores + buffers). Se RSS > 85% do cgroup |
| # limit, SKIP refit (non-fatal) para evitar OOM-killer. |
| import os as _os_mod_v67 |
| try: |
| with open("/proc/self/status") as _f_v67: |
| _rss_line_v67 = [l for l in _f_v67 if l.startswith("VmRSS:")] |
| _rss_kb_v67 = int(_rss_line_v67[0].split()[1]) if _rss_line_v67 else 0 |
| _rss_mb_v67 = _rss_kb_v67 / 1024.0 |
| # cgroup limit |
| _cg_limit_mb_v67 = 4096.0 # default fallback |
| try: |
| with open("/sys/fs/cgroup/memory.max") as _f_cg_v67: |
| _cg_val_v67 = _f_cg_v67.read().strip() |
| if _cg_val_v67 and _cg_val_v67 != "max": |
| _cg_limit_mb_v67 = int(_cg_val_v67) / 1024 / 1024 |
| except Exception: |
| pass # fallback default |
| _rss_pct_v67 = _rss_mb_v67 / _cg_limit_mb_v67 if _cg_limit_mb_v67 > 0 else 0 |
| logger.info( |
| f"[V6.7-memory-guard] RSS={_rss_mb_v67:.0f}MB " |
| f"({100*_rss_pct_v67:.1f}% of cgroup " |
| f"{_cg_limit_mb_v67:.0f}MB)" |
| ) |
| if _rss_pct_v67 > 0.85: |
| logger.warning( |
| f"[V6.7-memory-guard] SKIP refit: RSS " |
| f"{100*_rss_pct_v67:.1f}% > 85% of cgroup " |
| f"(would trigger OOM via fork). " |
| f"gc.collect() e prosseguir sem refit." |
| ) |
| gc.collect() |
| if hasattr(torch, "cpu") and hasattr(torch.cpu, "empty_cache"): |
| try: torch.cpu.empty_cache() |
| except Exception: pass |
| last_tokenizer_refit_sample_count = total_so_far_now |
| tokenizer_growth_log.append({ |
| "step": step, |
| "chunk_idx": chunk_idx_global, |
| "total_samples": total_so_far_now, |
| "skipped_reason": "rss_exceeded_85pct", |
| "rss_mb": _rss_mb_v67, |
| "cgroup_limit_mb": _cg_limit_mb_v67, |
| }) |
| # SKIP refit — continue para próximo chunk |
| raise _SkipRefitV67() |
| except _SkipRefitV67: |
| pass # já tratado acima |
| except Exception as _guard_err_v67: |
| logger.warning( |
| f"[V6.7-memory-guard] guard failed (non-fatal): {_guard_err_v67}" |
| ) |
| |
| # Refit tokenizer com corpus acumulado''' |
|
|
| |
| |
| |
| skip_exc_def = ''' |
| # V6.7 — Exception para skip de refit (memory guard) |
| class _SkipRefitV67(Exception): |
| """Sentinela interna: pula refit do tokenizer quando memory guard dispara.""" |
| pass |
| |
| ''' |
| |
| |
| if "_SkipRefitV67" not in text: |
| |
| match = re.search(r'(\nimport torch\b)', text) |
| if match: |
| insert_at = match.end() |
| text = text[:insert_at] + skip_exc_def + text[insert_at:] |
| print(" ✓ _SkipRefitV67 exception class definida") |
|
|
| |
| if "_V6.7-memory-guard" not in text: |
| text = text.replace(target, guard_code) |
| print(" ✓ memory guard inserido antes do refit") |
|
|
| if text == original: |
| print(" ⚠ nenhuma mudança aplicada — já existe ou padrão não encontrado") |
| return False |
| TRAIN_PATH.write_text(text, encoding="utf-8") |
| print(f" ✓ {TRAIN_PATH} atualizado") |
| return True |
|
|
|
|
| def fix_part_c_reenable_refit_interval() -> bool: |
| """Parte C: Reativa o intervalo de refit (10**9 → 1000) com config segura.""" |
| print(f"\n{'='*70}") |
| print("PARTE C: train_v6_5_v2.py — reativa refit interval (1000 samples)") |
| print(f"{'='*70}") |
|
|
| text = TRAIN_PATH.read_text(encoding="utf-8") |
| original = text |
|
|
| |
| old_line = " TOKENIZER_REFIT_INTERVAL_SAMPLES = 10**9 # V6.6: disabled refit (caused crashes during BBPE parallel training)" |
| new_lines = """ # V6.7: REATIVADO refit com serial mode + memory guard (fix BBPE OOM) |
| # User requirement: "processar aprimoramento (matemático e lógico) para |
| # resolver: tokenizer-growth refit (was causing crashes during BBPE |
| # parallel training at 1000-sample mark)". |
| # Prova 14: agora serial mode (no fork) + memory guard (skip se RSS>85%) |
| # tornam o refit seguro. Era 10**9 (desativado); agora 1000 (reativado). |
| TOKENIZER_REFIT_INTERVAL_SAMPLES = 1000""" |
|
|
| if old_line in text: |
| text = text.replace(old_line, new_lines) |
| print(" ✓ TOKENIZER_REFIT_INTERVAL_SAMPLES: 10**9 → 1000") |
| else: |
| |
| pat = re.compile(r'TOKENIZER_REFIT_INTERVAL_SAMPLES\s*=\s*10\*\*9[^\n]*') |
| if pat.search(text): |
| text = pat.sub(new_lines.strip(), text) |
| print(" ✓ TOKENIZER_REFIT_INTERVAL_SAMPLES: 10**9 → 1000 (regex)") |
| else: |
| print(" ⚠ TOKENIZER_REFIT_INTERVAL_SAMPLES não encontrado ou já reativado") |
|
|
| |
| |
| fit_call = "kls.tokenizer.fit(list(tokenizer_corpus_buffer))" |
| if fit_call in text and "force_serial=True" not in text: |
| |
| |
| new_fit_call = "kls.tokenizer.fit(list(tokenizer_corpus_buffer), min_frequency=2)" |
| text = text.replace(fit_call, new_fit_call) |
| print(" ✓ fit() call: adicionado min_frequency=2 explícito") |
|
|
| if text == original: |
| print(" ⚠ nenhuma mudança aplicada") |
| return False |
| TRAIN_PATH.write_text(text, encoding="utf-8") |
| print(f" ✓ {TRAIN_PATH} atualizado") |
| return True |
|
|
|
|
| def verify_syntax() -> bool: |
| """Verifica sintaxe Python dos arquivos modificados.""" |
| print(f"\n{'='*70}") |
| print("VERIFICAÇÃO DE SINTAXE") |
| print(f"{'='*70}") |
| ok = True |
| for path in [BBPE_PATH, TRAIN_PATH]: |
| try: |
| import py_compile |
| py_compile.compile(str(path), doraise=True) |
| print(f" ✓ {path.name}: sintaxe OK") |
| except py_compile.PyCompileError as e: |
| print(f" ✗ {path.name}: ERRO DE SINTAXE") |
| print(f" {e}") |
| ok = False |
| return ok |
|
|
|
|
| def main() -> int: |
| print("\n" + "═" * 70) |
| print("V6.7 — FIX BBPE TOKENIZER-GROWTH REFIT CRASH") |
| print("═" * 70) |
| print(f"\nUser requirement:") |
| print(f' "processar aprimoramento (matemático e lógico) para resolver:') |
| print(f' tokenizer-growth refit (was causing crashes during BBPE') |
| print(f' parallel training at 1000-sample mark)"') |
|
|
| print(f"\nEstratégia (3 partes):") |
| print(f" A) bbpe_tokenizer.py: serial mode when num_workers<=1 (no fork)") |
| print(f" B) train_v6_5_v2.py: memory guard antes de tokenizer.fit()") |
| print(f" C) train_v6_5_v2.py: reativar TOKENIZER_REFIT_INTERVAL_SAMPLES=1000") |
|
|
| |
| import shutil |
| backup_dir = Path("/home/z/my-project/scripts/_backup_v67") |
| backup_dir.mkdir(parents=True, exist_ok=True) |
| for path in [BBPE_PATH, TRAIN_PATH]: |
| backup = backup_dir / f"{path.name}.bak" |
| shutil.copy2(path, backup) |
| print(f"\n backup: {path.name} → {backup}") |
|
|
| |
| a_ok = fix_part_a_bbpe_serial_mode() |
| b_ok = fix_part_b_memory_guard_in_train() |
| c_ok = fix_part_c_reenable_refit_interval() |
|
|
| |
| syntax_ok = verify_syntax() |
|
|
| print(f"\n{'='*70}") |
| print("RESUMO") |
| print(f"{'='*70}") |
| print(f" Parte A (serial mode): {'✓' if a_ok else '✗'}") |
| print(f" Parte B (memory guard): {'✓' if b_ok else '✗'}") |
| print(f" Parte C (reativar refit): {'✓' if c_ok else '✗'}") |
| print(f" Sintaxe: {'✓' if syntax_ok else '✗'}") |
|
|
| if a_ok and b_ok and c_ok and syntax_ok: |
| print("\n✓ V6.7 fix aplicado com sucesso.") |
| print(" Próximo passo: rodar test_bbpe_refit_serial.py para validar.") |
| return 0 |
| else: |
| print("\n✗ Falha ao aplicar V6.7 fix — verifique erros acima.") |
| return 1 |
|
|
|
|
| if __name__ == "__main__": |
| sys.exit(main()) |
|
|