| """ |
| 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 |
| ) |
|
|
|
|
| |
| |
| |
|
|
| 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, |
| "|": 12, |
| "<bos>": 13, "<eos>": 14, "<pad>": 15, |
| "<demo_start>": 16, "<demo_end>": 17, |
| "<test_start>": 18, "<test_end>": 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"<demo>") |
| parts.append(f"I:\n{tokenize_grid(pair['input'])}") |
| parts.append(f"O:\n{tokenize_grid(pair['output'])}") |
| parts.append(f"</demo>") |
| |
| parts.append(f"<test>") |
| 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() |
| |
| |
| for marker in ["</test>", "<eos>", "<pad>"]: |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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: |
| |
| 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 |
| |
| |
| 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 |
| |
| |
| 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] |
|
|
|
|
| |
| |
| |
|
|
| 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) |
| |
| |
| 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: |
| |
| 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, |
| }) |
| |
| |
| 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: |
| |
| 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 [] |
| |
| |
| def grid_key(g): |
| return tuple(tuple(r) for r in g) |
| |
| |
| 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(): |
| |
| freq = Counter(grid_key(p) for p in preds) |
| |
| for key, count in freq.most_common(3): |
| stage1_candidates.append({ |
| "prediction": [list(r) for r in key], |
| "intra_count": count, |
| "transform": transform_name, |
| }) |
| |
| |
| 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"] |
| |
| |
| results = [] |
| for key, count in global_freq.most_common(top_k): |
| results.append(global_grids[key]) |
| |
| return results |
|
|
|
|
| |
| |
| |
|
|
| 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]: |
| |
| aug_task = copy.deepcopy(task) |
| aug_task["test"][0]["output"] = candidate |
| aug_task = augment_task(aug_task, transform_name=t_name) |
| |
| |
| aug_candidate = aug_task["test"][0]["output"] |
| |
| |
| prompt = format_task_for_model(aug_task, include_test_output=False) + "\n" |
| target = tokenize_grid(aug_candidate) |
| full_text = prompt + target |
| |
| |
| 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 |
| |
| |
| 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)) |
| |
| |
| scored.sort(key=lambda x: -x[1]) |
| return scored |
|
|
|
|
| |
| |
| |
|
|
| 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 |
| |
| |
| ttt_data = build_ttt_training_data(task, max_examples=self.max_ttt_examples) |
| |
| if not ttt_data: |
| return self.model |
| |
| |
| 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 = torch.optim.AdamW( |
| adapted_model.parameters(), |
| lr=self.ttt_lr, |
| weight_decay=0.01 |
| ) |
| |
| |
| 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()} |
| |
| |
| 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 |
| |
| 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: |
| |
| adapted_model = self.apply_ttt(task) |
| else: |
| adapted_model = self.model |
| |
| if adapted_model is None: |
| return [] |
| |
| |
| candidates = augmented_inference(adapted_model, self.tokenizer, task) |
| |
| |
| voted = hierarchical_vote(candidates, top_k=5) |
| |
| if use_poe and len(voted) > 2: |
| |
| scored = product_of_experts_score( |
| adapted_model, self.tokenizer, task, voted |
| ) |
| return [s[0] for s in scored[:2]] |
| |
| return voted[:2] |
|
|
|
|
| |
| |
| |
|
|
| 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") |
| |
| |
| 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}") |
| |
| |
| formatted = format_task_for_model(task, include_test_output=True) |
| print(f"\nFormatted task length: {len(formatted)} chars") |
| |
| |
| 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!") |
|
|