패딩 없는 훈련
패딩 없는 훈련
패딩 없는 훈련(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과 충돌하는 데이터 의존적 분기를 추가하기 때문에 오류나 경고가 발생하지 않습니다.
다음 단계
- 다른 collator에 대해서는 데이터 collator 가이드를 참조하세요.
- 전체 인자 집합에 대해서는 DataCollatorWithFlattening API 참조를 살펴보세요.
- 벤치마크와 더 깊은 설명은 Flash Attention으로 패킹하여 Hugging Face 훈련 효율성 개선을 읽어 보세요.