토큰마다 반복 깊이를 정하는 적응형 루프 트랜스포머 TaH2

Improving Test-Time Scaling with Adaptive Looped Transformers

arXiv2609.35748v1

Yichen You2026-09-28조회 3

무엇인가

이 논문은 루프 트랜스포머(looped transformer)가 테스트 시점 스케일링에서 실제로 이득을 주는지를 따진다. 루프 트랜스포머는 같은 레이어를 여러 번 재사용해 파라미터를 늘리지 않고 유효 깊이를 키우는 구조로, Huginn은 중간 블록을, Ouro는 전체 스택을 반복한다. 기존 연구는 파라미터 수나 토큰당 FLOPs를 맞춘 상태에서 루프 모델과 비루프 모델을 비교했지만, 출력이 길어질수록 루프가 정확도-연산 스케일링을 개선하는지는 다뤄지지 않았다. 저자들은 사전학습된 체크포인트에 반복 구조를 사후 도입하는 포스트트레이닝 방식으로 이 질문을 다룬다.

어떻게 동작하나

Qwen3-1.7B-Base를 같은 데이터로 포스트트레이닝해 AIME24–26을 4K~16K 출력 토큰 컷오프에서 평가한 결과, 디코딩 FLOPs가 두 배가 될 때의 정확도 상승폭(기울기)은 고정 깊이 Ouro 2.12, 적응형 Ouro 2.27, Huginn 2.26으로 비루프 베이스라인 Standard의 1.79보다 가팔랐다. 그러나 겹치는 연산 구간에서 세 모델 모두 Standard보다 정확도가 낮았다. 검증셋에서 첫 반복과 마지막 반복의 다음 토큰 손실을 비교해 보니, Ouro는 52.3%, Huginn은 33.5%의 토큰이 손실 변화가 10^-3 이하로 사실상 그대로였고, 각각 21.7%와 15.9%는 오히려 손실이 나빠졌다. 고정 깊이 루프는 이득이 없는 토큰에도 반복 연산을 쓰고 있었다.

무엇과 다른가

TaH2는 공유 Transformer 백본, 학습된 입력 주입 업데이터, 토큰 단위 반복 결정기(decider)로 구성된다. 업데이터와 결정기는 연구한 모든 스케일에서 파라미터의 3% 미만을 차지한다. 첫 반복은 토큰 임베딩을 그대로 입력으로 받고, 이후 반복은 소형 정규화 MLP인 업데이터가 임베딩을 다시 주입하면서 이전 반복의 최종 레이어 상태를 백본 입력 공간으로 되돌린다. 각 반복이 끝날 때 결정기는 임베딩, 은닉 상태, top-K 출력 확률을 받아 계속 진행할 조건부 확률을 내놓고, 이 확률이 임계값(τ_exit=0.5) 아래로 떨어지면 그 토큰은 멈춘다. 정지 확률들은 실행된 깊이들에 대한 가중치가 되어 출력 확률을 가중 혼합하며, 마지막으로 실행된 반복이 남은 확률 질량을 흡수한다. 학습 시에는 확장된 duo-causal 어텐션을 써서 깊이 m의 토큰이 위치 s≤t, 깊이 j≤m의 실행된 KV 상태만 보도록 하고, 멈춘 토큰에는 그래디언트 없는 추가 반복을 한 번 돌려 감독 신호를 만든다.

어떻게 쓰나

핵심 학습 기법은 룩어헤드 깊이 감독(lookahead depth supervision)이다. 각 반복에서 다음 토큰 예측 손실의 감소량 δ를 온라인으로 측정해, 그 토큰이 계속 반복할 가치가 있는지를 나타내는 라벨을 만든다. 양의 이득은 계속, 음의 이득은 정지를 뜻한다. 이득이 양수인 토큰들을 큰 값부터 정렬해 전체 이득의 ρ=0.99를 차지할 때까지 고르고, 선택된 가장 작은 이득을 컷오프로 삼아 그 이상인 토큰만 계속 라벨을 받는다. 한 번 정지 라벨을 받으면 이후 라벨은 0으로 고정된다. 목적함수는 다음 토큰 예측 손실에, 컷오프와의 거리로 가중한 비용 민감 BCE 결정기 손실을 계수 α_D=0.05로 더한 형태다. 깊이 선택은 학습 중에도 추론과 같은 규칙으로 결정기의 결정을 따르는 온폴리시 방식이다. Ouro가 전체 깊이로 학습한 뒤 추론에서 조기 종료를 적용해 학습-추론 불일치를 겪고, TaH가 오라클 정책으로 학습하는 것과 대비된다.

전제와 한계

실험은 Qwen3-1.7B/4B/8B-Base를 백본으로 쓴다. 학습 데이터는 AM-Qwen3-Distilled의 수학·코드·QA 프롬프트와 Nemotron-Agentic-v1의 툴 사용 샘플을 합친 것으로, 1.7B 실험은 273K 프롬프트에 Qwen3-8B가 재생성한 응답을 붙여 16,384 토큰 문맥으로 3에폭, 총 3.4B 토큰을 학습했다. 평가는 AIME24–26, AMC23, MATH500, OlympiadBench, IMO-AnswerBench, GPQA, SuperGPQA, HumanEval, MBPP, LiveCodeBench v6, BFCL v3에서 이뤄졌다. 1.7B, M=2에서 TaH2는 디코딩 FLOPs 두 배당 2.74점을 얻어 Standard의 1.79점보다 53% 가파른 기울기를 보였고, 평가를 32K까지 늘리면 Standard가 191.7 TFLOPs에서 12.0%로 포화되는 지점에서 TaH2는 같은 연산으로 15.4%를 기록해 3.4점 앞섰다. 최대 반복 깊이를 2에서 8로 올리면 AIME24–26에서 Standard 대비 이득이 +2.8점에서 +3.9점으로 커졌고, 10개 벤치마크 평균 이득은 +2.9점에서 +4.8점으로 늘었다. 다수결 투표(cons@32)에서도 27.3~29.3%로 Standard의 21.9%를 넘었다. 4B에서는 평균 3.2점(AIME 최대 6.9점), 8B에서는 평균 2.4점(AIME 최대 4.4점) 향상됐다.

설계 선택 실험에서 결정기 손실 가중치를 이득 크기 대신 균일하게 두면 평균 정확도가 3.6점 떨어져 가장 큰 손실을 냈고, 학습된 업데이터를 top-100 토큰 임베딩의 확률 가중합으로 바꾸면 2.1점, 모든 양의 이득을 유지(ρ=1)하면 1.9점, 학습 중 결정을 임계값 대신 샘플링하면 1.3점, top-1 불일치 라벨을 쓰면 1.1점, 최종 예측만 쓰면 1.0점 각각 하락했다. 분석에서는 계속 확률이 0에 가까운 토큰은 두 번째 반복의 손실 감소가 음수이거나 무시할 수준이고 확률이 높아질수록 평균 이득이 커져, 결정기가 실제로 이득이 큰 토큰에 반복을 몰아준다는 것을 보였다. 토큰별 깊이를 시각화하면 수학·코드 예시에서 수식과 최종 코드가 앞선 자연어 추론보다 반복을 적게 쓰고, QA 예시는 전반적으로 깊이를 유지했다.

실무 관점에서 이 논문은 이미 사전학습된 모델에 반복 구조를 얹어 추론 연산을 늘리는 선택지를 다룬다. 다만 TaH2는 토큰당 디코딩 FLOPs를 22%, 종단 지연을 30~34% 늘리므로, 지연에 민감한 서빙에서는 이득이 큰지 따져야 한다. 저자들은 서로 다른 반복 깊이의 요청을 한 번의 포워드로 배칭하는 확장 Mini-SGLang 엔진으로 서빙했고, 그 조건에서도 Standard보다 나은 테스트 시점 스케일링과 정확도를 얻었다고 보고한다. 도입을 검토한다면 백본 학습과 결정기 학습을 분리하지 않는 온폴리시 깊이 선택, 그리고 이득 기반 라벨이 실제 워크로드에서도 유효한지 확인하는 것이 핵심이다.

저자들이 밝힌 한계는 두 가지다. 첫째, TaH2는 표준 SFT보다 학습 FLOPs가 더 든다. 다만 포스트트레이닝은 사전학습보다 연산이 훨씬 적게 들어 기존 모델에 적응형 깊이를 붙이는 현실적인 방법이라고 주장한다. 둘째, 방법은 SFT 환경에서만 연구됐고 온폴리시 증류나 강화학습으로의 확장은 향후 과제로 남겨졌다. 학습·평가 코드와 체크포인트는 공개 예정이라고 밝혔지만, 원문에 실린 링크는 익명 처리되어 있어 저장소 주소는 확인할 수 없다.