패딩 없는 훈련

패딩 없는 훈련

패딩 없는 훈련(padding-free training, 패킹이라고도 함)은 각 샘플을 고정 길이로 패딩하는 대신 여러 샘플을 하나의 시퀀스로 결합합니다. 모델은 각 샘플이 끝나는 위치를 알아야 (linear) 어텐션이 샘플 간에 토큰을 섞지 않습니다.

출처: 문서

본문

패딩 없는 훈련(padding-free training, 패킹이라고도 함)은 각 샘플을 고정 길이로 패딩하는 대신 여러 샘플을 하나의 시퀀스로 결합합니다. 모델은 각 샘플이 끝나는 위치를 알아야 (linear) 어텐션이 샘플 간에 토큰을 섞지 않습니다.

이 경계를 제공하는 방법은 두 가지가 있습니다.

  • 데이터 collator로 미리 준비합니다.
  • 런타임에 position_ids에서 추론합니다.

권장하는 방법은 데이터 collator입니다. 이 가이드는 그 이유를 설명하고 position_ids 경로의 주의점을 다룹니다.

[!WARNING] position_ids에서 경계를 추론하는 것은 선호되는 방식이 아니며 표준 어텐션 모델에서만 동작합니다. Qwen3-Next, Qwen3.5(Gated DeltaNet) 같은 linear-attention 모델과 컨볼루션 기반 모델은 position_ids 경계를 무시하므로 데이터 collator가 필요합니다. Linear attention 및 컨볼루션 모델을 참조하세요.

데이터 collator로 경계 준비

경계 kwargs를 미리 준비하면 위의 문제를 제거하고 컴파일 여부와 무관하게 동일하게 동작합니다.

DataCollatorWithFlattening을 사용해 각 배치를 평탄화하고 경계 정보를 반환합니다. return_flash_attn_kwargs=True를 설정하면 collator가 경계를 런타임에 position_ids로 추론하도록 두지 않고 미리 계산합니다. 이를 Trainer에 전달하고 attention_mask를 추가하지 마세요. 평탄화된 배치가 이미 경계를 인코딩하고 있고, 마스크는 패킹된 배치와 충돌하기 때문입니다.

[!TIP] 패딩 없는 훈련은 표준 어텐션 모델에서 FlashAttention 구현에 의존합니다. 평탄화된 배치에 필요한 가변 길이 경로를 노출하는 것은 FlashAttention 커널뿐이기 때문입니다.

kernels 라이브러리를 설치하세요. 이 라이브러리는 로컬 빌드 없이 사전 빌드된 FlashAttention 커널을 가져옵니다. flash-attn(https://github.com/Dao-AILab/flash-attention)이 로컬에 설치되지 않은 경우에도 대체로 동작합니다. attn_implementation="kernels-community/flash-attn2"로 모델을 로드하세요.

import torch
from datasets import load_dataset
from transformers import AutoModelForCausalLM, AutoTokenizer, DataCollatorWithFlattening, Trainer, TrainingArguments

model_id = "meta-llama/Llama-3.2-1B"
tokenizer = AutoTokenizer.from_pretrained(model_id)
model = AutoModelForCausalLM.from_pretrained(
    model_id,
    attn_implementation="flash_attention_2",
    device_map="auto",
)

dataset = load_dataset("Salesforce/wikitext", "wikitext-2-raw-v1", split="train")
dataset = dataset.map(
    lambda example: tokenizer(example["text"], truncation=True, max_length=512),
    remove_columns=dataset.column_names,
)

# return_flash_attn_kwargs=True는 시퀀스 경계를 미리 계산함
data_collator = DataCollatorWithFlattening(return_flash_attn_kwargs=True)

trainer = Trainer(
    model=model,
    args=TrainingArguments(output_dir="padding-free-llama"),
    train_dataset=dataset,
    data_collator=data_collator,
)
trainer.train()

position_ids에서 경계 추론

FlashAttention은 position_ids만으로 패딩 없는 배치를 감지할 수 있으며, TRL 같은 하위 프레임워크가 그것에 의존하기 때문에 이 방식은 하위 호환성을 위해 유지됩니다.

position_ids에 의존하는 것에는 두 가지 문제가 있습니다.

  • position_ids에서 패킹된 시퀀스를 감지하는 것은 동적이고 데이터에 의존적인 검사입니다. 컴파일 없이도 동작하지만, torch.compile 아래에서는 그래프 중단(graph break)을 유발합니다. 실제 배치 크기는 대개 더 크기 때문에 이 검사가 실행되는 빈도를 제한하려고 현재 batch_size == 1로 제한되어 있습니다.
  • 컴파일된 FlashAttention은 일부 kwargs를 일반 Python int로 강제합니다. 런타임에 position_ids에서 이를 추론하면 장치-호스트 동기화가 강제되고, 더 오래된 PyTorch 버전에서는 텐서-정수 변환으로 인한 추가 그래프 중단이 발생합니다.

Linear attention 및 컨볼루션 모델

Gated DeltaNet(GDN), 기타 linear-attention 레이어, 인과 컨볼루션은 설계상 position_ids 전용 경로가 없습니다. 이러한 모델에서는 collator로 데이터를 준비하는 것만 지원되는 옵션입니다.

[!WARNING] GDN, linear-attention, 또는 인과 컨볼루션 모델에서는 position_ids만으로 의존하지 마세요. seq_idx를 포함한 경계 kwargs를 데이터 collator로 준비하세요.

이러한 모델에서는 return_flash_attn_kwargs=True와 return_seq_idx=True 둘 다 설정합니다.

from transformers import DataCollatorWithFlattening

data_collator = DataCollatorWithFlattening(
    return_flash_attn_kwargs=True,
    return_seq_idx=True,
)

정확한 커널 패키지는 모델의 원본 구현에 따라 다릅니다. Qwen3-Next, Qwen3.5 같은 Gated DeltaNet 모델은 flash-linear-attention을 사용하고, Bamba 같은 Mamba 기반 모델은 mamba-ssm을 사용합니다. 둘 다 컨볼루션에 causal-conv1d에 의존합니다. 올바른 커널이 없으면 모델은 경계 kwargs를 무시하고 샘플 간에 토큰을 섞는 참조 구현으로 대체됩니다.

[!TIP] 이러한 커널 중 상당수는 kernels 라이브러리를 통해서도 사용할 수 있으며, 호환 가능한 빌드를 가져와 줍니다. flash-linear-attention은 일반적으로 여전히 직접 설치가 필요합니다.

경계 kwargs가 없을 때 커널은 조용히 전체 배치를 하나의 시퀀스로 취급합니다. 런타임 검사가 torch.compile과 충돌하는 데이터 의존적 분기를 추가하기 때문에 오류나 경고가 발생하지 않습니다.

다음 단계

더 알아보기 (Learn more)