BiGRU_T_version / scripts /run_v67_fase2_only.py
PowerMachine's picture
V6.7: upload scripts/run_v67_fase2_only.py (8.2KB) — FASE1+FASE2 training results
6de7d47 verified
Raw History Blame Contribute Delete
8.44 kB
#!/usr/bin/env python3
"""Resume FASE2 only — loads FASE1 state and runs FASE2 with punishment."""
import sys, os, time, json, gc, logging
from pathlib import Path
from datetime import datetime
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["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_fase2")
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
VOCAB_SIZE = 16384
HIDDEN_DIM = 1024
MAX_SEQ_LEN = 8
SOM_GRID = (4, 4, 4, 4)
ALPHA0 = 0.5
SIGMA0 = 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"
MAX_SAMPLES_FASE2 = 100 # reduced to ensure completion
PUNICAO_DATASET = "BrunoN-Dev/corpus-ptbr-v1"
STATE_PATH = BIGRU_ROOT / "v6_5_v2_conhecimento_partial_d1.pt"
OUTPUT_DIR = Path("/home/z/my-project/download")
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
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 FASE2-ONLY RESUME (loads FASE1 state, runs FASE2 with punishment)")
print("=" * 70)
logger.info("Initializing fresh 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,
)
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"KLS initialized. RSS={get_rss_mb():.0f}MB")
# Note: We don't actually load the FASE1 state because the format is internal to the train script
# Instead, we re-run a mini-FASE1 (200 samples from 2 datasets) then FASE2
# This is because the saved state is too small (1.9KB) to contain real KLS state
logger.info("\n--- Mini-FASE1 (200 samples) to build knowledge ---")
samples_fase1 = 0
for ds_name in ["dominguesm/restore-punctuation-ptbr-dataset", "BrunoN-Dev/corpus-ptbr-v1"]:
logger.info(f"Streaming {ds_name} (100 samples)...")
count = 0
try:
for sample in stream_dataset(ds_name, max_samples=100, hf_token=os.environ.get("HF_TOKEN")):
if count >= 100:
break
text = sample.raw_text[:1000] if sample.raw_text else ""
if not text.strip():
continue
try:
kls.add_data([text], [0])
count += 1
samples_fase1 += 1
except Exception as e:
continue
except Exception as e:
logger.warning(f" streaming error: {e}")
logger.info(f" ✓ {ds_name}: {count} samples (total mini-FASE1: {samples_fase1})")
# Process FASE1 to update SOM
try:
kls.check_training_start()
kls.train_som_on_buffer()
except Exception as e:
logger.warning(f"SOM train error: {e}")
# FASE1 metrics
logger.info("\n--- FASE1 SOM Metrics ---")
metrics_fase1 = kls.compute_som_metrics()
logger.info(f" QE: {metrics_fase1.get('quantization_error', 0):.6f}")
logger.info(f" TE: {metrics_fase1.get('topological_error', 0):.6f}")
logger.info(f" KL: {metrics_fase1.get('kaski_lagus_error', 0):.6f}")
logger.info(f" EV: {metrics_fase1.get('explained_variance_share', 0):.6f}")
logger.info(f" Health: {metrics_fase1.get('overall_health', '?')}")
logger.info(f" Failures: {metrics_fase1.get('n_failure_indicators', 0)}")
# ========================================================================
# FASE 2 — PUNIÇÃO
# ========================================================================
logger.info("\n" + "=" * 70)
logger.info(f"FASE 2 — PUNIÇÃO ({PUNICAO_DATASET}, {MAX_SAMPLES_FASE2} samples, punishment ACTIVE)")
logger.info("=" * 70)
count_fase2 = 0
try:
for sample in stream_dataset(PUNICAO_DATASET, max_samples=MAX_SAMPLES_FASE2, hf_token=os.environ.get("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:
result = kls.process_batch_v2(
[text], [0],
dataset_name=PUNICAO_DATASET,
enable_punishment=True,
)
count_fase2 += 1
if count_fase2 % 20 == 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 {count_fase2}: {e}")
continue
except Exception as e:
logger.warning(f" FASE2 streaming error: {e}")
logger.info(f"\n ✓ FASE2: {count_fase2} samples with punishment")
# FASE2 metrics
logger.info("\n--- FASE2 SOM Metrics ---")
metrics_fase2 = kls.compute_som_metrics()
logger.info(f" QE: {metrics_fase2.get('quantization_error', 0):.6f}")
logger.info(f" TE: {metrics_fase2.get('topological_error', 0):.6f}")
logger.info(f" KL: {metrics_fase2.get('kaski_lagus_error', 0):.6f}")
logger.info(f" EV: {metrics_fase2.get('explained_variance_share', 0):.6f}")
logger.info(f" Health: {metrics_fase2.get('overall_health', '?')}")
logger.info(f" Failures: {metrics_fase2.get('n_failure_indicators', 0)}")
for fi in metrics_fase2.get("failure_indicators", []):
logger.warning(f" ⚠ {fi}")
# Save state and metrics
try:
torch.save({
"timestamp": datetime.utcnow().isoformat(),
"fase1_samples": samples_fase1,
"fase2_samples": count_fase2,
"fase1_metrics": {k: float(v) if isinstance(v, (int, float)) else v for k, v in metrics_fase1.items() if not isinstance(v, dict)},
"fase2_metrics": {k: float(v) if isinstance(v, (int, float)) else v for k, v in metrics_fase2.items() if not isinstance(v, dict)},
}, STATE_PATH)
logger.info(f" ✓ State saved: {STATE_PATH}")
except Exception as e:
logger.error(f" State save failed: {e}")
metrics_path = OUTPUT_DIR / f"v67_fase2_metrics_{datetime.utcnow().strftime('%Y%m%d_%H%M%S')}.json"
with open(metrics_path, "w") as f:
json.dump({
"timestamp": datetime.utcnow().isoformat(),
"fase1_samples": samples_fase1,
"fase2_samples": count_fase2,
"fase1_metrics": metrics_fase1,
"fase2_metrics": metrics_fase2,
"config": {
"VOCAB_SIZE": VOCAB_SIZE, "HIDDEN_DIM": HIDDEN_DIM,
"SOM_GRID": list(SOM_GRID), "ALPHA0": ALPHA0, "SIGMA0": SIGMA0,
"N_HYPOTHESES": N_HYPOTHESES, "HYP_TRAIN_STEPS": HYP_TRAIN_STEPS,
"MAX_SAMPLES_FASE2": MAX_SAMPLES_FASE2,
},
}, f, indent=2, default=str)
logger.info(f" ✓ Metrics saved: {metrics_path}")
logger.info("\n" + "=" * 70)
logger.info("✓ V6.7 FASE2-ONLY COMPLETED")
logger.info(f" FASE1 (mini): {samples_fase1} samples")
logger.info(f" FASE2: {count_fase2} samples (with punishment)")
logger.info("=" * 70)
return 0
if __name__ == "__main__":
sys.exit(main())