Attention 백엔드

Attention 백엔드 (Attention backends)

모든 attention 구현은 동일한 계산을 수행해요. 모든 토큰이 다른 모든 토큰과 비교되지요. 차이는 어떻게 계산을 수행하느냐에 있어요. 기본 attention은 전체 attention 행렬을 메모리에 materialize하기 때문에 확장성이 떨어져서, 추론을 느리게 만드는 병목이 생겨요. 최적화된 구현은 수학을 재배열해서 메모리 트래픽을 줄여 더 빠르고 저렴한 추론을 가능하게 해요.

출처: 문서

본문

AttentionInterface는 최적화된 attention 구현을 제공해요. attention 구현을 모델 구현에서 분리해서 서로 다른 함수로 실험하기 쉽게 만들어 주지요. 이 일관된 인터페이스로 새 백엔드를 쉽게 추가할 수 있어요.

Attention 백엔드 설명
"flash_attention_3" FlashAttention-2를 개선해, 연산을 겹치고(overlapping) forward·backward 패스를 더 밀접하게 융합해요
"flash_attention_2" 계산을 더 작은 블록으로 타일링하고 빠른 on-chip 메모리를 사용해요
"flex_attention" 로우레벨 커널을 직접 작성하지 않고도 커스텀 attention 패턴(스파스, block-local, sliding window)을 지정하는 프레임워크예요
"sdpa" scaled dot product attention의 PyTorch 내장 구현이에요
"paged|flash_attention_3" FlashAttention-3의 Paged 버전
"paged|flash_attention_2" FlashAttention-2의 Paged 버전
"paged|sdpa" SDPA의 Paged 버전
"paged|eager" eager의 Paged 버전

Attention 백엔드 설정하기 (Set an attention backend)

특정 attention 함수로 모델을 인스턴스화하려면 from_pretrained()의 attn_implementation 인자를 사용해요.

import torch
from transformers import AutoModelForCausalLM

model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-3.2-1B", attn_implementation="flash_attention_2"
)

set_attn_implementation()을 사용하면 모델을 다시 로드하지 않고도 런타임에 attention 백엔드를 전환할 수 있어요.

model.set_attn_implementation("sdpa")

커널 (Kernels)

Kernels 라이브러리로 런타임에 Hub에서 컴파일된 연산 커널을 직접 다운로드·로드할 수 있어요. 이렇게 하면 PyTorch나 CUDA 버전 불일치로 인한 패키징 문제를 피할 수 있어요.

Kernels는 감지되면 자동으로 AttentionInterface에 등록돼요. FlashAttention 패키지를 명시적으로 설치할 필요가 없어요. 이름으로 FlashAttention을 요청해도 Hub 커널로 폴백되니, FlashAttention 폴백을 참고해 주세요.

import torch
from transformers import AutoModelForCausalLM

model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-3.2-1B", attn_implementation="kernels-community/flash-attn2"
)

SDPA 컨텍스트 매니저

PyTorch의 scaled dot product attention(SDPA)은 CUDA 백엔드에서 가장 빠른 attention 함수를 자동으로 선택해요. 다른 백엔드에서는 기본적으로 PyTorch C++ 구현을 사용해요.

torch.nn.attention.sdpa_kernel 컨텍스트 매니저로 SDPA가 특정 구현을 쓰도록 강제할 수 있어요.

import torch
from torch.nn.attention import SDPBackend, sdpa_kernel
from transformers import AutoModelForCausalLM

model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-3.2-1B", attn_implementation="sdpa"
)

with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
    outputs = model.generate(**inputs)

백본별 attention (Backbone-specific attention)

멀티모달 모델은 양식(modality)마다 서로 다른 백본을 사용해요. 각 백본에 특정 attention 함수를 할당해서 성능을 최적화할 수 있어요. 예를 들어 일부 비전 백본은 FlashAttention이 지원하지 않는 fp32에서 더 좋은 성능을 보이기도 해요.

dict로 비전 백본을 서로 다른 attention 함수에 매핑하고, 텍스트 백본은 계속 FlashAttention을 사용하게 할 수 있어요. attention 구현의 키는 서브 config 이름과 일치해야 해요.

from transformers import AutoModelForImageTextToText

attention_implementation_per_backbone = {"vision_config": "sdpa", "text_config": "flash_attention_2"}

for key in attention_implementation_per_backbone:
    assert key in model.config.sub_configs, f"Invalid key in `attention_implementation`"

model = AutoModelForImageTextToText.from_pretrained(
    "facebook/chameleon-7b", attn_implementation=attention_implementation_per_backbone
)

특정 백본을 dict에서 빼면 기본 attention 함수(SDPA)를 사용해요.

model = AutoModelForImageTextToText.from_pretrained(
    "facebook/chameleon-7b", attn_implementation={"text_config": "flash_attention_2"}
)

단일 문자열로 모든 백본에 같은 attention 함수를 설정할 수도 있어요.

model = AutoModelForImageTextToText.from_pretrained(
    "facebook/chameleon-7b", attn_implementation="eager"
)

빈 키로 attention 함수를 전역 설정할 수도 있어요.

model = AutoModelForImageTextToText.from_pretrained(
    "facebook/chameleon-7b", attn_implementation={"": "eager"}
)

새 attention 함수 만들기 (Create a new attention function)

AttentionInterface.register()로 attention 레지스트리에 추가해서 attention 함수를 커스터마이즈하거나 새로 만들 수 있어요. 모델은 attn_implementation 인자를 통해 이 함수들을 사용해요.

[!WARNING] 커스텀 attention 함수를 등록할 때는 이에 맞는 attention mask 함수도 함께 등록해야 해요. 커스텀 attn_implementation 이름이 AttentionMaskInterface에 등록되어 있지 않으면, Transformers는 mask 생성을 건너뛰고 attention_mask=None을 attention 레이어에 전달해요. 그러면 attention 함수가 인과(causal), padding, packing, sliding-window 제약을 스스로 처리해야 합니다. 그렇지 않으면 그 제약들이 조용히 사라질 수 있어요.

이 예시는 각 레이어마다 문장을 출력하도록 attention 함수를 커스터마이즈해요. masking_utils.sdpa_mask를 attention mask 함수로 등록해서 원래 구현의 마스크를 유지해요.

import torch
from transformers import AutoModelForCausalLM, AttentionInterface, AttentionMaskInterface
from transformers.integrations.sdpa_attention import sdpa_attention_forward
from transformers.masking_utils import sdpa_mask

def my_new_sdpa(*args, **kwargs):
    print("I just entered the attention computation")
    return sdpa_attention_forward(*args, **kwargs)

AttentionInterface.register("my_new_sdpa", my_new_sdpa)
AttentionMaskInterface.register("my_new_sdpa", sdpa_mask)  # must have the same name as the registered attention function

model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.2-1B", attn_implementation="my_new_sdpa")
model(torch.ones(1, 5, dtype=int))

attention 함수에 새 인자를 추가할 수도 있어요. AttentionInterface를 지원하는 모델은 kwargs를 attention 레이어와 attention 함수로 전파해요. 모델의 forward 함수에서 인자를 kwargs로 전달하세요. 커스텀 attention 함수는 이 시그니처와 반환 형식을 따라야 해요.

import torch
from transformers import AutoModelForCausalLM, AttentionInterface, AttentionMaskInterface
from transformers.integrations.sdpa_attention import sdpa_attention_forward
from transformers.masking_utils import sdpa_mask

def custom_attention(
    module: torch.nn.Module,  # required arg
    query: torch.Tensor,  # required arg
    key: torch.Tensor,  # required arg
    value: torch.Tensor,  # required arg
    attention_mask: Optional[torch.Tensor],  # required arg
    a_new_kwargs = None,  # You can now add as many kwargs as you need
    another_new_kwargs = None,  # You can now add as many kwargs as you need
    **kwargs,  # You need to accept **kwargs as models will pass other args
) -> tuple[torch.Tensor, Optional[torch.Tensor]]
    ...  # do your magic!
    return attn_output, attn_weights  # attn_weights are optional here

AttentionInterface.register("custom", custom_attention)
AttentionMaskInterface.register("custom", sdpa_mask)  # to leave the existing mask untouched

model = AutoModelForCausalLM.from_pretrained(model_id, attn_implementation="custom")
model(torch.ones(1, 5, dtype=int), a_new_kwargs=..., another_new_kwargs=...)

모델이 attention 함수에 어떤 인자와 kwargs를 보내는지 확인하려면 모델의 modeling code를 확인해 주세요.

AttentionMaskInterface

AttentionMaskInterface는 create_*_mask 함수들이 mask를 활성 attention 백엔드가 기대하는 형식으로 변환할 때 참조하는 레지스트리예요. FlexAttention은 BlockMask가 필요하고, SDPA는 4D tensor가 필요하며, FlashAttention은 기본 2D padding mask가 필요해요. AttentionMaskInterface.register()로 커스텀 백엔드를 등록하거나 기존 백엔드의 포맷터를 오버라이드할 수 있어요.

import torch
from transformers import AttentionMaskInterface
from transformers.masking_utils import sdpa_mask

def my_new_sdpa_mask(*args, **kwargs):
    print("I just entered the attention mask computation")
    return sdpa_mask(*args, **kwargs)

AttentionMaskInterface.register("my_new_sdpa_mask", my_new_sdpa_mask)

활성 attn_implementation에 등록된 포맷터가 없으면 mask 생성을 건너뛰고 attention_mask=None이 attention 레이어로 전달돼요.

등록된 함수는 이 시그니처와 일치해야 해요.

def custom_attention_mask(
    batch_size: int,  # required arg
    q_length: int,  # required arg
    kv_length: int,  # required arg
    q_offset: int = 0,  # required arg
    kv_offset: int = 0,  # required arg
    mask_function: Callable = causal_mask_function,  # required arg
    attention_mask: Optional[torch.Tensor] = None,  # required arg
    **kwargs,  # a few additional args may be passed as kwargs, especially the model's config is always passed
) -> Optional[torch.Tensor]:

mask_function 인자는 PyTorch의 mask_mod 함수를 모방하는 Callable이에요. 4개의 인덱스 (batch_idx, head_idx, q_idx, kv_idx)를 받아 그 위치가 attention 계산에 기여하는지 나타내는 boolean을 반환해요. 이는 Build an attention mask의 or_mask_function과 and_mask_function이 사용하는 것과 같은 기본 형태예요.

[!TIP] mask_function이 mask를 만들지 못하면 torch.export용으로 이 workaround를 사용해 보세요.

Attention mask 만들기 (Build an attention mask)

transformers.masking_utils의 create_*_mask 함수들로 attention mask를 만들어요. 각 함수는 모델 config에서 활성 attention 백엔드를 읽고, AttentionMaskInterface에서 백엔드의 mask 포맷터를 찾아, 그 백엔드가 기대하는 형식을 반환해요. mask를 직접 invert, expand, cast 할 필요가 없어요.

attention 패턴에 맞는 함수를 골라요.

함수 용도
create_causal_mask 각 토큰이 자기 자신과 이전 토큰에 attend 하는 decoder-only 모델
create_bidirectional_mask encoder 모델, 또는 decoder에서 encoder 상태로의 cross-attention
create_sliding_window_causal_mask sliding-window attention 패턴을 가진 decoder 모델
create_chunked_causal_mask 시퀀스를 고정 크기 블록으로 청킹하는 decoder 모델
create_bidirectional_sliding_window_mask sliding-window attention 패턴을 가진 encoder 모델

[!WARNING] 레거시 callable mask 헬퍼들 — get_extended_attention_mask, create_extended_attention_mask_for_decoder, invert_attention_mask — 은 deprecated 경고를 내보내며 향후 릴리스에서 제거될 예정이에요. 대신 create_*_mask 함수를 사용하세요.

decoder forward 패스 안에서 create_causal_mask를 호출해요. config, 입력 임베딩, 사용자 제공 2D attention_mask, 캐시를 전달해요. 함수는 임베딩을 사용해서 배치 크기, 쿼리 길이, dtype, 디바이스를 읽고, 캐시를 사용해서 키 길이를 계산해요.

from transformers.masking_utils import create_causal_mask

attention_mask = create_causal_mask(
    config=self.config,
    inputs_embeds=inputs_embeds,
    attention_mask=attention_mask,
    past_key_values=past_key_values,
)

encoder self-attention에는 create_bidirectional_mask를 호출해요. encoder는 캐시를 하지 않으므로 past_key_values는 빼요.

from transformers.masking_utils import create_bidirectional_mask

attention_mask = create_bidirectional_mask(
    config=self.config,
    inputs_embeds=embedding_output,
    attention_mask=attention_mask,
)

cross-attention에서는 encoder_hidden_states로 encoder 상태를 전달해서 mask가 decoder의 쿼리 길이 대신 encoder의 키·밸류 길이를 사용하게 해요.

encoder_attention_mask = create_bidirectional_mask(
    config=self.config,
    inputs_embeds=embedding_output,
    attention_mask=encoder_attention_mask,
    encoder_hidden_states=encoder_hidden_states,
)

or_mask_function과 and_mask_function 인자로 기본 mask 위에 추가 제약을 얹을 수 있어요. or_mask_function은 추가 위치가 attend 하게 하고, and_mask_function은 기본 패턴을 더 제한해요. 둘 다 AttentionMaskInterface에서 설명한 4-인덱스 mask_function 시그니처를 따라요. (batch_idx, head_idx, q_idx, kv_idx)를 받아 boolean을 반환해요.

[!WARNING] or_mask_function과 and_mask_function은 어떤 attention 패턴이든 표현할 수 있지만, 내장 패턴보다 느리고 ExecuTorch와 호환되지 않아요. 오버헤드는 mask 생성이 forward 패스 시간에서 차지하는 비중이 큰 작은 모델(~200M 파라미터)에서 가장 두드러져요. 표준 create_*_mask 함수로 필요한 것을 표현할 수 없을 때만 사용하세요.

예를 들어 어디서나 True를 반환하는 함수를 인과 mask 위에 덧씌워 완전한 양방향(bidirectional) mask로 만들 수 있어요. 인과 패턴과의 합집합 덕분에 모든 토큰이 다른 모든 토큰에 attend 해요.

mask_kwargs = {
    "config": self.config,
    "inputs_embeds": inputs_embeds,
    "attention_mask": attention_mask,
    "past_key_values": past_key_values,
    "position_ids": position_ids,
    "or_mask_function": lambda *args: torch.tensor(True, dtype=torch.bool),
}

attention_mask = create_causal_mask(**mask_kwargs)

생성 중에는 generate()가 create_masks_for_generate를 통해 mask를 만드는데, 이는 모델 config에 따라 올바른 create_*_mask로 디스패치해요. 모델 클래스에서 이를 오버라이드하면 생성용 커스텀 masking 전략을 끼워 넣을 수 있어요.

커스텀 4D attention mask 전달하기

create_*_mask 함수들이 표현할 수 없는 attention 패턴이 필요할 때는 나만의 4D mask를 전달해요. 4D mask의 shape은 (batch_size, 1, query_length, kv_length)이고, 여기서 1은 모든 attention 헤드에 같은 mask를 브로드캐스트해요. Transformers는 이를 감지해서 그대로 사용하고 create_*_mask를 건너뛰어요.

4D mask는 두 가지 값 규약 중 하나를 사용해요.

dtype attend mask out
boolean True False
float 0.0 -inf

float 규약은 softmax 전에 mask를 attention 점수에 더해요. 점수에 0.0을 더하면 그대로이므로 그 위치가 기여해요. 점수에 -inf를 더하면 softmax 후 0으로 떨어지므로 그 위치는 제외돼요.

[!IMPORTANT] 허용되는 규약은 attention 백엔드에 따라 달라요. sdpa는 boolean 또는 float mask를 받아요. eager는 mask를 점수에 더하므로 float mask만 받아요. flash_attention_2와 flex_attention은 각자만의 형식(2D padding mask와 BlockMask)을 사용하며 원시 4D mask를 받아들이지 않아요.

흔한 실수는 2D padding mask의 1/0 규약을 float 4D mask에 그대로 쓰는 거예요. mask가 점수에 더해지므로 0.0은 위치를 유지하고 1.0은 아주 작은 편향만 더해요.

아래 예시는 틀린 mask와 올바른 mask를 대조해요. 둘 다 같은 1/0 인과 패턴에서 시작해요.

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

model = AutoModelForCausalLM.from_pretrained("TinyLlama/TinyLlama-1.1B-Chat-v1.0", attn_implementation="sdpa")
tokenizer = AutoTokenizer.from_pretrained("TinyLlama/TinyLlama-1.1B-Chat-v1.0")

input_ids = tokenizer("my favorite condiment on a", return_tensors="pt").input_ids
seq_len = input_ids.shape[1]

# 1 attends, 0 masks
causal = torch.tril(torch.ones(seq_len, seq_len))

# wrong: 1.0/0.0 floats are added to the scores, so 0.0 keeps a token and 1.0 barely changes it
wrong_mask = causal[None, None]

# correct: 0.0 attends, -inf masks
correct_mask = torch.where(causal.bool(), 0.0, float("-inf"))[None, None]

틀린 mask는 0.0이 mask하는 값이라 아무것도 제외하지 못해서 모든 위치를 유지해요.

        wrong_mask                          correct_mask
   (1 attends, 0 masks)              (0 attends, -inf masks)

      k0 k1 k2 k3 k4                     k0   k1   k2   k3   k4
   q0  1  0  0  0  0                  q0  0  -inf -inf -inf -inf
   q1  1  1  0  0  0                  q1  0   0   -inf -inf -inf
   q2  1  1  1  0  0                  q2  0   0    0   -inf -inf
   q3  1  1  1  1  0                  q3  0   0    0    0   -inf
   q4  1  1  1  1  1                  q4  0   0    0    0    0

양방향 attention (Bidirectional attention)

Decoder-only 모델은 기본적으로 인과(단방향) attention을 사용하는데, 각 토큰은 자기 자신과 이전 토큰에만 attend 해요. is_causal=False로 설정하면 모든 토큰이 다른 모든 토큰에 attend 하는 양방향 attention으로 전환돼요. 이렇게 하면 예를 들어 embedding을 생성하기 위해 decoder-only 모델을 텍스트 encoder로 사용할 수 있어요.

[!NOTE] 이는 인과(decoder) 모델에서만 동작해요. encoder 모델을 decoder 모델로 바꿔주지는 않아요.

모델 config에서 is_causal=False를 설정하면 모든 forward 패스의 기본값이 양방향 attention이 돼요.

from transformers import AutoModel, AutoConfig

config = AutoConfig.from_pretrained("meta-llama/Llama-3.2-1B")
config.is_causal = False

model = AutoModel.from_pretrained("meta-llama/Llama-3.2-1B", config=config)

# all forward passes now use bidirectional attention
outputs = model(**inputs)

모델 config 대신 forward 호출에서 is_causal을 전달하면 모델을 두 번 로드하지 않고도 인과와 양방향 attention을 오갈 수 있어요. 이 kwarg는 config를 일시적으로 오버라이드하고 호출 후에 복원돼요.

from transformers import AutoModel

model = AutoModel.from_pretrained("meta-llama/Llama-3.2-1B")

# run with bidirectional attention
outputs = model(**inputs, is_causal=False)

# run with default causal attention
outputs = model(**inputs)

더 알아보기 (Learn more)