확산 언어모델에서 결정론적 PRM 유도는 같은 연산량의 ORM 재정렬에 진다
Why Deterministic PRM Guidance Underperforms in Discrete Diffusion Reasoning
무엇인가
확산 언어모델(dLLM)은 한 번에 한 토큰씩 확정하는 자기회귀 디코딩과 달리, 부분적으로 마스킹된 전체 시퀀스를 유지하면서 조금씩 복원한다. 그래서 매 디노이징 단계마다 완성되지 않은 해답 상태가 드러나고, 이 중간 상태를 프로세스 리워드 모델(PRM)로 채점해 유망한 궤적을 살려두면 테스트 시점 연산을 추론 정확도로 바꿀 수 있어 보인다. 문제는 자기회귀용 PRM이 왼쪽에서 오른쪽으로 자라는 접두사를 전제하는데, dLLM의 스냅샷은 앞부분이 아직 가려진 채 뒷부분이 먼저 드러나는 흩어진 위치 부분집합이라는 점이다. 이 논문은 그 불일치가 실제로 얼마나 손해인지를 같은 연산 예산 위에서 측정한다.
어떻게 동작하나
비교의 핵심은 forward pass 회계다. 시퀀스 길이와 무관하게 디노이징 스텝 T=128로 두면 바닐라 샘플 하나는 128패스, ORM Rerank@N은 128N+N패스가 든다. PRM Guided는 b스텝마다 K개로 분기해 전부 디노이징하고 전부 채점한 뒤 최고점 하나만 남기므로 KT+K⌈T/b⌉패스가 든다. b=64일 때 K=8은 1,040패스, K=32는 4,160패스이고, 맞춰둔 ORM Rerank 예산은 각각 1,032와 4,128패스다. 약 0.8% 차이로 오히려 PRM Guided 쪽에 유리하게 잡혀 있다. PRM 자체는 Dream-v0-Instruct-7B 백본을 얼린 채 LoRA 어댑터와 2층 MLP 리워드 헤드, 256차원 sinusoidal 스텝 인덱스 임베딩만 학습한 모델로, 최종 정답의 정오 라벨을 중간 상태에 물려주는 방식(outcome-supervised intermediate-state value model)으로 훈련된다. 양방향 어텐션 PRM과 causal 어텐션 PRM을 함께 비교한다.
무엇과 다른가
결과는 일관되게 ORM 재정렬의 승리다. Dream-7B, GSM8K에서 문제당 8개 후보일 때 PRM Guided는 65.18%, ORM Rerank는 75.13%로 9.95pp 차이다. 후보를 32개로 늘리면 70.02% 대 82.71%로 격차가 12.69pp까지 벌어진다. 과제별 검증기를 따로 학습한 MATH에서는 20.80% 대 30.65%(9.85pp), MBPP에서는 50.88% 대 63.04%(12.16pp)다. PRM Guided가 Majority@N은 앞서므로 PRM에 신호가 아예 없는 것은 아니지만, 같은 연산을 독립 샘플링과 최종 상태 검증기에 쓰는 편이 항상 낫다. 벽시계 시간으로 맞춰도 순서는 유지된다. ORM Rerank@6이 문제당 약 126초에 72.40%를 내는 동안 PRM Guided K=8은 156초를 쓰고 65.18%에 머문다. 참고로 ORM Rerank@32의 82.71%도 Oracle@32의 91.13%에는 8.4pp 못 미치므로, 최종 검증기 쪽에도 개선 여지가 남아 있다.
어떻게 쓰나
격차의 첫 번째 원인은 유도가 후보 풀 자체를 망가뜨린다는 것이다. 마스크 비율이 높아질수록 PRM의 ROC-AUC가 0.77에서 0.54까지 단조 하락한다. 논문은 이것이 학습 결함이 아니라 정보량의 한계라고 설명한다. 상태와 최종 정오 사이의 상호정보량이 줄면 어떤 채점기든 ROC-AUC 상한이 1/2+√(I(x_t;y)/(2π₀π₁))로 내려간다는 명제를 제시한다. 라벨을 새 롤아웃 8개로 다시 붙여 1만 개 상태를 재라벨링해 재학습해도 pooled ROC-AUC는 0.011, GSM8K 정확도는 0.61pp밖에 오르지 않는다. 결정론적 top-1 가지치기는 다양성도 무너뜨린다. K=N=8에서 PRM Hybrid는 문제당 고유 답이 1.75개로 독립 샘플링의 4.31개에 크게 못 미치고 답 엔트로피는 4분의 1로 줄며, Oracle@8이 81.05%에서 67.30%로 13.75pp 떨어진다. 즉 단순 샘플링이면 살아남았을 정답 후보를 유도가 미리 잘라낸다. 저장된 궤적에 대한 오프라인 반사실 분석에서도 top-1 컷은 초기 상태에서 정답 계보를 46%, 최종 상태에서 20% 확률로 제거하며, 4개를 남기면 각각 10%와 6%로 줄어든다.
전제와 한계
두 번째 원인은 최종 선택의 품질이다. 같은 1,040패스로 K=8 입자를 PRM 점수로 리샘플링하는 SMC 샘플러를 쓰면 Oracle@8이 67.30%에서 77.89%로 회복되어, top-1 가지치기가 잃은 13.75pp 중 10.59pp를 되찾는다. 그런데도 PRM이 고른 답의 정확도는 가중 투표로 65.48%, 최고 점수 입자로 66.34%에 그쳐 결정론적 유도와 같은 수준이고 ORM Rerank@8의 75.13%에는 크게 못 미친다. 풀을 복구해도 최종 선택이 나아지지 않는다는 뜻이다. 같은 후보 풀을 놓고 재정렬만 시키면 최종 상태로 재학습한 PRM은 ORM과 동등해지지만, 마스크 전 구간으로 학습한 cross-mask PRM은 N=8과 N=32에서 각각 32.3pp, 17.4pp 뒤진다. N=8에서는 무작위 선택과 다를 바 없다. 논문은 pooled ROC-AUC가 1에 가까워도 문제 내 top-1 선택이 무작위와 같을 수 있다는 명제를 들어, 난이도만 구분하고 같은 문제 안 후보를 구분하지 못하는 채점기의 위험을 지적한다. MBPP는 두 실패를 깔끔하게 분리한다. 완성된 프로그램을 재정렬할 때 PRM은 65.47%로 ORM과 대등하지만, 디노이징을 유도할 때는 50.88%로 떨어진다. 여기서의 14.6pp 손실은 유도 자체의 비용이다.
부수 진단으로 causal PRM의 readout 문제도 다룬다. 최종 상태에서 평균 풀링 causal PRM의 ROC-AUC는 0.61로 양방향의 0.78에 못 미치지만, 자기회귀 리워드 모델의 표준인 마지막 토큰 풀링으로 바꾸면 0.73까지 올라 격차의 약 70%를 메운다. N=32 재정렬에서 평균 풀링 causal은 40.56%로 단일 샘플보다 낮아지고, 마지막 토큰 변형은 50.64%, 양방향 PRM은 65.35%, ORM은 82.71%다. LLaDA-8B-Base에서도 양방향 유도가 31.64%로 causal 유도의 22.25%, 바닐라 20.77%를 앞선다. 오른쪽 문맥을 그냥 읽어서 얻는 이득인지도 확인했는데, 스냅샷의 오른쪽 절반을 0으로 만들어도 ROC-AUC 변화는 0.001 수준이고, 입력을 뒤집어 causal PRM에 평균 내면 역방향 패스가 RoPE 백본에서 사실상 무작위(ROC-AUC 0.500)라 아무것도 얻지 못한다.
실무 관점의 결론은 다섯 가지 권고로 정리된다. 중간 상태 채점기와 최종 상태 검증기의 비용을 같은 forward pass 단위로 청구할 것, 정확도 옆에 Oracle@K를 함께 보고해 후보 풀 붕괴를 드러낼 것, 유도는 늦게 하고 후보를 하나 이상 남길 것, 최종 선택은 최종 상태로 학습한 검증기에 맡길 것, causal PRM을 쓴다면 readout부터 고칠 것이다. 논문은 이런 분해가 확산 모델에 국한되지 않는다고 본다. 중간 점수로 가지치기하고 같은 채점기로 최종 선택하는 자기회귀 PRM 빔 서치도 풀 손상과 선택 오류로 똑같이 쪼개진다는 것이다.
저자들이 밝힌 한계도 분명하다. 핵심 결과는 Dream-7B와 결정론적 top-1 PRM Guided 레시피, GSM8K를 주 과제로 하고 MATH와 MBPP를 과제별 대조군으로 삼은 범위 안에 있다. SMC 대조군은 다양성을 보존하는 샘플러 하나만 다루므로 다른 확률적 디코더는 공개된 툴킷으로 벤치마크해야 한다. ORM의 우위는 과제에 맞춰 학습한 검증기를 전제로 한다. GSM8K로 학습한 ORM은 MATH500으로 전이되지 않는다. LLaDA에서 양방향 우위는 재현됐지만 LLaDA 전용 ORM 비교와 permutation-LM 검증은 다음 과제로 남겨두었다. 저자들은 정오 라벨이 붙은 디노이징 상태 코퍼스와 평가 툴킷을 공개해 같은 프로토콜에서 재현 가능한 비교를 할 수 있게 했다.