March 10, 2025 — 📝 Paper

[논문 리뷰] Twist-SMC #1 Introduction


1. Introduction

LLM은 기본적으로 ARM(Auto-Regressive Model) 즉, “다음 토큰 확률 분포 예측기” 이다.

그렇다는 것은, ‘LLM의 다음 단어 확률’에 ‘원하는 대로 말하게 하는 녀석’을 붙이면, 우리가 원하는 대로 대답 할 수 있는 LLM이 된다는 것이다.

논문에서는 이를 “Model의 Output을 Steering 한다” 라고 이야기 한다.

논문에서는 Steering 하는 방법에 RLHF, Prompt Engineering 등이 있다고 한다.


그리고는 논문에서 다음과 같은 핵심 문장이 등장한다.

We view the above tasks as instances of probabilistic inference: sampling from a target unnormalized density and estimating its intractable (log) normalization constant

위의 작업은 정규화되지 않은 대상 밀도로부터 샘플링하고 그 난해한 (로그) 정규화 상수를 추정하는 확률적 추론의 인스턴스로 간주합니다.

정말 어려운 말이다.

이제부터 우리는 저 문장을 이해하는 여정을 떠나보도록 하겠다.


여기서 논문의 제목(Probabilistic Inference in Language Model via Twisted SMC)에도 포함된 단어인 Probabilistic inference라는 단어가 등장한다.

Probabilistic Inference (확률적 추론)


그렇다면 우리는 논문 속 Probabilistic Inference에 대한 설명을 다시 보면, 이 논문에서 뭘 하고 싶은 건 지에 대한 “윤곽”을 잡을 수 있게 된다.

sampling from a target unnormalized density and estimating its intractable (log) normalization constant

‘sampling from a target unnormalized density’

→ ‘estimating its intractable (log) normalization constant’

크게 2 단계로 나누어 볼 수 있다.

Target unnormalized density에서 샘플링을 해서, Normalization constant를 추정한다는 것이다

그러면 과연 Target unnormalized density, Normalization constant는 뭔지를 알아보도록 하자.

일단 이 정도로 윤곽을 잡아볼 수 있겠다.


그러고는 중요한 (1)번 수식이 등장한다.

이 수식과 설명을 이해한다면, 우리는 Target unnormalized density, Normalization constant이 뭔지, 그리고 왜 구하는 건지 등을 알 수 있을 것이다.

우리가 알고 싶어 하던 Target unnormalized density, Normalization constant가 모두 등장했다


‘sampling from a target unnormalized density’

→ ‘estimating its intractable (log) normalization constant’

Target unnormalized density에서 샘플링을 해서, Normalization constant를 추정

드디어 우리는 Target unnormalized density, Normalization constant 를 알았다


결국 우리가 궁극적으로 알고 싶은 건 ⁍이다.

⁍은 결국 Target normalized distribution이자, 베이지안 추론 관점(?)으로는 사후 확률이라고도 할 수 있다.

결국에는 다 ⁍ 이거를 잘 알아내기 위해 하는 일이다.

사후 확률은 다음과 같이 정의되기 때문에, 결국 다음 3가지를 다 알아야 한다.

  1. ⁍: Partition Function or Normalization Constant
  2. ⁍: **Base Model **
  3. ⁍: Potential function

여기서 2번과 3번은 우리가 알 수 있다

결국 문제가 되는 건 1번 : Partition Function** **인 것이다

여기서 ⁍ = ⁍ 이다.

우리는 2번과 3번을 알 수 있기 때문에, 자연스레 ⁍: Unnormalized Target Distribution을 알고 있다.

결국에는 ⁍를 안다는 것은 ⁍를 모든 시퀀스에 대해서 알아야 한다는 것이고, 이게 어렵다는 것이다.

이것이 어려운 이유가 곧 논문에서 “문제”라고 이야기 하는 것인데, 이는 조금 뒤에 더 자세히 살펴보도록 하자.

그래서 논문에서는 Probabilistic Inference을 다음과 같이 이야기 했다.

‘sampling from a target unnormalized density’

→ ‘estimating its intractable (log) normalization constant’

⁍ = ⁍ 에서 샘플링 → ⁍을 추론

결국, 우리가 알 수 있는 ⁍: Unnormalized Target Distribution 에서 샘플링을 통해 ⁍를 를 “추론” 하겠다는 뜻이다.

샘플링이란?

샘플링을 통해 추론한다는 것에서 확률적 추론(Probabilistic Inference)인 것이고,

그래서 우리 논문의 제목이 **Probabilistic Inference in Language Models **via Twisted Sequential Monte Carlo 인 것이다.


무엇이 문제인가?

다시 한번 짚고 넘어가자

우리는 ⁍: Target Distribution를 알고 싶다

그러기 위해선 ⁍: **Partition Function or **Normalization Constant 를 알아야 한다

⁍를 안다는 것은, ⁍를 모든 시퀀스에 대해서 알아야 한다는 것이고, 이게 어렵다는 것이다.

그렇다면 ⁍를 모든 시퀀스에 대해서 알아야 한다는 것은 왜 어렵다는 걸까?

A particular challenge in sampling from Eq. (1) is that the target distribution ⁍ is non-causal.

바로 우리의 목표인 ⁍가 non-causal하기 때문이라고 한다.

non-causal 이란?

왜 ⁍가 non-causal 하다는 걸까?

바로 다음에 나온다.

In order to sample tokens sequentially, one needs to infer the marginal distribution ⁍ = ⁍ which involves an intractable marginalization.

⁍⁍

⁍ 이기 때문에,

결국 marginal distribution ⁍ 을 구하려면, ⁍ 까지 marginalization 해야 하는데, 이는 곧 1:t(초록 동그라미 들)를 위해 **미래의 정보인 t+1:T(분홍색 네모들) **까지를 사용해야 한다는 것이다.

그래서 non-causal 하다는 것이다.

그래서 ⁍가 non-causal하기 때문에, marginalization 하는 과정이 intractable(매우 더러운)하다는 것이다.

marginal distribution 이란?


그래서 논문에서는 ⁍가 아니라 미래정보를 사용하지 않고, ⁍ 를 이용해서 구할(추정?) 수 있도록 하는

이 엄청난 twist function인 ⁍를 지금부터 소개하겠다고 한다.


Introduction에서는 “문제”와 앞으로 얘기할 ”twist function인 ****⁍“에 대해서 간단히 설명하고 있다.

기본적인 틀을 잘 잡고 싶어 자세히 설명하다보니, 너무 길어져 버렸다.

다음은 2. Background이다.