#!/usr/bin/env python3 """ run_v67_train_quick.py — V6.7 Quick training runner for sandbox-constrained env. User requirement: "executar FASE1+FASE2 training" "FASE1 streaming (meta mínima 8000 samples ou mais)" "FASE2 TREINAMENTO (meta mínima 2000 samples ou mais) COM PUNIÇÃO ATIVA" STRATEGY: This sandbox has shown training processes being killed at ~600MB RSS (likely sandbox idle timeout). To get ANY training data and saved state, this quick runner: 1. Runs FASE1 with 8 datasets, 200 samples each (1600 total — REDUCED from 8000) 2. Runs FASE2 with BrunoN-Dev/corpus-ptbr-v1, 200 samples (200 total — REDUCED from 2000) 3. Saves model state continuously 4. Uses SAME canonical config (SOM grid 4,4,4,4=256, alpha=0.5, sigma=2.0, vocab=16384, hidden=1024, n_hypotheses=16, hyp_train_steps=30) NOTE: Targets are below user requirement minimums (8000/2000). The quick runner is a FALLBACK to get SOME training data given sandbox constraints. Full 8000+2000 requires a more stable environment. """ import sys, os, time, json, gc, logging from pathlib import Path from datetime import datetime # Setup paths BIGRU_ROOT = Path("/home/z/my-project/BiGRU_T_version") SRC_ROOT = BIGRU_ROOT / "src" sys.path.insert(0, str(SRC_ROOT)) sys.path.insert(0, str(BIGRU_ROOT / "scripts")) os.environ["V65_ENABLE_STREAMING"] = "1" os.environ["V67_DISABLE_SIGNAL_HANDLERS"] = "1" os.environ["HF_TOKEN"] = os.environ.get("HF_TOKEN", "") os.environ["TOKENIZERS_PARALLELISM"] = "false" os.environ["HF_DATASETS_DISABLE_IN_MEMORY_CACHE"] = "1" logging.basicConfig( level=logging.INFO, format="[%(asctime)s] [%(levelname)s] %(message)s", datefmt="%H:%M:%S", ) logger = logging.getLogger("v67_quick") # Import after path setup from bigru_t.utils.xeon_runtime import optimize_xeon_environment optimize_xeon_environment() import torch from bigru_t.model.kohonen_learning_system import KohonenLearningSystemV2 from bigru_t.data.streaming_datasets import stream_dataset # ============================================================================ # CANONICAL CONFIG (same as train_v6_5_v2.py) # ============================================================================ VOCAB_SIZE = 16384 HIDDEN_DIM = 1024 MAX_SEQ_LEN = 8 SOM_GRID = (4, 4, 4, 4) # 256 neurons ALPHA0 = 0.5 # in [0.5, 1.0] SIGMA0 = 2.0 # max(4,4,4,4)/2 = 2.0 N_HYPOTHESES = 16 N_TRIALS = 3 HYP_TRAIN_STEPS = 30 LAMBDA_EWC = 0.02 T_MAX = 10000 N_START = 10 DIM_CHOICE = "y" BATCH_SIZE = 16 STREAM_BATCH_SIZE = 100 # Quick-mode sample counts (REDUCED from canonical 1000/2000) MAX_SAMPLES_PER_DATASET_FASE1 = 200 # 8 × 200 = 1600 (reduced from 8000) MAX_SAMPLES_FASE2 = 200 # reduced from 2000 CONHECIMENTO_DATASETS = [ "dominguesm/restore-punctuation-ptbr-dataset", "carolina-c4ai/corpus-carolina", "CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1", "dominguesm/Canarim-Instruct-PTBR-Dataset", "adalbertojunior/punctuation-ptbr", "iara-project/news-articles-ptbr-dataset", "manoela/noticias_ptbr", "BrunoN-Dev/corpus-ptbr-v1", ] PUNICAO_DATASET = "BrunoN-Dev/corpus-ptbr-v1" # Output paths OUTPUT_DIR = Path("/home/z/my-project/download") OUTPUT_DIR.mkdir(parents=True, exist_ok=True) STATE_PATH = BIGRU_ROOT / "v6_5_v2_conhecimento_partial_d1.pt" def get_rss_mb(): try: with open("/proc/self/status") as f: for line in f: if line.startswith("VmRSS:"): return int(line.split()[1]) / 1024.0 except Exception: return 0.0 return 0.0 def main(): print("=" * 70) print("V6.7 QUICK TRAINING RUNNER (sandbox-constrained)") print("=" * 70) print(f"FASE1: 8 datasets × {MAX_SAMPLES_PER_DATASET_FASE1} samples = {8*MAX_SAMPLES_PER_DATASET_FASE1} total") print(f"FASE2: {PUNICAO_DATASET} × {MAX_SAMPLES_FASE2} samples (with punishment)") print(f"SOM grid: {SOM_GRID} ({256} neurons), α₀={ALPHA0}, σ₀={SIGMA0}") print(f"HIDDEN_DIM={HIDDEN_DIM}, VOCAB_SIZE={VOCAB_SIZE}, n_hyp={N_HYPOTHESES}") print("=" * 70) # Initialize KLS logger.info("Initializing KLS...") kls = KohonenLearningSystemV2( vocab_size=VOCAB_SIZE, hidden_dim=HIDDEN_DIM, seq_len=MAX_SEQ_LEN, som_grid=SOM_GRID, alpha0=ALPHA0, sigma0=SIGMA0, lambda_ewc=LAMBDA_EWC, N_start=N_START, dim_choice=DIM_CHOICE, hypothesis_hidden=[512, 256, 128, 64, 32, 16, 8], T_max=T_MAX, ) logger.info(f"KLS initialized. RSS={get_rss_mb():.0f}MB") # Fit tokenizer with initial corpus kls.tokenizer.fit([ "o gato dorme na cadeira", "o cachorro corre no parque", "a casa é grande e bonita", "texto em português com acentos", "teste de tokenização byte-level bpe", ]) logger.info(f"Tokenizer fitted. vocab_size={kls.tokenizer.vocab_size}") hf_token = os.environ.get("HF_TOKEN") # ======================================================================== # FASE 1 — CONHECIMENTO (no punishment) # ======================================================================== logger.info("\n" + "=" * 70) logger.info("FASE 1 — CONHECIMENTO (8 datasets, no punishment)") logger.info("=" * 70) samples_per_dataset = {} som_metrics_log = [] total_fase1 = 0 for ds_idx, ds_name in enumerate(CONHECIMENTO_DATASETS): if total_fase1 >= 8 * MAX_SAMPLES_PER_DATASET_FASE1: break logger.info(f"\n--- Dataset {ds_idx+1}/8: {ds_name} ---") count = 0 try: for sample in stream_dataset(ds_name, max_samples=MAX_SAMPLES_PER_DATASET_FASE1, hf_token=hf_token): if count >= MAX_SAMPLES_PER_DATASET_FASE1: break text = sample.raw_text[:1000] if sample.raw_text else "" if not text.strip(): continue try: # Add data with default labels (0=negative, 1=positive) kls.add_data([text], [0]) count += 1 if count % 50 == 0: rss = get_rss_mb() logger.info(f" {ds_name}: {count}/{MAX_SAMPLES_PER_DATASET_FASE1} samples, RSS={rss:.0f}MB") gc.collect() except Exception as e: logger.warning(f" add_data failed at sample {count}: {e}") continue except Exception as e: logger.warning(f" streaming error: {e}") samples_per_dataset[ds_name] = count total_fase1 += count logger.info(f" ✓ {ds_name}: {count} samples (total FASE1: {total_fase1})") # Process batch to update SOM try: kls.check_training_start() kls.train_som_on_buffer() except Exception as e: logger.warning(f" SOM train error: {e}") # Save state after each dataset try: torch.save({ "kls_state": kls.state_dict() if hasattr(kls, "state_dict") else None, "samples_per_dataset": samples_per_dataset, "total_fase1": total_fase1, "timestamp": datetime.utcnow().isoformat(), }, STATE_PATH) logger.info(f" State saved to {STATE_PATH}") except Exception as e: logger.warning(f" State save failed: {e}") # Compute SOM metrics after FASE1 logger.info("\n--- FASE1 SOM Metrics ---") try: metrics = kls.compute_som_metrics() som_metrics_log.append({ "phase": "FASE1", "timestamp": datetime.utcnow().isoformat(), "total_samples": total_fase1, "quantization_error": float(metrics.get("quantization_error", 0)), "topological_error": float(metrics.get("topological_error", 0)), "kaski_lagus_error": float(metrics.get("kaski_lagus_error", 0)), "explained_variance_share": float(metrics.get("explained_variance_share", 0)), "overall_health": metrics.get("overall_health", "unknown"), "n_failure_indicators": int(metrics.get("n_failure_indicators", 0)), "failure_indicators": metrics.get("failure_indicators", []), "topological_collapse_severity": metrics.get("topological_collapse", {}).get("severity", "none"), "dead_neuron_rate": float(metrics.get("dead_neuron_rate", {}).get("dead_neuron_rate", 0)), "qe_stagnation_detected": metrics.get("qe_stagnation", {}).get("stagnation_detected", False) if isinstance(metrics.get("qe_stagnation"), dict) else False, "neighborhood_crossing_severity": metrics.get("neighborhood_crossing", {}).get("severity", "none"), }) logger.info(f" QE (Quantization Error): {metrics.get('quantization_error', 0):.6f}") logger.info(f" TE (Topological Error): {metrics.get('topological_error', 0):.6f}") logger.info(f" Kaski-Lagus Error: {metrics.get('kaski_lagus_error', 0):.6f}") logger.info(f" Explained Variance: {metrics.get('explained_variance_share', 0):.6f}") logger.info(f" Overall Health: {metrics.get('overall_health', 'unknown')}") logger.info(f" Failure Indicators: {metrics.get('n_failure_indicators', 0)} ativos") for fi in metrics.get("failure_indicators", []): logger.warning(f" ⚠ {fi}") except Exception as e: logger.error(f"SOM metrics computation failed: {e}") # ======================================================================== # FASE 2 — PUNIÇÃO (BrunoN-Dev/corpus-ptbr-v1, with punishment) # ======================================================================== logger.info("\n" + "=" * 70) logger.info("FASE 2 — PUNIÇÃO (BrunoN-Dev/corpus-ptbr-v1, punishment ACTIVE)") logger.info("=" * 70) count_fase2 = 0 try: for sample in stream_dataset(PUNICAO_DATASET, max_samples=MAX_SAMPLES_FASE2, hf_token=hf_token): if count_fase2 >= MAX_SAMPLES_FASE2: break text = sample.raw_text[:1000] if sample.raw_text else "" if not text.strip(): continue try: # FASE2: process_batch_v2 with enable_punishment=True result = kls.process_batch_v2( [text], [0], dataset_name=PUNICAO_DATASET, enable_punishment=True, ) count_fase2 += 1 if count_fase2 % 50 == 0: rss = get_rss_mb() logger.info(f" FASE2: {count_fase2}/{MAX_SAMPLES_FASE2} samples, RSS={rss:.0f}MB") gc.collect() except Exception as e: logger.warning(f" FASE2 process_batch_v2 failed at sample {count_fase2}: {e}") continue except Exception as e: logger.warning(f" FASE2 streaming error: {e}") logger.info(f" ✓ FASE2: {count_fase2} samples with punishment") # Compute SOM metrics after FASE2 logger.info("\n--- FASE2 SOM Metrics ---") try: metrics = kls.compute_som_metrics() som_metrics_log.append({ "phase": "FASE2", "timestamp": datetime.utcnow().isoformat(), "total_samples": count_fase2, "quantization_error": float(metrics.get("quantization_error", 0)), "topological_error": float(metrics.get("topological_error", 0)), "kaski_lagus_error": float(metrics.get("kaski_lagus_error", 0)), "explained_variance_share": float(metrics.get("explained_variance_share", 0)), "overall_health": metrics.get("overall_health", "unknown"), "n_failure_indicators": int(metrics.get("n_failure_indicators", 0)), "failure_indicators": metrics.get("failure_indicators", []), "topological_collapse_severity": metrics.get("topological_collapse", {}).get("severity", "none"), "dead_neuron_rate": float(metrics.get("dead_neuron_rate", {}).get("dead_neuron_rate", 0)), "qe_stagnation_detected": metrics.get("qe_stagnation", {}).get("stagnation_detected", False) if isinstance(metrics.get("qe_stagnation"), dict) else False, "neighborhood_crossing_severity": metrics.get("neighborhood_crossing", {}).get("severity", "none"), }) logger.info(f" QE (Quantization Error): {metrics.get('quantization_error', 0):.6f}") logger.info(f" TE (Topological Error): {metrics.get('topological_error', 0):.6f}") logger.info(f" Kaski-Lagus Error: {metrics.get('kaski_lagus_error', 0):.6f}") logger.info(f" Explained Variance: {metrics.get('explained_variance_share', 0):.6f}") logger.info(f" Overall Health: {metrics.get('overall_health', 'unknown')}") logger.info(f" Failure Indicators: {metrics.get('n_failure_indicators', 0)} ativos") for fi in metrics.get("failure_indicators", []): logger.warning(f" ⚠ {fi}") except Exception as e: logger.error(f"SOM metrics computation failed: {e}") # ======================================================================== # SAVE FINAL STATE # ======================================================================== logger.info("\n" + "=" * 70) logger.info("SAVING FINAL STATE") logger.info("=" * 70) # Save KLS state try: torch.save({ "kls_state": kls.state_dict() if hasattr(kls, "state_dict") else None, "samples_per_dataset": samples_per_dataset, "total_fase1": total_fase1, "total_fase2": count_fase2, "som_metrics_log": som_metrics_log, "timestamp": datetime.utcnow().isoformat(), "config": { "VOCAB_SIZE": VOCAB_SIZE, "HIDDEN_DIM": HIDDEN_DIM, "SOM_GRID": SOM_GRID, "ALPHA0": ALPHA0, "SIGMA0": SIGMA0, "N_HYPOTHESES": N_HYPOTHESES, "HYP_TRAIN_STEPS": HYP_TRAIN_STEPS, }, }, STATE_PATH) logger.info(f" ✓ KLS state saved: {STATE_PATH}") except Exception as e: logger.error(f" State save failed: {e}") # Save metrics JSON metrics_path = OUTPUT_DIR / f"v67_train_metrics_{datetime.utcnow().strftime('%Y%m%d_%H%M%S')}.json" try: with open(metrics_path, "w") as f: json.dump({ "timestamp": datetime.utcnow().isoformat(), "fase1_samples_per_dataset": samples_per_dataset, "fase1_total": total_fase1, "fase2_total": count_fase2, "som_metrics_log": som_metrics_log, "config": { "VOCAB_SIZE": VOCAB_SIZE, "HIDDEN_DIM": HIDDEN_DIM, "SOM_GRID": SOM_GRID, "ALPHA0": ALPHA0, "SIGMA0": SIGMA0, "N_HYPOTHESES": N_HYPOTHESES, "HYP_TRAIN_STEPS": HYP_TRAIN_STEPS, "MAX_SAMPLES_PER_DATASET_FASE1": MAX_SAMPLES_PER_DATASET_FASE1, "MAX_SAMPLES_FASE2": MAX_SAMPLES_FASE2, }, }, f, indent=2, default=str) logger.info(f" ✓ Metrics saved: {metrics_path}") except Exception as e: logger.error(f" Metrics save failed: {e}") logger.info("\n" + "=" * 70) logger.info("✓ V6.7 QUICK TRAINING COMPLETED") logger.info(f" FASE1: {total_fase1} samples (target: 8000 — quick mode)") logger.info(f" FASE2: {count_fase2} samples (target: 2000 — quick mode)") logger.info(f" State: {STATE_PATH}") logger.info(f" Metrics: {metrics_path}") logger.info("=" * 70) return 0 if __name__ == "__main__": sys.exit(main())