File size: 16,099 Bytes
1f25375
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
#!/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())