ARC-AGI / trm_solver /trm_nano /train_task.py
rogermt's picture
Fix: force unbuffered stdout for Kaggle notebook output + use print() with flush"
0e4cf22 verified
Raw
History Blame Contribute Delete
18.4 kB
#!/usr/bin/env python -u
"""
TRM-Nano Batch Runner for NeuroGolf 2026.
Usage (2 GPUs):
python -u train_task.py --num_gpus 2 --batch_size 32 --log_every 50 --log_batch_every 10
Single task:
python -u train_task.py --task_file task001.json --log_every 1 --log_batch_every 5
NOTE: Always run with 'python -u' or set PYTHONUNBUFFERED=1 for Kaggle output.
"""
# Force unbuffered stdout/stderr so output appears immediately in Kaggle
import sys
import os
os.environ["PYTHONUNBUFFERED"] = "1"
if hasattr(sys.stdout, 'reconfigure'):
sys.stdout.reconfigure(line_buffering=True)
import json
import glob
import time
import argparse
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
from datetime import datetime
import multiprocessing as mp
from model import TRMNano, TRMNanoConfig, count_params
from onnx_builder import build_trm_nano_onnx
import onnx
import onnxruntime as ort
# ─── Logging (uses print with flush for Kaggle compatibility) ─────────────────
def log(msg, file=None):
"""Print with flush + optionally write to log file."""
ts = datetime.now().strftime("%H:%M:%S")
line = f"{ts} | {msg}"
print(line, flush=True)
if file:
file.write(line + "\n")
file.flush()
# ─── ARC Data Augmentation ────────────────────────────────────────────────────
def dihedral_transform(grid, trans_id):
if trans_id == 0: return grid
if trans_id == 1: return np.rot90(grid, 1)
if trans_id == 2: return np.rot90(grid, 2)
if trans_id == 3: return np.rot90(grid, 3)
if trans_id == 4: return np.fliplr(grid)
if trans_id == 5: return np.flipud(grid)
if trans_id == 6: return np.rot90(np.fliplr(grid), 1)
if trans_id == 7: return np.rot90(np.flipud(grid), 1)
return grid
def augment_pair(inp, out):
trans_id = np.random.randint(0, 8)
mapping = np.concatenate([[0], np.random.permutation(np.arange(1, 10))])
return dihedral_transform(mapping[inp], trans_id), dihedral_transform(mapping[out], trans_id), trans_id, mapping
def grid_to_tokens(grid):
h, w = grid.shape
offset_h = np.random.randint(0, max(1, 30 - h + 1))
offset_w = np.random.randint(0, max(1, 30 - w + 1))
padded = np.zeros((30, 30), dtype=np.int64)
padded[offset_h:offset_h+h, offset_w:offset_w+w] = grid + 2
eos_row, eos_col = offset_h + h, offset_w + w
if eos_row < 30:
padded[eos_row, offset_w:eos_col] = 1
if eos_col < 30:
padded[offset_h:eos_row, eos_col] = 1
return padded.flatten()
# ─── Dataset ──────────────────────────────────────────────────────────────────
class ARCTaskDataset(Dataset):
def __init__(self, task_data, num_augmentations=1000):
self.pairs = []
for pair in task_data.get("train", []):
self.pairs.append((np.array(pair["input"], dtype=np.int64),
np.array(pair["output"], dtype=np.int64)))
self.num_augmentations = num_augmentations
self.total_size = len(self.pairs) * (1 + num_augmentations)
def __len__(self):
return self.total_size
def __getitem__(self, idx):
pair_idx = idx % len(self.pairs)
aug_idx = idx // len(self.pairs)
inp, out = self.pairs[pair_idx]
if aug_idx > 0:
inp, out, _, _ = augment_pair(inp, out)
return (torch.tensor(grid_to_tokens(inp), dtype=torch.long),
torch.tensor(grid_to_tokens(out), dtype=torch.long))
# ─── Training ─────────────────────────────────────────────────────────────────
def train_on_task(task_data, config, epochs=500, lr=1e-3, batch_size=64,
num_aug=1000, device="cuda", log_every=100, log_batch_every=0, logfile=None):
model = TRMNano(config).to(device)
dataset = ARCTaskDataset(task_data, num_augmentations=num_aug)
loader = DataLoader(dataset, batch_size=batch_size, shuffle=True,
num_workers=0, pin_memory=True, drop_last=True)
optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=0.1, betas=(0.9, 0.95))
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
ema_decay = 0.999
ema_params = {n: p.data.clone() for n, p in model.named_parameters()}
best_loss = float("inf")
best_state = None
global_batch = 0
for epoch in range(epochs):
model.train()
total_loss, n_batches = 0, 0
epoch_start = time.time()
for inp_tokens, out_tokens in loader:
batch_start = time.time()
inp_tokens, out_tokens = inp_tokens.to(device), out_tokens.to(device)
logits = model(inp_tokens)
loss = F.cross_entropy(logits.view(-1, config.vocab_size), out_tokens.view(-1), ignore_index=0)
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
with torch.no_grad():
for n, p in model.named_parameters():
ema_params[n].mul_(ema_decay).add_(p.data, alpha=1 - ema_decay)
total_loss += loss.item()
n_batches += 1
global_batch += 1
if log_batch_every > 0 and global_batch % log_batch_every == 0:
bt = (time.time() - batch_start) * 1000
log(f" batch {global_batch} | loss={loss.item():.4f} | {bt:.0f}ms", logfile)
scheduler.step()
avg_loss = total_loss / max(n_batches, 1)
epoch_time = time.time() - epoch_start
if avg_loss < best_loss:
best_loss = avg_loss
best_state = {n: p.clone() for n, p in ema_params.items()}
if (epoch + 1) % log_every == 0:
log(f" epoch {epoch+1}/{epochs} loss={avg_loss:.4f} best={best_loss:.4f} ({n_batches}b, {epoch_time:.1f}s)", logfile)
if best_state:
with torch.no_grad():
for n, p in model.named_parameters():
p.copy_(best_state[n])
return model, best_loss
# ─── Evaluation ───────────────────────────────────────────────────────────────
def verify_with_onnx(onnx_path, task_data):
try:
sess = ort.InferenceSession(onnx_path)
except Exception:
return 0, 999
passes, fails = 0, 0
for split in ["train", "test", "arc-gen"]:
for example in task_data.get(split, []):
if "output" not in example:
continue
inp_grid = np.array(example["input"], dtype=np.int64)
out_grid = np.array(example["output"], dtype=np.int64)
h, w = inp_grid.shape
if h > 30 or w > 30:
fails += 1
continue
inp_onehot = np.zeros((1, 10, 30, 30), dtype=np.float32)
for r in range(h):
for c in range(w):
inp_onehot[0, inp_grid[r, c], r, c] = 1.0
result = sess.run(None, {"input": inp_onehot})
pred_onehot = (result[0] > 0.0).astype(np.float32)
oh, ow = out_grid.shape
exp_onehot = np.zeros((1, 10, 30, 30), dtype=np.float32)
for r in range(oh):
for c in range(ow):
exp_onehot[0, out_grid[r, c], r, c] = 1.0
if np.array_equal(pred_onehot, exp_onehot):
passes += 1
else:
fails += 1
return passes, fails
# ─── ONNX Export ──────────────────────────────────────────────────────────────
def export_onnx_manual(model, config, output_path):
onnx_model = build_trm_nano_onnx(model, config)
onnx.save(onnx_model, output_path)
return os.path.getsize(output_path)
# ─── Run tasks on a device ────────────────────────────────────────────────────
def run_tasks(tasks, args, config, device, logfile=None):
results = []
solved_count = 0
for i, (task_id, task_num, task_file) in enumerate(tasks):
log(f"[{i+1}/{len(tasks)}] {task_id}", logfile)
try:
with open(task_file) as f:
task_data = json.load(f)
n_train = len(task_data.get("train", []))
n_test = len(task_data.get("test", []))
n_arcgen = len(task_data.get("arc-gen", []))
log(f" train={n_train}, test={n_test}, arc-gen={n_arcgen}", logfile)
t0 = time.time()
model, best_loss = train_on_task(
task_data, config, epochs=args.epochs, lr=args.lr,
batch_size=args.batch_size, num_aug=args.num_aug,
device=device, log_every=args.log_every,
log_batch_every=args.log_batch_every, logfile=logfile)
train_time = time.time() - t0
onnx_path = os.path.join(args.output_dir, f"{task_id}.onnx")
size_bytes = export_onnx_manual(model, config, onnx_path)
passes, fails = verify_with_onnx(onnx_path, task_data)
solved = (fails == 0 and passes > 0)
if solved:
solved_count += 1
status = "βœ“ SOLVED" if solved else f"βœ— {passes}p/{fails}f"
log(f" {status} | loss={best_loss:.4f} | {size_bytes/1024:.0f}KB | {train_time:.0f}s", logfile)
results.append({"task_id": task_id, "task_num": task_num, "solved": solved,
"passes": passes, "fails": fails, "best_loss": best_loss,
"onnx_bytes": size_bytes, "train_time_s": train_time})
if not solved and os.path.exists(onnx_path):
os.remove(onnx_path)
except Exception as e:
log(f" ERROR: {e}", logfile)
results.append({"task_id": task_id, "task_num": task_num, "solved": False, "error": str(e)})
# Save progress
with open(os.path.join(args.output_dir, "progress.json"), "w") as f:
json.dump({"results": results, "solved": solved_count, "total": len(tasks)}, f, indent=2)
return results, solved_count
def gpu_worker_fn(gpu_id, tasks, args, config):
device = f"cuda:{gpu_id}"
log_path = os.path.join(args.output_dir, f"log_gpu{gpu_id}.txt")
with open(log_path, "w") as logfile:
log(f"GPU{gpu_id} started on {device} with {len(tasks)} tasks", logfile)
results, solved = run_tasks(tasks, args, config, device, logfile)
with open(os.path.join(args.output_dir, f"progress_gpu{gpu_id}.json"), "w") as f:
json.dump({"results": results, "solved": solved, "total": len(tasks)}, f, indent=2)
log(f"GPU{gpu_id} done: {solved}/{len(tasks)} solved", logfile)
# ─── Main ─────────────────────────────────────────────────────────────────────
def main():
parser = argparse.ArgumentParser(description="TRM-Nano Batch Trainer for NeuroGolf")
parser.add_argument("--data_dir", type=str, default="/kaggle/input/competitions/neurogolf-2026/")
parser.add_argument("--skip_solved", type=str, default=None)
parser.add_argument("--task_file", type=str, default=None)
parser.add_argument("--output_dir", type=str, default="/kaggle/working/trm_onnx/")
parser.add_argument("--epochs", type=int, default=500)
parser.add_argument("--lr", type=float, default=1e-3)
parser.add_argument("--batch_size", type=int, default=64)
parser.add_argument("--num_aug", type=int, default=500)
parser.add_argument("--hidden_size", type=int, default=96)
parser.add_argument("--H_cycles", type=int, default=3)
parser.add_argument("--L_cycles", type=int, default=4)
parser.add_argument("--log_every", type=int, default=100,
help="Log epoch summary every N epochs (default: 100)")
parser.add_argument("--log_batch_every", type=int, default=0,
help="Log per-batch loss every N batches (0=off)")
parser.add_argument("--max_tasks", type=int, default=None)
parser.add_argument("--start_task", type=int, default=1)
parser.add_argument("--num_gpus", type=int, default=1)
parser.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu")
args = parser.parse_args()
os.makedirs(args.output_dir, exist_ok=True)
config = TRMNanoConfig(
hidden_size=args.hidden_size,
num_heads=max(2, args.hidden_size // 24),
H_cycles=args.H_cycles,
L_cycles=args.L_cycles,
)
# ─── Single task mode ─────────────────────────────────────────────────────
if args.task_file:
task_id = os.path.basename(args.task_file).replace(".json", "")
log(f"Single task: {task_id}")
log(f"Config: h={config.hidden_size} T={config.H_cycles} n={config.L_cycles} params={count_params(TRMNano(config)):,}")
with open(args.task_file) as f:
task_data = json.load(f)
t0 = time.time()
model, best_loss = train_on_task(task_data, config, epochs=args.epochs, lr=args.lr,
batch_size=args.batch_size, num_aug=args.num_aug,
device=args.device, log_every=args.log_every,
log_batch_every=args.log_batch_every)
log(f" Trained in {time.time()-t0:.1f}s, loss={best_loss:.4f}")
onnx_path = os.path.join(args.output_dir, f"{task_id}.onnx")
size_bytes = export_onnx_manual(model, config, onnx_path)
log(f" ONNX: {size_bytes:,} bytes")
passes, fails = verify_with_onnx(onnx_path, task_data)
log(f" Verify: {passes} pass, {fails} fail {'βœ“ SOLVED!' if fails==0 else 'βœ—'}")
return
# ─── Batch mode ───────────────────────────────────────────────────────────
log("=" * 60)
log("TRM-Nano Batch Trainer for NeuroGolf 2026")
log("=" * 60)
log(f"Config: h={config.hidden_size}, heads={config.num_heads}, T={config.H_cycles}, n={config.L_cycles}")
log(f"Params: {count_params(TRMNano(config)):,}")
log(f"GPUs: {args.num_gpus}, Batch: {args.batch_size}, Epochs: {args.epochs}")
log(f"Log: every {args.log_every} epochs, every {args.log_batch_every} batches")
skip_tasks = set()
if args.skip_solved and os.path.exists(args.skip_solved):
with open(args.skip_solved) as f:
skip_data = json.load(f)
if isinstance(skip_data, list):
skip_tasks = set(skip_data)
elif isinstance(skip_data, dict):
skip_tasks = set(skip_data.keys())
log(f"Skipping {len(skip_tasks)} solved tasks")
task_files = sorted(glob.glob(os.path.join(args.data_dir, "task*.json")))
if not task_files:
task_files = sorted(glob.glob(os.path.join(args.data_dir, "*.json")))
tasks_to_run = []
for f in task_files:
task_id = os.path.basename(f).replace(".json", "")
task_num = None
try:
task_num = int(task_id.replace("task", ""))
except ValueError:
pass
if task_id in skip_tasks or str(task_num) in skip_tasks or task_num in skip_tasks:
continue
if task_num is not None and task_num < args.start_task:
continue
tasks_to_run.append((task_id, task_num, f))
if args.max_tasks:
tasks_to_run = tasks_to_run[:args.max_tasks]
log(f"Tasks: {len(tasks_to_run)}")
# ─── Dispatch ─────────────────────────────────────────────────────────────
num_gpus = min(args.num_gpus, torch.cuda.device_count()) if torch.cuda.is_available() else 1
if num_gpus > 1:
log(f"Parallel: {num_gpus} GPUs")
chunks = [[] for _ in range(num_gpus)]
for i, task in enumerate(tasks_to_run):
chunks[i % num_gpus].append(task)
mp.set_start_method("spawn", force=True)
processes = []
for gpu_id in range(num_gpus):
p = mp.Process(target=gpu_worker_fn, args=(gpu_id, chunks[gpu_id], args, config))
p.start()
processes.append(p)
for p in processes:
p.join()
# Merge
all_results, total_solved = [], 0
for gpu_id in range(num_gpus):
pf = os.path.join(args.output_dir, f"progress_gpu{gpu_id}.json")
if os.path.exists(pf):
with open(pf) as f:
d = json.load(f)
all_results.extend(d["results"])
total_solved += d["solved"]
with open(os.path.join(args.output_dir, "progress.json"), "w") as f:
json.dump({"results": all_results, "solved": total_solved, "total": len(tasks_to_run)}, f, indent=2)
log(f"ALL DONE: {total_solved}/{len(tasks_to_run)} solved")
else:
log(f"Single device: {args.device}")
results, solved_count = run_tasks(tasks_to_run, args, config, args.device)
log(f"DONE: {solved_count}/{len(tasks_to_run)} solved")
onnx_files = sorted(glob.glob(os.path.join(args.output_dir, "*.onnx")))
total_kb = sum(os.path.getsize(f) for f in onnx_files) / 1024
log(f"ONNX files: {len(onnx_files)} ({total_kb:.1f} KB)")
if __name__ == "__main__":
main()