arc-agi-2-solver / ttt_engine.py
Interstellar007's picture
Upload ttt_engine.py
f1cfb2f verified
Raw
History Blame Contribute Delete
17.6 kB
"""
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
"<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()
# Remove any markers
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
# ============================================================
# 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!")