""" Test-Time Training (TTT) Engine for ARC-AGI Implements per-task LoRA fine-tuning with D8 augmentations. Based on: - Akyürek et al. (2411.07279): TTT design + augmented inference - Franzen et al. (2505.07859): DFS + Product-of-Experts scoring """ import copy import json import time import random from typing import List, Dict, Tuple, Optional, Any from collections import Counter, defaultdict import numpy as np from arc_data import ( grids_equal, grid_to_string, string_to_grid, task_to_prompt, create_ttt_dataset, create_leave_one_out_tasks, augment_task, reverse_d8, reverse_color_permutation, D8_TRANSFORMS, D8_INVERSES, get_d8_transform, create_color_permutation, apply_color_permutation, grid_to_numpy, numpy_to_grid ) # ============================================================ # Grid tokenization (Franzen-style: minimal vocab) # ============================================================ ARC_VOCAB = { "0": 0, "1": 1, "2": 2, "3": 3, "4": 4, "5": 5, "6": 6, "7": 7, "8": 8, "9": 9, "\n": 10, " ": 11, # Row separator and cell separator "|": 12, # Input/output separator "": 13, "": 14, "": 15, "": 16, "": 17, "": 18, "": 19, } def tokenize_grid(grid: List[List[int]]) -> str: """Tokenize a grid using minimal representation.""" return "\n".join(" ".join(str(c) for c in row) for row in grid) def format_task_for_model(task: Dict, include_test_output: bool = False) -> str: """ Format an ARC task for LLM input using structured markers. """ parts = [] for i, pair in enumerate(task["train"]): parts.append(f"") parts.append(f"I:\n{tokenize_grid(pair['input'])}") parts.append(f"O:\n{tokenize_grid(pair['output'])}") parts.append(f"") parts.append(f"") parts.append(f"I:\n{tokenize_grid(task['test'][0]['input'])}") parts.append(f"O:") if include_test_output and "output" in task["test"][0] and task["test"][0]["output"]: parts.append(f"\n{tokenize_grid(task['test'][0]['output'])}") return "\n".join(parts) def parse_grid_from_output(text: str) -> Optional[List[List[int]]]: """Parse a grid from model output text.""" text = text.strip() # Remove any markers for marker in ["", "", ""]: text = text.replace(marker, "") text = text.strip() if not text: return None try: lines = text.split("\n") grid = [] for line in lines: line = line.strip() if not line: continue cells = line.split() row = [int(c) for c in cells] if all(0 <= c <= 9 for c in row): grid.append(row) if len(grid) > 0 and all(len(row) == len(grid[0]) for row in grid): return grid except (ValueError, IndexError): pass return None # ============================================================ # TTT Dataset Construction # ============================================================ def build_ttt_training_data(task: Dict, max_examples: int = 250) -> List[Tuple[str, str]]: """ Build TTT training data from leave-one-out tasks with augmentations. Returns list of (input_text, target_text) pairs. """ loo_tasks = create_leave_one_out_tasks(task) data = [] for loo_task in loo_tasks: for t_name, _ in D8_TRANSFORMS: # D8 transform only aug_task = augment_task(loo_task, transform_name=t_name) input_text = format_task_for_model(aug_task, include_test_output=False) target_text = tokenize_grid(aug_task["test"][0]["output"]) data.append((input_text, target_text)) if len(data) >= max_examples: break # D8 + color permutation color_perm = create_color_permutation(seed=hash(t_name) % 10000) aug_task_c = augment_task(loo_task, transform_name=t_name, color_perm=color_perm) input_text_c = format_task_for_model(aug_task_c, include_test_output=False) target_text_c = tokenize_grid(aug_task_c["test"][0]["output"]) data.append((input_text_c, target_text_c)) if len(data) >= max_examples: break # D8 + permuted examples aug_task_p = augment_task(loo_task, transform_name=t_name, permute_examples=True) input_text_p = format_task_for_model(aug_task_p, include_test_output=False) target_text_p = tokenize_grid(aug_task_p["test"][0]["output"]) data.append((input_text_p, target_text_p)) if len(data) >= max_examples: break if len(data) >= max_examples: break random.shuffle(data) return data[:max_examples] # ============================================================ # Augmented Inference with Hierarchical Voting # ============================================================ def augmented_inference(model, tokenizer, task: Dict, n_augmentations: int = 16, use_color_perms: bool = True) -> List[Dict]: """ Generate predictions using augmented inference. Returns list of {transform_name, color_perm, prediction, raw_prediction} """ import torch candidates = [] for t_name, _ in D8_TRANSFORMS: aug_task = augment_task(task, transform_name=t_name) prompt = format_task_for_model(aug_task, include_test_output=False) # Generate prediction inputs = tokenizer(prompt + "\n", return_tensors="pt", truncation=True, max_length=4096) inputs = {k: v.to(model.device) for k, v in inputs.items()} with torch.no_grad(): outputs = model.generate( **inputs, max_new_tokens=900, temperature=0.1, do_sample=False, pad_token_id=tokenizer.eos_token_id, ) raw_text = tokenizer.decode(outputs[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True) raw_pred = parse_grid_from_output(raw_text) if raw_pred is not None: # Reverse the D8 transform to get prediction in original space pred = [] for row in raw_pred: pred.append(row) pred = reverse_d8(pred, t_name) candidates.append({ "transform_name": t_name, "color_perm": None, "prediction": pred, "raw_prediction": raw_pred, }) # Also try with a color permutation if use_color_perms and len(candidates) < n_augmentations: color_perm = create_color_permutation(seed=hash(t_name + "color") % 10000) aug_task_c = augment_task(task, transform_name=t_name, color_perm=color_perm) prompt_c = format_task_for_model(aug_task_c, include_test_output=False) inputs_c = tokenizer(prompt_c + "\n", return_tensors="pt", truncation=True, max_length=4096) inputs_c = {k: v.to(model.device) for k, v in inputs_c.items()} with torch.no_grad(): outputs_c = model.generate( **inputs_c, max_new_tokens=900, temperature=0.1, do_sample=False, pad_token_id=tokenizer.eos_token_id, ) raw_text_c = tokenizer.decode(outputs_c[0][inputs_c["input_ids"].shape[1]:], skip_special_tokens=True) raw_pred_c = parse_grid_from_output(raw_text_c) if raw_pred_c is not None: # Reverse color permutation then D8 pred_c = reverse_color_permutation(raw_pred_c, color_perm) pred_c = reverse_d8(pred_c, t_name) candidates.append({ "transform_name": t_name + "_color", "color_perm": color_perm, "prediction": pred_c, "raw_prediction": raw_pred_c, }) return candidates def hierarchical_vote(candidates: List[Dict], top_k: int = 2) -> List[List[List[int]]]: """ Hierarchical voting from Akyürek et al.: 1. Intra-transformation voting: group by transform, pick top-3 per group 2. Global voting: pick top-2 across all groups """ if not candidates: return [] # Helper to make grid hashable def grid_key(g): return tuple(tuple(r) for r in g) # Stage 1: Intra-transformation voting transform_groups = defaultdict(list) for cand in candidates: base_transform = cand["transform_name"].replace("_color", "") transform_groups[base_transform].append(cand["prediction"]) stage1_candidates = [] for transform_name, preds in transform_groups.items(): # Count frequencies freq = Counter(grid_key(p) for p in preds) # Top-3 per transform group for key, count in freq.most_common(3): stage1_candidates.append({ "prediction": [list(r) for r in key], "intra_count": count, "transform": transform_name, }) # Stage 2: Global voting global_freq = Counter() global_grids = {} for cand in stage1_candidates: key = grid_key(cand["prediction"]) global_freq[key] += cand["intra_count"] global_grids[key] = cand["prediction"] # Select top-k results = [] for key, count in global_freq.most_common(top_k): results.append(global_grids[key]) return results # ============================================================ # Product-of-Experts Scoring (Franzen et al.) # ============================================================ def product_of_experts_score(model, tokenizer, task: Dict, candidates: List[List[List[int]]], n_augmentations: int = 8) -> List[Tuple[List[List[int]], float]]: """ Score candidate solutions using Product of Experts. For each candidate, compute likelihood under multiple augmented views. Final score = product of likelihoods. """ import torch scored = [] for candidate in candidates: log_score = 0.0 n_valid = 0 for t_name, _ in D8_TRANSFORMS[:n_augmentations]: # Create augmented task with this candidate as output aug_task = copy.deepcopy(task) aug_task["test"][0]["output"] = candidate aug_task = augment_task(aug_task, transform_name=t_name) # Get augmented candidate aug_candidate = aug_task["test"][0]["output"] # Format as full sequence (input + output) prompt = format_task_for_model(aug_task, include_test_output=False) + "\n" target = tokenize_grid(aug_candidate) full_text = prompt + target # Compute log-likelihood of target tokens inputs = tokenizer(full_text, return_tensors="pt", truncation=True, max_length=4096) inputs = {k: v.to(model.device) for k, v in inputs.items()} prompt_len = len(tokenizer(prompt, truncation=True, max_length=4096)["input_ids"]) with torch.no_grad(): outputs = model(**inputs) logits = outputs.logits # Log-probability of target tokens target_logits = logits[0, prompt_len-1:-1, :] target_ids = inputs["input_ids"][0, prompt_len:] if len(target_ids) > 0: log_probs = torch.nn.functional.log_softmax(target_logits, dim=-1) token_log_probs = log_probs.gather(1, target_ids.unsqueeze(1)).squeeze(1) avg_log_prob = token_log_probs.mean().item() log_score += avg_log_prob n_valid += 1 if n_valid > 0: scored.append((candidate, log_score / n_valid)) # Geometric mean # Sort by score (higher is better) scored.sort(key=lambda x: -x[1]) return scored # ============================================================ # Full TTT Pipeline # ============================================================ class TTTEngine: """ Test-Time Training engine for ARC tasks. Per-task LoRA adaptation + augmented inference + voting. """ def __init__(self, model=None, tokenizer=None, lora_rank: int = 32, ttt_steps: int = 64, ttt_lr: float = 2e-4, max_ttt_examples: int = 250): self.model = model self.tokenizer = tokenizer self.lora_rank = lora_rank self.ttt_steps = ttt_steps self.ttt_lr = ttt_lr self.max_ttt_examples = max_ttt_examples def apply_ttt(self, task: Dict) -> Any: """ Apply test-time training for a specific task. Returns the adapted model (with task-specific LoRA). """ import torch from peft import LoraConfig, get_peft_model, TaskType # Build TTT dataset ttt_data = build_ttt_training_data(task, max_examples=self.max_ttt_examples) if not ttt_data: return self.model # Apply LoRA lora_config = LoraConfig( r=self.lora_rank, lora_alpha=self.lora_rank * 2, target_modules=["q_proj", "v_proj"], lora_dropout=0.05, bias="none", task_type=TaskType.CAUSAL_LM, ) adapted_model = get_peft_model(self.model, lora_config) adapted_model.train() # Optimizer optimizer = torch.optim.AdamW( adapted_model.parameters(), lr=self.ttt_lr, weight_decay=0.01 ) # Training loop for step in range(min(self.ttt_steps, len(ttt_data))): input_text, target_text = ttt_data[step % len(ttt_data)] full_text = input_text + "\n" + target_text inputs = self.tokenizer( full_text, return_tensors="pt", truncation=True, max_length=4096, padding=False, ) inputs = {k: v.to(adapted_model.device) for k, v in inputs.items()} # Compute loss only on target tokens prompt_len = len(self.tokenizer(input_text + "\n", truncation=True, max_length=4096)["input_ids"]) labels = inputs["input_ids"].clone() labels[0, :prompt_len] = -100 # Mask input tokens outputs = adapted_model(**inputs, labels=labels) loss = outputs.loss loss.backward() optimizer.step() optimizer.zero_grad() adapted_model.eval() return adapted_model def solve_task(self, task: Dict, use_ttt: bool = True, use_poe: bool = True) -> List[List[List[int]]]: """ Full solve pipeline: 1. Apply TTT (per-task LoRA) 2. Augmented inference 3. Hierarchical voting or PoE scoring """ if use_ttt and self.model is not None: # Apply TTT adapted_model = self.apply_ttt(task) else: adapted_model = self.model if adapted_model is None: return [] # Augmented inference candidates = augmented_inference(adapted_model, self.tokenizer, task) # Hierarchical voting voted = hierarchical_vote(candidates, top_k=5) if use_poe and len(voted) > 2: # Score with Product of Experts scored = product_of_experts_score( adapted_model, self.tokenizer, task, voted ) return [s[0] for s in scored[:2]] return voted[:2] # ============================================================ # Test # ============================================================ if __name__ == "__main__": from arc_data import load_arc_dataset_from_hf print("Loading ARC-AGI-2 tasks...") tasks = load_arc_dataset_from_hf("arc-agi-community/arc-agi-2", "train") # Test TTT dataset construction task = tasks[0] ttt_data = build_ttt_training_data(task, max_examples=50) print(f"TTT dataset: {len(ttt_data)} examples") if ttt_data: inp, tgt = ttt_data[0] print(f" Input length: {len(inp)} chars") print(f" Target length: {len(tgt)} chars") print(f" Sample input:\n{inp[:300]}...") print(f" Sample target:\n{tgt}") # Test format formatted = format_task_for_model(task, include_test_output=True) print(f"\nFormatted task length: {len(formatted)} chars") # Test grid parsing test_grid_str = "1 2 3\n4 5 6\n7 8 9" parsed = parse_grid_from_output(test_grid_str) print(f"Parsed grid: {parsed}") assert parsed == [[1,2,3],[4,5,6],[7,8,9]] print("\n✅ TTT engine tests passed!")