모델 중간 출력 추적하기

모델 중간 출력 추적하기

모델의 forward() 메서드는 예전에 config 기본값에서 output_attentions 같은 None 플래그를 수동으로 해석하고, 레이어별 어텐션 가중치와 hidden state를 튜플로 누적하며, return_dict=False일 때 ModelOutput 데이터클래스를 일반 튜플로 변환하는 작업을 했습니다. 두 개의 데코레이터가 그 모든 보일러플레이트를 대체합니다.

출처: 문서

본문

모든 모델의 forward() 메서드는 예전에 config 기본값에서 output_attentions 같은 None 플래그를 수동으로 해석하고, 레이어별 어텐션 가중치와 hidden state를 튜플로 누적하며, return_dict=False일 때 ModelOutput 데이터클래스를 일반 튜플로 변환하는 작업을 했습니다. 두 개의 데코레이터가 그 모든 보일러플레이트를 대체합니다.

  • @capture_outputs는 출력 플래그를 해석하고, 중간 값을 수집하며, return_dict 변환을 처리합니다.
  • @merge_with_config_defaults는 config에서 use_cache를 해석합니다. CLIPModel처럼 캐시하지 않는 모델에서는 생략하세요.

이 데코레이터들은 주로 새 모델을 통합할 때 접하게 됩니다. 단계별 안내는 🤗 Transformers에 모델 추가하기를 참조하세요.

캡처할 서브모듈 선언

베이스 모델의 forward() 메서드에 @capture_outputs를 적용합니다. 이 데코레이터는 다음 작업을 수행하는 forward 훅을 연결합니다.

  • 서브모듈이 관찰되고 있다는 것을 알 필요 없이, forward pass 중 지정된 서브모듈 클래스의 출력을 가로챕니다.
  • 레이어별 어텐션 가중치와 hidden state를 튜플로 수집합니다.
  • 수집된 값을 반환된 ModelOutput 데이터클래스에 주입합니다.
  • return_dict=False일 때 데이터클래스를 일반 튜플로 변환합니다.
  • None일 때 kwargs 또는 self.config에서 output_attentions와 output_hidden_states를 해석합니다.

출력 필드를 서브모듈에 매핑

@capture_outputs는 어떤 서브모듈이 어떤 출력을 생성하는지 알아야 합니다. PreTrainedModel 서브클래스에 클래스 수준 딕셔너리인 _can_record_outputs를 선언합니다. 각 키는 출력 필드 이름("hidden_states", "attentions", "cross_attentions")이고, 각 값은 모듈 클래스, 클래스 이름 문자열, OutputRecorder 인스턴스, 또는 하나의 키 아래 여러 모듈 유형을 기록하기 위한 그 목록입니다.

OutputRecorder로 세밀한 제어

OutputRecorder는 target_class(출력을 수집할 nn.Module 서브클래스)와 선택적 index를 받아 모듈의 출력 튜플에서 어떤 요소를 가져올지 선택합니다. 특정 속성 이름을 가진 모듈에만 훅을 연결하려면 layer_name을 전달하세요. self-attention과 cross-attention처럼 두 레이어가 같은 클래스를 공유할 때 layer_name을 사용합니다.

아래 예시는 서로 다른 출력 제어 수준으로 두 데코레이터를 실제로 사용하는 방법을 보여줍니다. 실제 사례는 LlamaModel을 참조하세요.

from ...processing_utils import Unpack
from ...utils import TransformersKwargs
from ...utils.generic import merge_with_config_defaults
from ...utils.output_capturing import capture_outputs, OutputRecorder

class MyPreTrainedModel(PreTrainedModel):
    _can_record_outputs = {
        # hidden_states 캡처: 각 MyDecoderBlock forward 후 훅 실행,
        # 첫 출력을 가져옴 (기본적으로 index 0).
        "hidden_states": MyDecoderBlock,

        # self-attention 가중치 캡처: 각 MyAttention forward 후 훅 실행,
        # 두 번째 출력을 가져옴 (기본적으로 index=1).
        "attentions": MyAttention,

        # cross-attention 가중치 캡처: 같은 클래스, 다른 서브모듈.
        # layer_name은 블록 안의 속성 `self.crossattention`을 대상으로 함.
        # 요청대로 두 번째 출력을 캡처 (index=1)
        "cross_attentions": OutputRecorder(
            MyAttention, layer_name="crossattention", index=1
        ),
    }

# 이제 베이스 모델의 forward에서 데코레이터와 `Unpack` `kwargs`가 필요함
class MyModel(MyPreTrainedModel):

    @merge_with_config_defaults # ← use_cache를 해석함
    @capture_outputs            # ← 출력 수집 + return_dict 처리
    def forward(
        self,
        input_ids: torch.LongTensor | None = None,
        past_key_values: Cache = None,
        **kwargs: Unpack[TransformersKwargs],
    ) -> BaseModelOutputWithPast:

        # 수동 수집은 필요 없음. 레이어를 평소처럼 실행하기만 하면 됨.
        hidden_states = self.embed_tokens(input_ids)
        for layer in self.layers:
            hidden_states = layer(hidden_states, **kwargs)

        # 기본 출력 반환. 데코레이터가
        # hidden_states/attentions/cross_attentions를 자동으로 채워줌.
        return BaseModelOutputWithPast(
            last_hidden_state=hidden_states,
            past_key_values=past_key_values
        )

레이어 클래스 패치

출력 추적은 _can_record_outputs가 모델의 레이어가 인스턴스화하는 정확한 클래스를 가리키는 것에 의존합니다. 레이어 구현을 커스텀 어텐션 커널, 양자화된 전문가 레이어, 또는 아키텍처 변형으로 교체한다면 그 포인터들은 동기화된 상태를 유지해야 합니다. 패치 API는 _can_record_outputs를 일관되게 유지하기 위한 깔끔한 전역 레지스트리를 제공합니다.

register_patch_mapping은 원본 클래스 이름을 대체 nn.Module 서브클래스에 매핑합니다. 키는 정확한 클래스 이름 또는 정규식 패턴일 수 있습니다. 정확한 일치가 우선합니다. 패턴은 re.search()로 테스트되므로, 앵커가 없는 패턴은 클래스 이름 어디에서나 일치합니다. overwrite=True를 전달하지 않는 한 같은 키를 두 번 등록하면 ValueError가 발생합니다.

unregister_patch_mapping으로 항목을 제거합니다.

from transformers.monkey_patching import register_patch_mapping, unregister_patch_mapping

# 정확한 이름 – Qwen2MoeExperts만 대체함
register_patch_mapping({"Qwen2MoeExperts": SequentialExperts})

# 정규식 – 이름이 "Attention"으로 끝나는 모든 클래스를 대체함
register_patch_mapping({".*Attention$": FusedAttention})

# 앵커 버전 – Llama2Attention, Llama3Attention, …만 일치함
register_patch_mapping({"^Llama\\d+Attention$": CustomLlamaAttention})

# 같은 방식으로, 등록된 이름을 전달해 레지스트리에서 커스텀 키를 제거할 수 있음
unregister_patch_mapping(["Qwen2MoeExperts", ".*Attention$"])

매핑이 등록되면 patch_output_recorders가 모든 서브모듈을 순회하며 각 OutputRecorder.target_class를 등록된 대체 클래스로 업데이트합니다.

[!TIP] from_pretrained() 메서드는 patch_output_recorders를 자동으로 호출합니다. 모델을 직접 구성할 때만 직접 호출하면 됩니다.

from transformers.monkey_patching import patch_output_recorders
# from_pretrained 밖에서 수동으로 구성함
model = Qwen2MoeModel(config)

# 이게 없으면 _can_record_outputs는 여전히 원본 Qwen2MoeExperts 클래스를 가리켜
# hooks가 CustomExperts 인스턴스에서는 절대 실행되지 않음.
patch_output_recorders(model)

더 알아보기 (Learn more)