본문으로 건너뛰기

[LlamaFactory] LLaMA Factory v1: 멀티모달 및 메모리 효율적인 SFT를 위한 Ulysses CP와 Chunk Loss 지원

PR 링크: hiyouga/LlamaFactory#10762 상태: Merged | 변경: +962 / -188

들어가며

최근 대규모 언어 모델(LLM)의 발전은 텍스트뿐만 아니라 이미지, 오디오 등 다양한 양식의 데이터를 이해하고 생성하는 멀티모달(Multimodal) 능력으로 확장되고 있습니다. 이러한 멀티모달 모델을 효율적으로 학습시키기 위해서는 기존의 텍스트 기반 모델 학습 방식과는 다른 접근 방식이 필요합니다. 특히, 모델의 크기가 커지고 처리해야 할 시퀀스 길이가 길어질수록 메모리 사용량과 계산 복잡성은 기하급수적으로 증가합니다.

Hiyouga의 LLaMA Factory 레포지토리에서 이번 PR은 이러한 문제를 해결하기 위해 두 가지 주요 기능을 도입했습니다. 첫째, Ulysses Context Parallelism(CP)을 멀티모달 모델까지 확장하여 시퀀스 병렬 처리의 효율성을 높였습니다. 둘째, 메모리 사용량을 획기적으로 줄일 수 있는 Chunk Loss를 SFT(Supervised Fine-Tuning) 작업에 적용했습니다. 이 글에서는 이 PR의 코드 변경 사항을 상세히 분석하고, 각 개선 사항이 왜 좋은 최적화인지, 그리고 실제 성능 향상에 어떻게 기여하는지 살펴보겠습니다.

코드 분석

이번 PR은 주로 설정 파일(yaml)과 핵심 학습 로직(base_trainer.py, chunk_loss.py)에 걸쳐 변경 사항을 포함하고 있습니다.

1. 설정 파일 (.yaml)

새로운 학습 시나리오를 지원하기 위한 두 개의 설정 파일이 추가되었습니다.

  • examples/v1/train_full/train_full_chunk_loss.yaml: Chunk Loss를 사용한 SFT 학습 설정을 정의합니다. 특히 chunk_loss_size 파라미터가 추가되어, 손실 계산 시 시퀀스를 얼마나 작은 단위로 나눌지 결정합니다. 이 예시에서는 chunk_loss_size: 256으로 설정되어 있습니다.

    # Before (기존 방식, chunk_loss_size 미설정)
    # ...
    cutoff_len: 2048
    # ...
    
    # After (Chunk Loss 추가)
    model: Qwen/Qwen3-0.6B
    trust_remote_code: true
    model_class: llm
    # ...
    cutoff_len: 2048
    # Maximum flattened token rows per logits/CE chunk; this is not the sequence length.
    chunk_loss_size: 256
    # ...
    
  • examples/v1/train_full/train_full_multimodal_ulysses_cp.yaml: 멀티모달 모델에 Ulysses Context Parallelism을 적용하기 위한 설정을 정의합니다. cp_mode: ulyssescp_size: 2가 설정되어, 2개의 CP 그룹을 사용하여 시퀀스 병렬 처리를 수행함을 나타냅니다.

    # Before (기존 방식, 멀티모달 CP 미지원)
    # ...
    
    # After (멀티모달 Ulysses CP 추가)
    model: Qwen/Qwen3.5-0.8B
    trust_remote_code: true
    model_class: llm
    # ...
    cp_mode: ulysses
    cp_size: 2
    # ...
    

2. TrainingArguments 클래스 업데이트 (src/llamafactory/v1/config/training_args.py)

Chunk Loss 기능을 활성화하기 위한 chunk_loss_size 인자가 TrainingArguments 클래스에 추가되었습니다. 이 인자는 None일 경우 Chunk Loss가 비활성화되며, 양수 값으로 설정될 경우 해당 크기의 청크로 손실을 계산합니다. 또한, chunk_loss_size가 0 이하일 경우 ValueError를 발생시켜 잘못된 설정을 방지합니다.

--- a/src/llamafactory/v1/config/training_args.py
+++ b/src/llamafactory/v1/config/training_args.py
@@ -144,6 +144,10 @@
         default=1,
         metadata={"help": "Log metrics every N optimizer steps."}),
     )
+    chunk_loss_size: int | None = field(
+        default=None,
+        metadata={"help": "Maximum flattened token rows per Chunk Loss chunk. None disables Chunk Loss."}),
+    )
     pref_loss: Literal["sigmoid", "orpo", "simpo"] = field(
         default="sigmoid",
         metadata={"help": "The type of DPO loss to use."}),
@@ -173,6 +177,8 @@
         self.dist_config = get_plugin_config(self.dist_config)
         self.optim_config = get_plugin_config(self.optim_config)
         self.lr_scheduler_config = get_plugin_config(self.lr_scheduler_config)
+        if self.chunk_loss_size is not None and self.chunk_loss_size <= 0:
+            raise ValueError("`chunk_loss_size` must be positive.")
         try:
             from ..plugins.model_plugins.deepspeed_utils import register_deepspeed_dist_config
 

3. BaseTrainer 수정 (src/llamafactory/v1/core/base_trainer.py)

이 PR은 Context Parallelism(CP) 관련 손실 계산 로직을 BaseTrainercompute_loss 메서드에서 제거하고, 각 Trainer 구현체가 자체적으로 처리하도록 변경했습니다. 이는 CP 환경에서의 손실 계산 방식이 Trainer마다 다를 수 있음을 반영하며, 특히 Ulysses CP와 같은 고급 병렬 처리 기법과의 통합을 유연하게 만듭니다.

리뷰어 khazic의 지적에 따라, compute_loss 메서드의 docstring이 업데이트되어 CP 환경에서의 손실 계산 및 집계에 대한 계약(contract)을 명확히 했습니다.

--- a/src/llamafactory/v1/core/base_trainer.py
+++ b/src/llamafactory/v1/core/base_trainer.py
@@ -239,7 +239,12 @@
 
     @abstractmethod
     def compute_loss(self, batch: BatchInput) -> Tensor:
-        """Compute the scalar loss."""
+        """Compute the scalar loss.
+
+        Subclasses must handle sequence-parallel layout and loss aggregation when
+        `self.cp_size > 1`, or reject context parallelism during initialization.
+        The shared training loop does not dispatch sequence-parallel loss.
+        """
         ...
 
     def fit(self) -> None:
@@ -265,14 +270,7 @@
                 step_valid_tokens = DistributedInterface().all_reduce(step_valid_tokens, op=ReduceOp.SUM)
                 num_micro = len(micro_batches)
                 for i, micro_batch in enumerate(micro_batches):
-                    if self.args.cp_size > 1:
-                        from ..plugins.model_plugins.parallelization.sequence_parallel import \
-                            SequenceParallelLossPlugin,
-
-                        loss = SequenceParallelLossPlugin("sequence_parallel_loss")(self.model, micro_batch)
-                    else:
-                        loss = self.compute_loss(micro_batch)
+                    loss = self.compute_loss(micro_batch)
                     mini_step_valid_tokens = compute_valid_tokens([micro_batch])
                     # fsdp uses mean reduction so we need to scale the loss by dp_size
                     loss = loss * mini_step_valid_tokens * self.dp_size / (step_valid_tokens + 1e-6)

4. Chunk Loss 구현 (src/llamafactory/v1/plugins/model_plugins/chunk_loss.py)

이 PR의 핵심 중 하나는 chunk_loss.py 파일에 새로 구현된 Chunk Loss 로직입니다. 이 구현은 torch.autograd.Function을 사용하여 커스텀 역전파를 정의하고, 모델의 최종 선형 레이어(output head)에서 직접 교차 엔트로피(Cross-Entropy) 손실을 계산합니다.

  • _ChunkedLinearCrossEntropy: 이 클래스는 실제 청크 단위 손실 계산 및 기울기 계산을 담당합니다. forward 메서드에서는 전체 시퀀스의 logits을 생성하는 대신, chunk_size만큼의 작은 단위로 나누어 torch.nn.functional.cross_entropy를 적용합니다. backward 메서드는 이 과정에서 발생하는 기울기를 효율적으로 집계합니다.

    리뷰어 khazic의 지적에 따라, grad_weightgrad_bias가 FP32로 누적된 후 최종적으로 원래의 dtype으로 캐스팅되도록 수정되었습니다. 이는 BF16과 같은 저정밀도 부동소수점 연산에서 발생할 수 있는 정밀도 손실을 최소화하여, 기존의 Eager Loss 방식과 유사한 정확도를 보장합니다.

    # Before (정밀도 문제 가능성)
    # grad_weight = torch.zeros_like(head_weight) if needs_weight_grad else None
    # ...
    # if grad_weight is not None:
    #     grad_weight.add_(chunk_grads[grad_index])
    
    # After (FP32 누적 및 최종 캐스팅)
    grad_weight = torch.zeros_like(head_weight, dtype=torch.float32) if needs_weight_grad else None
    # ...
    ctx.save_for_backward(
        # ...
        grad_weight.to(head_weight.dtype) if grad_weight is not None else None,
        # ...
    )
    
  • ChunkLoss 클래스: 이 클래스는 torch.nn.Moduleforward 메서드를 가로채는 방식으로 동작합니다. __init__ 메서드에서 모델의 output_head (일반적으로 nn.Linear 타입)의 forward 메서드를 커스텀 _head_forward로 교체합니다. _head_forwardchunk_loss_size에 따라 입력을 분할하고 _ChunkedLinearCrossEntropy를 호출하여 손실을 계산합니다. 또한, IGNORE_INDEX (-100) 하드코딩 대신 v1/utils/constants.IGNORE_INDEX를 사용하도록 수정되었습니다 (리뷰어 khazic의 피드백 반영).

    # Before (하드코딩된 IGNORE_INDEX)
    # token_loss = F.cross_entropy(
    #     logits,
    #     labels_flat[start:end],
    #     reduction="none",
    #     ignore_index=-100, # <-- 하드코딩
    # )
    
    # After (상수 사용)
    token_loss = F.cross_entropy(
        logits,
        labels_flat[start:end],
        reduction="none",
        ignore_index=IGNORE_INDEX, # <-- 상수 사용
    )
    
  • 멀티모달 CP 관련 로직: PR 설명에 따르면, 멀티모달 CP는 미디어 입력은 복제된 상태로 유지하면서 인코더, 멀티모달 퓨전, mRoPE 준비를 수행합니다. 이후 언어 모델 경계에서 퓨즈된 언어 시퀀스를 샤딩합니다. 이는 멀티모달 데이터의 특성을 고려하여 효율적인 시퀀스 병렬 처리를 가능하게 합니다. 구체적인 코드 변경은 base_trainer.pycompute_loss 메서드 수정과 관련 플러그인에서 이루어졌을 것으로 추정됩니다.

왜 이게 좋은가?

1. 메모리 효율성 향상 (Chunk Loss)

기존의 Eager Loss 방식은 전체 시퀀스 길이에 대한 logits을 메모리에 올려 손실을 계산합니다. 시퀀스 길이가 길어질수록 이 logits 텐서의 크기는 [batch_size, sequence_length, vocab_size]로 기하급수적으로 커져 메모리 병목 현상을 일으킵니다. Chunk Loss는 이 logits 텐서를 chunk_loss_size 단위의 작은 청크로 나누어 처리함으로써 메모리 사용량을 크게 줄입니다.

PR의 벤치마크 결과는 이를 명확히 보여줍니다:

  • Qwen3.5-35B-A3B 모델, 4K 시퀀스 길이, 16개 카드 기준
    • Baseline (Eager Loss): Peak Torch allocated: 22.19 GiB/device
    • Chunk Loss 512: Peak Torch allocated: 13.53 GiB/device (-39.01%)
    • Chunk Loss 2048: Peak Torch allocated: 17.49 GiB/device (-21.16%)

특히 chunk_loss_size를 512로 설정했을 때, 피크 메모리 사용량이 약 39% 감소했습니다. 이는 더 큰 모델이나 더 긴 시퀀스를 제한된 하드웨어 자원으로 학습할 수 있게 해줍니다.

2. 멀티모달 학습 효율성 증대 (Ulysses CP)

멀티모달 모델은 텍스트 외에 이미지 임베딩 등 추가적인 데이터를 처리해야 합니다. Ulysses CP를 멀티모달 모델까지 확장함으로써, 이러한 복잡한 데이터 흐름 속에서도 시퀀스 병렬 처리의 이점을 누릴 수 있게 되었습니다. PR 설명의 표는 Ulysses CP 적용 시 손실 값의 큰 변화 없이 병렬 처리 효율을 높일 수 있음을 시사합니다.

  • Qwen3.5-0.8B 모델, DP4 x CP1 vs DP4 x CP2
    • Baseline (DP4 x CP1): Mean loss: 4.2374
    • Ulysses CP (DP4 x CP2): Mean loss: 4.2444

손실 값의 차이가 매우 작다는 것은, 멀티모달 CP가 모델의 학습 능력에 부정적인 영향을 미치지 않으면서 병렬 처리로 인한 학습 속도 향상 또는 더 큰 배치 크기 사용 가능성을 열어준다는 것을 의미합니다.

3. 정확도 유지 및 코드 품질 향상

  • 정확도: Chunk Loss 구현 시 FP32로 기울기를 누적하고 최종적으로 캐스팅하는 방식은 기존 Eager Loss 방식과의 수치적 차이를 최소화하여 학습 정확도를 유지합니다. 벤치마크의 손실 비교 그래프에서도 Chunk Loss 설정(512, 2048)이 Baseline과 매우 유사한 손실 곡선을 보여줍니다.
  • 코드 품질: 리뷰 피드백을 적극적으로 반영하여 IGNORE_INDEX 상수 사용, 기울기 계산 정밀도 향상 등 코드의 견고성과 유지보수성을 높였습니다. 또한, BaseTrainercompute_loss 메서드 docstring 업데이트는 API 사용 계약을 명확히 하여 향후 개발의 편의성을 증진시킵니다.

4. 일반적인 교훈

  • 메모리 병목 해결: 시퀀스 길이가 길어질 때 발생하는 메모리 병목은 logits 계산 단계에서 발생하기 쉽습니다. 이를 해결하기 위해 계산을 청크 단위로 나누는 것은 효과적인 전략입니다. 다만, 청크 크기 설정은 메모리 절감과 처리 속도 간의 트레이드오프를 고려해야 합니다.
  • 멀티모달 학습의 복잡성: 멀티모달 데이터는 텍스트 외의 데이터를 포함하므로, 기존의 병렬 처리 전략을 그대로 적용하기 어렵습니다. 각 데이터 양식의 특성을 고려한 맞춤형 병렬 처리 전략이 필요합니다.
  • 리뷰의 중요성: 코드 리뷰는 단순히 버그를 찾는 것을 넘어, 성능, 정확도, 코드 품질 등 다양한 측면에서 개선점을 발견하고 적용하는 중요한 과정입니다. 특히 autograd.Function과 같이 복잡한 로직에서는 정밀도 문제나 하드코딩된 값 사용과 같은 잠재적 이슈가 발생하기 쉬우므로 더욱 중요합니다.

결론

이번 LLaMA Factory v1의 PR은 Ulysses Context Parallelism을 멀티모달 모델로 확장하고, SFT를 위한 메모리 효율적인 Chunk Loss를 도입함으로써 대규모 멀티모달 모델 학습의 효율성을 크게 향상시켰습니다. 벤치마크 결과는 메모리 사용량 감소와 정확도 유지라는 두 마리 토끼를 잡았음을 보여줍니다. 또한, 코드 리뷰를 통해 제기된 문제점들을 해결하며 코드의 품질과 견고성을 더욱 높였습니다. 이러한 개선은 LLM 연구자들이 더 크고 복잡한 모델을 더 효율적으로 학습할 수 있도록 지원하는 중요한 발걸음입니다.

참고 자료

⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.

댓글

관련 포스트

PR Analysis 의 다른글