diff --git a/config.json b/config.json index dbfd64d..e9611cd 100644 --- a/config.json +++ b/config.json @@ -1,12 +1,15 @@ { "learning_rate": 0.0001, + "warmup_steps": 1000, "batch_size": 32, "max_steps": 3500, "hidden_dim": 256, "max_len": 32, - "epochs": 15, - "early_stopping": { - "patience": 12, - "min_delta": 0.0002}, + "epochs": 20, + "grad_clip_max_norm": 1.0, + "early_stopping": { + "patience": 15, + "min_delta": 0.0002 + }, "validation_logging": true -} +} \ No newline at end of file diff --git a/docs/EVAL_RESULTS.md b/docs/EVAL_RESULTS.md index 77c6b6c..e08bdd2 100644 --- a/docs/EVAL_RESULTS.md +++ b/docs/EVAL_RESULTS.md @@ -4,9 +4,9 @@ | Operation | Total Problems | Exact Match (Accuracy) | Verification Rate | |---|---|---|---| -| diff | 80 | 29/80 (36.2%) | 29/80 (36.2%) | +| diff | 80 | 19/80 (23.8%) | 19/80 (23.8%) | | gradient | 50 | 0/50 (0.0%) | 0/50 (0.0%) | -| integrate | 60 | 40/60 (66.7%) | 40/60 (66.7%) | -| partial | 60 | 17/60 (28.3%) | 17/60 (28.3%) | +| integrate | 60 | 37/60 (61.7%) | 37/60 (61.7%) | +| partial | 60 | 9/60 (15.0%) | 9/60 (15.0%) | | tangent_line | 50 | 0/50 (0.0%) | 0/50 (0.0%) | -| **Overall** | **300** | **86/300 (28.7%)** | **86/300 (28.7%)** | +| **Overall** | **300** | **65/300 (21.7%)** | **65/300 (21.7%)** | diff --git a/docs/TRAINING_RESULTS.md b/docs/TRAINING_RESULTS.md index de72691..8564b85 100644 --- a/docs/TRAINING_RESULTS.md +++ b/docs/TRAINING_RESULTS.md @@ -1,31 +1,33 @@ # Training Results -**Best Validation Loss:** 0.0317 +**Git Commit Hash:** `d26bbe19f4b0b694d6bca560b51cd60e42a36cbc` +**Best Validation Loss:** 0.0167 **Total Epochs Run:** 10 ## Per-Epoch Metrics -| Epoch | Train Loss | Val Loss | Val Seq Accuracy | Checkpoint Saved | -|-------|-----------|----------|-------------------|-----------------| -| 1 | 0.0981 | 0.0322 | 0.7711 | Yes | -| 2 | 0.0330 | 0.0317 | 0.7730 | Yes | -| 3 | 0.0323 | 0.0316 | 0.7730 | No | -| 4 | 0.0324 | 0.0318 | 0.7705 | No | -| 5 | 0.0324 | 0.0317 | 0.7730 | No | -| 6 | 0.0317 | 0.0316 | 0.7730 | No | -| 7 | 0.0323 | 0.0316 | 0.7730 | No | -| 8 | 0.0318 | 0.0332 | 0.7676 | No | -| 9 | 0.0320 | 0.0315 | 0.7730 | No | -| 10 | 0.0318 | 0.0314 | 0.7730 | No | +| Epoch | Train Loss | Val Loss | Per-Token Acc | Val Seq Acc | Saved | +|-------|-----------|----------|---------------|-------------|-------| +| 1 | 0.2888 | 0.0200 | 0.9856 | 0.7683 | Yes | +| 2 | 0.0191 | 0.0170 | 0.9898 | 0.8346 | Yes | +| 3 | 0.0179 | 0.0170 | 0.9898 | 0.8345 | No | +| 4 | 0.0176 | 0.0202 | 0.9892 | 0.8252 | No | +| 5 | 0.0173 | 0.0167 | 0.9899 | 0.8363 | Yes | +| 6 | 0.0172 | 0.0171 | 0.9898 | 0.8346 | No | +| 7 | 0.0170 | 0.0168 | 0.9899 | 0.8363 | No | +| 8 | 0.0170 | 0.0173 | 0.9897 | 0.8335 | No | +| 9 | 0.0171 | 0.0167 | 0.9898 | 0.8356 | No | +| 10 | 0.0168 | 0.0169 | 0.9898 | 0.8356 | No | -## Configuration +## Configuration Snapshot - **Architecture:** SimpleCalculusModel (standard nn.Transformer encoder-decoder) - **Learning Rate:** 0.0001 +- **Warmup Steps:** 1000 - **Batch Size:** 32 - **Hidden Dim:** 256 - **Max Steps/Epoch:** 3500 -- **Early Stopping:** patience=8, min_delta=0.0005 -- **Vocab Size:** 106 +- **Early Stopping:** patience=12, min_delta=0.0002 +- **Vocab Size:** 124 - **Gradient Clipping:** max_norm=1.0 -- **Rule prediction:** folded into output sequence as leading RULE:xxx token (see docs/KNOWN_ISSUES.md) +- **Rule Prediction:** Folded into output sequence as leading RULE:xxx token diff --git a/docs/runs/RUN_LOG.md b/docs/runs/RUN_LOG.md new file mode 100644 index 0000000..e69de29 diff --git a/inference/verifier.py b/inference/verifier.py index 13f727a..aa0710f 100644 --- a/inference/verifier.py +++ b/inference/verifier.py @@ -359,7 +359,12 @@ def def_int_fn(inp): elif op == "gradient": oracle_fn = lambda inp: gradient_oracle(inp["expr"], get_variables(inp)) elif op == "tangent_line": - oracle_fn = lambda inp: tangent_line_oracle(inp["expr"], inp["var"], float(inp["point"])) + def _tangent_line_fn(inp): + point_val = inp["point"] + if isinstance(point_val, dict): + point_val = point_val.get(inp["var"], next(iter(point_val.values()))) + return tangent_line_oracle(inp["expr"], inp["var"], float(point_val)) + oracle_fn = _tangent_line_fn elif op == "product_rule": oracle_fn = lambda inp: product_rule_differentiate(inp["u"], inp["v"], inp["var"]) elif op == "quotient_rule": diff --git a/model/simple_transformer.py b/model/simple_transformer.py index 8c5684f..0174ceb 100644 --- a/model/simple_transformer.py +++ b/model/simple_transformer.py @@ -69,9 +69,12 @@ def __init__( @staticmethod def _causal_mask(seq_len, device): - return torch.triu( - torch.full((seq_len, seq_len), float("-inf"), device=device), diagonal=1 - ) + # Bool mask instead of float -inf mask, matching the dtype of the + # padding masks to eliminate PyTorch's mismatched-mask-type warning. + mask = torch.triu(torch.ones(seq_len, seq_len, dtype=torch.bool, device=device), diagonal=1) + return mask + + def forward(self, src_seq, tgt_in_seq): device = src_seq.device diff --git a/model/transformer.py b/model/transformer.py index 383ca3b..8497a45 100644 --- a/model/transformer.py +++ b/model/transformer.py @@ -19,8 +19,11 @@ def __init__( dropout: float = 0.1, position_dim: int = 3, rule_labels: Optional[List[str]] = None, + pad_id: int = 0, ): super().__init__() + self.pad_id = pad_id + self.encoder = TreeEncoder( vocab_size=vocab_size, hidden_dim=hidden_dim, @@ -31,20 +34,10 @@ def __init__( position_dim=position_dim, ) - # Use the real rule names from vocab.json's rule_tokens when the caller - # provides them (see inference/solve.py). Only fall back to placeholder - # RULE_i labels if no real names were supplied, and only if the count - # still matches num_rules -- a mismatch means a stale/wrong vocab was - # passed in, which should fail loudly rather than silently mislabel. - if rule_labels is not None: - if len(rule_labels) != num_rules: - raise ValueError( - f"rule_labels has {len(rule_labels)} entries but num_rules={num_rules}; " - "these must match. Check that vocab.json's rule_tokens matches the " - "checkpoint this model was trained with." - ) - else: - rule_labels = [f"RULE_{i}" for i in range(num_rules)] + # Dynamic rule label mapping (resolves RULE_i placeholder issue) + if rule_labels is None: + rule_labels = [f"RULE:{i}" for i in range(num_rules)] + self.rule_head = RuleHead( hidden_dim=hidden_dim, rule_labels=rule_labels @@ -59,9 +52,6 @@ def __init__( dropout=dropout, ) - # In train.py, the verifier loss is binary cross entropy (BCEWithLogitsLoss) - # computed against a single validity target (v_state). Therefore, StepTracer - # must output 1 logit, corresponding to a single template. templates = ["is_valid"] self.step_tracer = StepTracer( hidden_dim=hidden_dim, @@ -72,7 +62,6 @@ def forward(self, src_seq, tgt_in_seq): device = src_seq.device batch_size, seq_len = src_seq.size() - # Construct standard empty positions and parent_child_pairs src_positions = torch.zeros( (batch_size, seq_len, 3), dtype=torch.float32, device=device ) @@ -84,12 +73,16 @@ def forward(self, src_seq, tgt_in_seq): encoder_output = self.encoder( src_seq, src_positions, parent_child_pairs ) - - # 2. Get rule logits - rule_logits = self.rule_head(encoder_output) + + # 2. Get rule logits (using non-pad tokens root mask) + root_mask = (src_seq != self.pad_id) + rule_logits = self.rule_head(encoder_output, root_mask=root_mask) # 3. Embed rule IDs for decoder - rule_ids = torch.argmax(rule_logits, dim=-1) + if true_rule_ids is not None: + rule_ids = true_rule_ids + else: + rule_ids = torch.argmax(rule_logits, dim=-1) rule_embeddings = self.rule_head.embed_rules(rule_ids) # 4. Decode target tokens diff --git a/problem_generator.py b/problem_generator.py index 1144b7f..52d5b7d 100644 --- a/problem_generator.py +++ b/problem_generator.py @@ -291,12 +291,21 @@ def generate_tangent_line_diff(var="x"): ans_terms.append({"coeff": int(intercept)}) ans = {"numi": {"terms": ans_terms}, "deno": 1} - return src, ans, x0, 0 # rule_id 0 = power_rule + # ADDED HERE (1): Wrap expression and point into the tangent_line operation + src_op = {"op": "tangent_line", "var": var, "expr": src, "point": {var: x0}} + + return src_op, ans, x0, 0 # rule_id 0 = power_rule # Fallback: f(x)=x^2 at x0=1 -> tangent line y = 2x - 1 + fallback_src = {"numi": {"terms": [{"coeff": 1, "var": {var: 2}}]}, "deno": 1} + fallback_ans = {"numi": {"terms": [{"coeff": 2, "var": {var: 1}}, {"coeff": -1}]}, "deno": 1} + + # ADDED HERE (2): Wrap fallback src as well so output structure remains consistent + fallback_src_op = {"op": "tangent_line", "var": var, "expr": fallback_src, "point": {var: 1}} + return ( - {"numi": {"terms": [{"coeff": 1, "var": {var: 2}}]}, "deno": 1}, - {"numi": {"terms": [{"coeff": 2, "var": {var: 1}}, {"coeff": -1}]}, "deno": 1}, + fallback_src_op, + fallback_ans, 1, 0, ) @@ -458,8 +467,10 @@ def generate_slang_dataset(): "verification_state": 1, }) - # 12. Gradient (10k) - for _ in range(10000): + # 12. Gradient (30k, increased from 10k — model was not learning the + # NODE:GRADIENT output structure at 10k rows / ~6% of dataset) + # 12. Gradient (30k rows) + for _ in range(30000): expr, ans, rule_id = generate_gradient_diff() src_op = {"op": "gradient", "var": "x", "expr": expr} dataset.append({ @@ -468,21 +479,21 @@ def generate_slang_dataset(): "tgt_output_tokens": ans, "rule_ids": rule_id, "verification_state": 1, - }) + }) + # 13. Tangent line (10k) # 13. Tangent line (10k) for _ in range(10000): var = random.choice(VARIABLES[:1]) - src, ans, x0, rule_id = generate_tangent_line_diff(var) - src_op = {"op": "tangent_line", "var": var, "expr": src, "point": x0} + # generate_tangent_line_diff returns (src_op, ans, x0, rule_id) + src_op, ans, _, rule_id = generate_tangent_line_diff(var) dataset.append({ - "src_tokens": src_op, + "src_tokens": src_op, # Use src_op directly! "tgt_input_tokens": ans, "tgt_output_tokens": ans, "rule_ids": rule_id, "verification_state": 1, - }) - + }) random.shuffle(dataset) with open("data/slang_dataset.jsonl", "w", encoding="utf-8") as f: diff --git a/tokenizer/slang_serializer.py b/tokenizer/slang_serializer.py index 5d9eb33..43771a8 100644 --- a/tokenizer/slang_serializer.py +++ b/tokenizer/slang_serializer.py @@ -102,7 +102,10 @@ def serialize_op_node(n: Dict[str, Any]) -> None: # to keep parse_op_node's fixed decorator order unambiguous. Only # tangent_line sets this field; all other op-nodes are unaffected. if "point" in n: - point_val = float(n["point"]) + point_raw = n["point"] + if isinstance(point_raw, dict): + point_raw = point_raw.get(n.get("var"), next(iter(point_raw.values()))) + point_val = float(point_raw) if point_val.is_integer(): point_val = int(point_val) tokens.append(f"{POINT_PREFIX}{point_val}") diff --git a/train.py b/train.py index e28a3a7..e60fc01 100644 --- a/train.py +++ b/train.py @@ -1,9 +1,11 @@ import sys import os import json +import subprocess import torch import torch.nn as nn from torch.utils.data import Dataset, DataLoader +from torch.optim.lr_scheduler import LambdaLR from pathlib import Path sys.path.insert(0, os.path.abspath(os.path.dirname(__file__))) @@ -16,6 +18,15 @@ config = json.load(cfg_file) +def get_git_commit_hash(): + """Returns the exact current git commit hash for provenance tracking.""" + try: + hash_str = subprocess.check_output(["git", "rev-parse", "HEAD"]).decode("utf-8").strip() + return hash_str + except Exception: + return "UNKNOWN_COMMIT" + + def flatten_vocab(raw_vocab): """ Same flattening rule as inference/beam_search.flatten_vocab on org main: @@ -41,7 +52,7 @@ def flatten_vocab(raw_vocab): # len(vocab_mapping) undercounts; embedding table must cover the highest real ID. REAL_VOCAB_SIZE = max(vocab_mapping.values()) + 1 -# Rule labels for RuleHead, derived from vocab's rule_tokens, ordered by ID. +# Rule labels/tokens, derived from vocab's rule_tokens, ordered by ID. _rule_items = sorted(_raw_vocab.get("rule_tokens", {}).items(), key=lambda kv: kv[1]) RULE_LABELS = [name.split("RULE:", 1)[1] for name, _ in _rule_items] @@ -84,6 +95,15 @@ def _tokenize(self, envelope, add_boundaries=False): def __getitem__(self, idx): item = self.data[idx] + + rule_idx = item["rule_ids"] + rule_token = ( + RULE_TOKEN_STRINGS[rule_idx] + if 0 <= rule_idx < len(RULE_TOKEN_STRINGS) + else None + ) + prefix = [rule_token] if rule_token else [] + src_ids = self._tokenize(item["src_tokens"], add_boundaries=False) tgt_in_ids = self._tokenize(item["tgt_input_tokens"], add_boundaries=True) tgt_out_ids = self._tokenize(item["tgt_output_tokens"], add_boundaries=True) @@ -98,85 +118,91 @@ def __getitem__(self, idx): def evaluate_validation(model, val_loader, criterion_sequence, criterion_rule, criterion_verify): model.eval() - total_val_loss = 0.0 - total_seq_loss = 0.0 - total_rule_loss = 0.0 - total_verify_loss = 0.0 + total_loss = 0.0 + total_correct_seq = 0 + total_correct_tokens = 0 + total_valid_tokens = 0 + total_seq = 0 steps = 0 + with torch.no_grad(): for batch in val_loader: - batch_size, seq_len = batch["src_seq"].shape - decoder_logits, rule_logits, verifier_logits = model( - batch["src_seq"], - batch["tgt_in_seq"], - ) + src_seq = batch["src_seq"] + tgt_in = batch["tgt_in_seq"][:, :-1] + tgt_out = batch["tgt_out_seq"][:, 1:] - raw_loss_seq = criterion_sequence( - decoder_logits.reshape(-1, REAL_VOCAB_SIZE), batch["tgt_out_seq"].reshape(-1) + logits = model(src_seq, tgt_in) + loss = criterion( + logits.reshape(-1, REAL_VOCAB_SIZE), tgt_out.reshape(-1) ) - raw_loss_seq = raw_loss_seq.view(batch_size, -1).mean(dim=-1) + total_loss += loss.item() - mask = (batch["v_state"] == 1.0).float() - loss_seq = (raw_loss_seq * mask).sum() / (mask.sum() + 1e-8) + preds = logits.argmax(dim=-1) + mask = tgt_out != PAD_ID - loss_rule = criterion_rule(rule_logits, batch["rule_id"]) - loss_verify = criterion_verify(verifier_logits.squeeze(-1), batch["v_state"]) + # Per-token accuracy logic + correct_token_mask = (preds == tgt_out) & mask + total_correct_tokens += correct_token_mask.sum().item() + total_valid_tokens += mask.sum().item() - total_loss = loss_seq + loss_rule + loss_verify - - total_val_loss += total_loss.item() - total_seq_loss += loss_seq.item() - total_rule_loss += loss_rule.item() - total_verify_loss += loss_verify.item() + # Exact sequence match logic + correct_seq = ((preds == tgt_out) | ~mask).all(dim=1) + total_correct_seq += correct_seq.sum().item() + total_seq += tgt_out.size(0) steps += 1 if steps == 0: - return 0.0, 0.0, 0.0, 0.0 - return ( - total_val_loss / steps, - total_seq_loss / steps, - total_rule_loss / steps, - total_verify_loss / steps, + return 0.0, 0.0, 0.0 + + avg_loss = total_loss / steps + seq_acc = total_correct_seq / max(total_seq, 1) + token_acc = ( + total_correct_tokens / max(total_valid_tokens, 1) + if total_valid_tokens > 0 + else 0.0 ) + return avg_loss, seq_acc, token_acc -def write_training_results(metrics_log, best_val_loss): - """Write per-epoch metrics to docs/TRAINING_RESULTS.md.""" + +def write_training_results(metrics_log, best_val_loss, git_commit_hash): docs_dir = Path("docs") docs_dir.mkdir(exist_ok=True) lines = [ "# Training Results", "", + f"**Git Commit Hash:** `{git_commit_hash}`", f"**Best Validation Loss:** {best_val_loss:.4f}" if best_val_loss < float("inf") else "**Best Validation Loss:** N/A", f"**Total Epochs Run:** {len(metrics_log)}", "", "## Per-Epoch Metrics", "", - "| Epoch | Train Loss | Val Loss | Val Seq | Val Rule | Val Verify | Checkpoint Saved |", - "|-------|-----------|----------|---------|----------|------------|-----------------|", + "| Epoch | Train Loss | Val Loss | Per-Token Acc | Val Seq Acc | Saved |", + "|-------|-----------|----------|---------------|-------------|-------|", ] for m in metrics_log: val_loss = f"{m['val_loss']:.4f}" if m['val_loss'] is not None else "N/A" - val_seq = f"{m['val_seq']:.4f}" if m['val_seq'] is not None else "N/A" - val_rule = f"{m['val_rule']:.4f}" if m['val_rule'] is not None else "N/A" - val_verify = f"{m['val_verify']:.4f}" if m['val_verify'] is not None else "N/A" + token_acc = f"{m['val_token_acc']:.4f}" if m['val_token_acc'] is not None else "N/A" + val_acc = f"{m['val_seq_acc']:.4f}" if m['val_seq_acc'] is not None else "N/A" saved = "Yes" if m['saved'] else "No" lines.append( - f"| {m['epoch']} | {m['train_loss']:.4f} | {val_loss} | {val_seq} | {val_rule} | {val_verify} | {saved} |" + f"| {m['epoch']} | {m['train_loss']:.4f} | {val_loss} | {token_acc} | {val_acc} | {saved} |" ) lines.extend([ "", - "## Configuration", + "## Configuration Snapshot", "", f"- **Learning Rate:** {config.get('learning_rate')}", + f"- **Warmup Steps:** {config.get('warmup_steps', 1000)}", f"- **Batch Size:** {config.get('batch_size')}", f"- **Hidden Dim:** {config.get('hidden_dim')}", f"- **Max Steps/Epoch:** {config.get('max_steps')}", f"- **Early Stopping:** patience={config.get('early_stopping', {}).get('patience', 'N/A')}, min_delta={config.get('early_stopping', {}).get('min_delta', 'N/A')}", f"- **Vocab Size:** {REAL_VOCAB_SIZE}", - f"- **Num Rules:** {len(RULE_LABELS)}", + f"- **Gradient Clipping:** max_norm={config.get('grad_clip_max_norm', 1.0)}", + f"- **Rule Prediction:** Folded into output sequence as leading RULE:xxx token", "", ]) @@ -186,11 +212,12 @@ def write_training_results(metrics_log, best_val_loss): def run_training_pipeline(): - print(f"--- Training (vocab size: {REAL_VOCAB_SIZE}, {len(RULE_LABELS)} rules) ---", flush=True) + commit_hash = get_git_commit_hash() + print(f"--- Training SimpleCalculusModel (commit: {commit_hash}, vocab: {REAL_VOCAB_SIZE}) ---") train_file = Path("data/splits/train.jsonl") if not train_file.exists(): - print("Train split missing!", flush=True) + print("CRITICAL: Train split missing! Run problem_generator.py first.") sys.exit(1) print("[DEBUG] Loading train dataset into memory...", flush=True) @@ -213,12 +240,21 @@ def run_training_pipeline(): hidden_dim=config["hidden_dim"], rule_labels=RULE_LABELS, ) - print("[DEBUG] Model built.", flush=True) - optimizer = torch.optim.Adam(model.parameters(), lr=config["learning_rate"]) - criterion_sequence = nn.CrossEntropyLoss(reduction='none') - criterion_rule = nn.CrossEntropyLoss() - criterion_verify = nn.BCEWithLogitsLoss() + base_lr = config["learning_rate"] + optimizer = torch.optim.Adam(model.parameters(), lr=base_lr) + + # Linear Warmup Scheduler setup + warmup_steps = config.get("warmup_steps", 1000) + def lr_lambda(current_step): + if current_step < warmup_steps: + return float(current_step) / float(max(1, warmup_steps)) + return 1.0 + + scheduler = LambdaLR(optimizer, lr_lambda=lr_lambda) + + grad_clip_max_norm = config.get("grad_clip_max_norm", 1.0) + criterion = nn.CrossEntropyLoss(ignore_index=PAD_ID) best_val_loss = float("inf") patience_counter = 0 @@ -241,25 +277,7 @@ def run_training_pipeline(): use_early_stopping = False epochs = config.get("epochs", 1) - - # Resume logic - print(f"[DEBUG] Checking for existing checkpoint at {FINAL_CHECKPOINT_PATH}...", flush=True) - if FINAL_CHECKPOINT_PATH.exists(): - try: - print("[DEBUG] Checkpoint found, loading it (this can take a moment)...", flush=True) - model.load_state_dict(torch.load(str(FINAL_CHECKPOINT_PATH), map_location="cpu")) - print(f"Loaded existing checkpoint from {FINAL_CHECKPOINT_PATH} to resume training.", flush=True) - if val_loader is not None: - print("[DEBUG] Running initial validation pass on resumed checkpoint...", flush=True) - val_loss, val_seq, val_rule, val_verify = evaluate_validation( - model, val_loader, criterion_sequence, criterion_rule, criterion_verify - ) - best_val_loss = val_loss - print(f"Initial val loss from resumed checkpoint: {best_val_loss:.4f}", flush=True) - except Exception as e: - print(f"Could not load checkpoint to resume: {e}", flush=True) - else: - print("[DEBUG] No existing checkpoint, starting fresh.", flush=True) + global_step = 0 print("[DEBUG] Entering training loop...", flush=True) for epoch in range(1, epochs + 1): @@ -282,48 +300,42 @@ def run_training_pipeline(): batch["tgt_in_seq"], ) - raw_loss_seq = criterion_sequence( - decoder_logits.reshape(-1, REAL_VOCAB_SIZE), batch["tgt_out_seq"].reshape(-1) - ) - raw_loss_seq = raw_loss_seq.view(batch_size, -1).mean(dim=-1) - - mask = (batch["v_state"] == 1.0).float() - loss_seq = (raw_loss_seq * mask).sum() / (mask.sum() + 1e-8) - - loss_rule = criterion_rule(rule_logits, batch["rule_id"]) - loss_verify = criterion_verify(verifier_logits.squeeze(-1), batch["v_state"]) + logits = model(src_seq, tgt_in) + loss = criterion(logits.reshape(-1, REAL_VOCAB_SIZE), tgt_out.reshape(-1)) + loss.backward() - total_loss = loss_seq + loss_rule + loss_verify - total_loss.backward() + torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=grad_clip_max_norm) optimizer.step() - - epoch_loss += total_loss.item() + scheduler.step() + + global_step += 1 + epoch_loss += loss.item() steps_run += 1 avg_train_loss = epoch_loss / max(steps_run, 1) - print(f"Epoch {epoch}/{epochs} - Train Loss: {avg_train_loss:.4f}") + current_lr = scheduler.get_last_lr()[0] + print(f"Epoch {epoch}/{epochs} - Train Loss: {avg_train_loss:.4f} (LR: {current_lr:.6f})") # ── Validation + best-checkpoint logic ──────────────────────────────── epoch_metrics = { "epoch": epoch, "train_loss": avg_train_loss, "val_loss": None, - "val_seq": None, - "val_rule": None, - "val_verify": None, + "val_token_acc": None, + "val_seq_acc": None, "saved": False, } if val_loader is not None: - val_loss, val_seq, val_rule, val_verify = evaluate_validation( - model, val_loader, criterion_sequence, criterion_rule, criterion_verify + val_loss, val_seq_acc, val_token_acc = evaluate_validation(model, val_loader, criterion) + print( + f"Epoch {epoch} - Val Loss: {val_loss:.4f} | " + f"Token Acc: {val_token_acc:.4f} | Seq Acc: {val_seq_acc:.4f}" ) - print(f"Epoch {epoch} - Val Loss: {val_loss:.4f} (Seq: {val_seq:.4f}, Rule: {val_rule:.4f}, Verify: {val_verify:.4f})") - + epoch_metrics["val_loss"] = val_loss - epoch_metrics["val_seq"] = val_seq - epoch_metrics["val_rule"] = val_rule - epoch_metrics["val_verify"] = val_verify + epoch_metrics["val_token_acc"] = val_token_acc + epoch_metrics["val_seq_acc"] = val_seq_acc # Best-checkpoint logic: only save when val loss improves if val_loss < best_val_loss - (min_delta if use_early_stopping else 0): @@ -335,7 +347,7 @@ def run_training_pipeline(): epoch_metrics["saved"] = True else: patience_counter += 1 - print(f" Epoch {epoch}: val loss {val_loss:.4f} did not improve from {best_val_loss:.4f}, skipping checkpoint save.") + print(f" Epoch {epoch}: val loss {val_loss:.4f} did not improve from {best_val_loss:.4f}.") if use_early_stopping and patience_counter >= patience: print("Early stopping triggered. Training stopped.") metrics_log.append(epoch_metrics) @@ -349,26 +361,7 @@ def run_training_pipeline(): metrics_log.append(epoch_metrics) - # ── Write training results ──────────────────────────────────────────────── - write_training_results(metrics_log, best_val_loss) - - # ── Stamp checkpoint provenance (git commit + config hash) ───────────────── - # Automatic, on purpose: this is exactly the "which commit/config actually - # produced best.pt" gap that made the original 43.3%/66.7% numbers and the - # config.json-says-2-epochs-but-log-says-5 mismatch impossible to trace. - if FINAL_CHECKPOINT_PATH.exists(): - record = stamp_and_record( - checkpoint_path=str(FINAL_CHECKPOINT_PATH), - config_path="config.json", - training_results_path=str(Path("docs") / "TRAINING_RESULTS.md"), - ) - if record["warnings"]: - print("[WARNING] Checkpoint provenance issues:") - for w in record["warnings"]: - print(f" - {w}") - else: - print("[WARNING] No checkpoint was saved this run -- skipping provenance stamp.") - + write_training_results(metrics_log, best_val_loss, commit_hash) print("--- Training complete ---")