하루 종일 컨텍스트가 10분 컨텍스트보다 낫습니까? 플래시 어텐션과 시퀀스 길이 질문
이 기사에서 답을 찾기 위해 존재하는 질문은 다음과 같습니다. 변환기가 10분 창 대신 1초 해상도로 전체 거래일에 참석할 수 있다면 더 나은 예측을 할 수 있을까요?
최근까지는 물어볼 수도 없었습니다. 표준적인 관심 요구 따라서 float16의 12개 헤드에서 23,400단계의 하루에는 점수 매트릭스에만 대략 12.9GB가 필요합니다. 이는 모델 매개변수보다 많고 대부분의 카드가 제공하는 것보다 더 많은 것입니다. 질문은 누군가가 테스트하기 전에 산술적으로 종료되었습니다.
Flash Attention(Dao et al., 2022)이 이를 엽니다. 주의를 근사화하는 것이 아니라 정확히 동일한 결과를 계산합니다. 하지만 IO 인식으로 계산을 재구성하여 GPU 메모리 수준 간의 트래픽을 최소화합니다. 이것이 바로 여기서 정말 흥미로운 내용이며, 이 기사의 대부분은 타일링, 온라인 소프트맥스 재발, IO 바운드 및 역방향 통과 재계산.
그러나 메커니즘은 주장이 아니라 조력자입니다. "Longer context is better"는 시장에 대한 경험적 진술이며 이 블로그의 입장은 Temporal Fusion Transformers, 이는 금융 시리즈 과적합 및 단기 검토 반복 모델의 바닐라 변환기가 높은 빈도에서 경쟁력을 유지한다는 사실을 발견했습니다. 따라서 이 기사는 메커니즘이 아닌 측정에 대해 마무리됩니다.
주의가 기억에 묶여 있는 이유
주의 계산 — 거래 맥락에서 기본 요소 자체는 Multi-Horizon 포트폴리오 예측을 위한 임시 융합 변환기. 전체 문제는 그 중 한 줄입니다: 중간 점수 행렬 ~이다 , 메모리에 기록되고, 소프트맥스를 위해 다시 읽고, 다시 쓰고, 최종 matmul을 위해 다시 읽혀지며, 역전파를 위해 보관되어야 합니다.
어텐션의 산술 강도는 다음과 같습니다. , 따라서 약 64 FLOP/바이트 — A100 능선 지점의 왼쪽에 있습니다. 평평한 컴퓨팅 천장이 아닌 경사진 대역폭 천장에 위치합니다. GPU는 이동하는 데 더 많은 시간을 소비합니다. 무엇이든 곱하는 것보다 주변에. 이것이 사용하는 지붕선 프레임워크(능선 지점, 경사진 천장과 평평한 천장, 그리고 동일한 추론으로 GPU를 구매할 가치가 있는지 여부를 결정하는 이유)는 GPU가 성과를 거둘 때.
알고리즘이 활용하는 메모리 계층 구조
| 메모리 레벨 | 사이즈 | 대역폭 | 대기 시간 |
|---|---|---|---|
| HBM(고대역폭 메모리) | 40~80GB | 2.0TB/초 | ~400ns |
| SRAM(온칩, 공유 메모리) | 20MB | 19TB/초 | ~4ns |
SRAM은 약 1000분의 1 용량으로 대역폭은 10배, 대기 시간은 100배 더 낮습니다. Flash Attention이 수행하는 모든 작업은 용량을 포기하고 대역폭과 대기 시간을 구매하는 것에서 비롯됩니다. CPU 백테스트에서 측정된 동일한 "하드웨어를 구입하는 대신 알고리즘을 재구성하는" 움직임이 백테스트 속도 사다리.
플래시 어텐션 알고리즘
Flash Attention은 SRAM에 맞는 크기의 타일에서 주의를 처리하며 전체를 구체화하지 않습니다. HBM의 매트릭스.
분할 ~ 안으로 행 블록 및 ~ 안으로 기둥 블록, 타일과 그 축전지가 칩에 맞도록 선택되었습니다. 각 쿼리 블록에 대해 모든 키-값 블록을 반복합니다.
For each query block Q_i:
Initialize: O_i = 0, m_i = -inf, l_i = 0 # output, running max, running sum
For each KV block (K_j, V_j):
1. Load Q_i, K_j, V_j from HBM to SRAM
2. Compute S_ij = Q_i @ K_j^T / sqrt(d) # in SRAM
3. Compute local max: m_ij = rowmax(S_ij)
4. Compute P_ij = exp(S_ij - m_ij) # in SRAM
5. Compute local sum: l_ij = rowsum(P_ij)
6. Update running statistics:
m_new = max(m_i, m_ij)
l_new = l_i * exp(m_i - m_new) + l_ij * exp(m_ij - m_new)
O_i = O_i * (l_i * exp(m_i - m_new) / l_new)
+ P_ij @ V_j * (exp(m_ij - m_new) / l_new)
m_i = m_new, l_i = l_new
Write O_i to HBM
온라인 소프트맥스 재발
타일링을 가능하게 하는 비결은 온라인 소프트맥스입니다. 순진한 소프트맥스에는 행에 대해 두 번의 패스가 필요합니다. 하나는 최대값(수치적 안정성을 위해)을 찾고, 다른 하나는 지수화 및 정규화를 위한 것입니다. 저장을 거부한 행에 두 번 패스하는 것은 모순입니다. 따라서 Flash Attention은 계속해서 통계를 실행하고 진행되는 대로 크기를 조정합니다.
블록 이후 :
출력 누산기는 동일한 요소로 수정됩니다.
새로운 블록이 실행 중인 최대값을 올릴 때마다 이전에 누적된 출력은 다음과 같이 소급하여 재조정됩니다. — 마치 새로운 최대값이 처음부터 알려진 것처럼. 결과는 대수적으로 2패스 소프트맥스와 동일합니다. 정확한 산술에서 이는 근사치가 아닙니다. 그것은 재결합이다. (유한 정밀도에서는 다른 반올림 경로가 중요합니다. 아래의 정확성 확인을 참조하세요.)
IO 복잡성
이것은 승리에 대한 공식 성명입니다. 플래시 어텐션 수행
HBM 액세스, 여기서 SRAM 크기입니다. 표준 구현을 위해. 참고하세요 분모에 나타납니다. 온칩 스크래치패드가 클수록 왕복 횟수가 줄어듭니다. 이것이 바로 알고리즘이 FLOP 개수가 아닌 메모리 계층 구조의 관점에서 기술되는 이유입니다. 일반적인 경우 그리고 KB의 경우 이 비율은 대략 5~10배 적은 액세스로 Flash Attention을 선호합니다.
역방향 전달: 저장 대신 다시 계산
주의를 통한 역전파에는 일반적으로 다음이 필요합니다. 정방향 패스가 유지를 거부한 행렬입니다. Flash Attention은 타일을 재계산합니다. 역방향 패스 동안 출력만 저장 그리고 소프트맥스 통계 - 둘 다 , 아니다 . 전체 문제였던 메모리 용어에 대해 적당한 양의 중복된 산술을 교환합니다. 이는 단일 연산자 내부의 타일 세분성에서 적용되는 그라데이션 체크포인트와 동일한 할인입니다.
FA2: 병렬성
Flash Attention 2(Dao, 2023)는 알고리즘을 유지하고 일정을 수정했습니다.
- 비매트뮬 FLOP 수가 적습니다. FA1은 텐서 코어가 아닌 CUDA 코어에서 실행되는 작업인 크기 조정, 최대값 찾기 및 지수화에 실시간을 소비했습니다. FA2는 내부 루프의 끝으로 크기 조정을 연기합니다.
- 시퀀스 길이에 대한 병렬성. FA1은 배치 및 헤드에 대해서만 병렬화됩니다. FA2는 쿼리 블록에 대해서도 병렬화합니다. 이는 특히 자산당 하나의 매우 긴 시퀀스*와 배치 크기가 1~4인 거래 사례에 특히 중요합니다. 바로 배치 및 헤드 병렬 처리가 GPU를 고갈시키는 체제입니다.
- 워프 작업 분할. 각 워프는 점수 계산을 분할하고 워프 전체에 걸쳐 축소하여 교차 워프 감소를 제거하는 대신 쿼리 블록의 서로 다른 하위 집합을 사용합니다.
보고된 결과: A100의 이론적 피크 FLOP의 ~70%, FA1의 경우 ~35%.
FA3: 호퍼 역학
Flash Attention 3(Dao, Shah, 2024)은 H100의 아키텍처에 따라 다릅니다.
- 비동기 워프 전문화. Hopper의 TMA(Tensor Memory Accelerator)는 HBM→SRAM을 비동기적으로 이동합니다. FA3는 워프를 다음 KV 블록에 대해 TMA 로드를 발행하는 생산자와 현재 KV 블록을 계산하는 소비자로 분할하므로 데이터 이동이 산술 뒤에 숨겨집니다.
- 인터리브된 matmul 및 소프트맥스. 한 블록의 경우 텐서 코어에서 실행되는 반면 이전 블록의 소프트맥스는 CUDA 코어에서 실행됩니다. 두 개의 서로 다른 하드웨어 장치는 시간 분할이 아닌 실제로 동시입니다.
- 일관되지 않은 처리를 사용하는 FP8. H100은 FP16 처리량의 2배로 FP8을 수행합니다. 순진한 FP8 관심은 이상값으로 인해 손상됩니다. FA3는 블록 단위 양자화 전에 벡터를 무작위로 회전하여 좌표 전체에 걸쳐 이상치 크기를 분산시키며, 순진한 FP8보다 2.6배 낮은 수치 오류로 보고됩니다.
| 버전 | GPU | 활용도 | 속도 향상 대 표준 |
|---|---|---|---|
| FA1 | A100 | ~35% | 2-4배 |
| FA2 | A100 | ~70% | 5-7배 |
| FA3(FP16) | H100 | ~75% | 3-5x 대 FA2 |
| FA3(FP8) | H100 | ~75% | 1.6x 대 FA3 FP16 |
인과 마스킹은 거래가 할인되는 곳입니다.
인과 마스킹은 시계열에 필수입니다. 모델은 미래에 주의를 기울여서는 안 됩니다. 타일링에서는 추가 비용이 아니라 절약되는 비용입니다. 키가 쿼리와 관련하여 완전히 미래에 있는 타일은 완전히 건너뛰고 로드되지도 계산되지도 않으므로 작업이 대략 절반으로 줄어듭니다. PyTorch에서는 다음과 같습니다. is_causal=True; 다른 것은 필요하지 않습니다.
통합은 8줄입니다
필요한 코드 중 Flash Attention에 관한 코드는 거의 없습니다. 융합 커널에 대한 명시적 점수 매트릭스 경로를 바꿉니다.
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_head)
scores = scores.masked_fill(causal_mask, float("-inf"))
out = torch.matmul(torch.softmax(scores, dim=-1), v)
out = torch.nn.functional.scaled_dot_product_attention(
q, k, v,
dropout_p=self.dropout if self.training else 0.0,
is_causal=True,
)
그것이 전체 변화입니다. q, k, v 모양이 있다 (batch, heads, seq, head_dim); 인과 마스크는 커널이 빌드하기 때문에 사라졌습니다. 완전한 주석이 달린 PyTorch 거래 모델을 입력 투영, 블록 및 3클래스 업/플랫/다운 헤드에 적용하려면 DeepLOB, 전체 학습 파이프라인은 [Temporal Fusion Transformers](를 참조하세요./en/blog/post/temporal-fusion-transformer-trading). 여기에 비계의 네 번째 복사본을 만드는 것은 아무 것도 가르쳐주지 않을 것입니다.
요구사항: 컴퓨팅 능력 >= 8.0(A100, H100, RTX 3090+), 반정밀도 입력, PyTorch >= 2.0. 실제로 연결된 빠른 경로를 확인합니다. torch.backends.cuda.sdp_kernel 진단 및 torch.cuda.max_memory_allocated() — SDPA는 전제 조건이 실패할 경우 자동으로 수학 커널로 폴백하며, 자동 폴백은 단지 느린 작업 모델과 똑같아 보입니다.
측정: 긴 컨텍스트가 도움이 됩니까?
위의 모든 내용은 이제 32K 또는 128K 컨텍스트가 저렴한 수준임을 나타냅니다. 유용한지 여부에 대해서는 아무 말도 하지 않습니다. 정직한 실험:
동일한 아키텍처를 교육합니다. 이 시리즈의 다른 곳에서 사용되는 BTC 시리즈의 SDPA-Flash를 사용하면 매개변수, 최적화 프로그램 및 대상이 고정되어 시퀀스 길이가 유일한 변수가 됩니다. 두 가지 사항을 보고합니다.
- 비용. 에포크당 측정된 벽시계 및
torch.cuda.max_memory_allocated()각각 . - 이점. 표본 외 예측 성능 대 , 앞으로 나아가는 분할에서.
이 기사의 초기 초안에는 다음에서 분석적으로 파생된 시퀀스 길이별 메모리 수치 표가 포함되어 있습니다. 활성화 공식. 해당 행은 제거됩니다. 측정된 적이 없으며 기사 자체의 메모리 예산 산술에 동의하지 않습니다. 결과표에 제시된 파생 수치는 조작된 결과이며, 본 블로그에서는 이를 제공하지 않습니다.
이 실험의 흥미로운 특징은 어느 방향으로든 게시할 수 있다는 것입니다. 샘플 외 성능이 단조롭게 증가하는 경우 , 이는 전체 긴 컨텍스트 프로그램을 정당화합니다. 그것이 수천 단계에서 정체되거나 저하된다면, 그것은 [정직한 부정](의 동반자인 더 강력한 부분입니다./en/blog/post/honest-negative-no-robust-edge) — 이는 메모리 벽이 변압기 거래에 대한 구속력이 결코 아니었다는 것을 의미합니다.
컨텍스트가 많을수록 용량이 많아지므로 표면이 더 과적합됩니다.
평탄하거나 부정적인 결과를 기대하는 데에는 특별한 이유가 있습니다. 시간 융합 트랜스포머 바닐라 변환기가 순진하게 금융 시리즈 과적합에 적용되었음을 이미 문서화했습니다. 시간적 귀납적 편향이 부족하고 단기 검토 순환 모델이 높은 빈도에서 경쟁력을 유지합니다. 컨텍스트를 512단계에서 32,768단계로 확장해도 길이에 비례하는 정보가 추가되지 않습니다. 거의 효율적인 가격 시리즈의 한계 32,000번째 지연은 거의 전달되지 않습니다. 안정적으로 추가하는 것은 매개변수에 맞는 항목의 가치입니다.
그래서 대청소 이 블로그가 다른 모든 검색에 적용되는 것과 동일한 메커니즘을 사용하는 모델 선택 검색으로 취급되어야 합니다. 3개의 시퀀스 길이와 그 밖의 다른 모든 것이 시도 횟수이며, 승자는 단순히 이웃을 이기는 것이 아니라 해당 시도 횟수와 PBO 게이트에 대해 계산된 수축된 샤프 비율을 클리어해야 합니다. 그렇지 않으면 "긴 컨텍스트 승리"는 세 번의 시끄러운 실행 중 가장 좋은 것을 선택하는 것과 구별할 수 없습니다.
정확성 확인, "정확함"은 많은 작업을 수행하기 때문입니다.
Flash Attention은 정확한 산술에서 정확합니다. 첨부된 권장 사항(fp16 또는 bf16에서 실행하고 H100에서는 FP8을 고려)은 그렇지 않습니다. 이는 별도의 주장이며 두 번째 주장이 실제로 지배적입니다. 합계를 다시 연결하고 정밀도를 절반으로 낮추는 것은 모두 혼란스러운 일이며 순서 보장을 도입한 기사는 정밀도를 손에 흔들어서는 안 됩니다.
블로그에는 이미 올바른 도구가 있습니다. GPU 정밀 함정은 표준을 설정합니다. 정밀도가 낮으면 경고하지 않고 그럴듯한 쓰레기를 반환하며 눈에 띄는 곡선이 아닌 다운스트림 이산 수량(거래 횟수)에 대한 패리티 오라클을 사용하여 정확성을 증명합니다. 여기에 적용됨:
- bf16의 SDPA-Flash와 동일한 입력에 대한 fp64 참조 구현을 사용하여 주의를 계산합니다. 출력 텐서에 대한 최대 상대 오류를 보고합니다.
- 결정까지 밀어붙입니다. 위쪽/평평/아래쪽 라벨을 내보내는 모델의 경우 두 경로 사이에 플립되는 라벨 수를 전체 결정의 일부로 보고합니다.
작고 제한적이며 설명 가능한 불일치는 올바른 빠른 경로의 특징입니다. 제한이 없다는 것은 FP8 권장 사항이 이 모델에 결코 안전하지 않다는 것을 의미합니다. 실행될 때까지 두 번호 모두 알 수 없습니다.
언제 도달해야 할까요?
- CUDA GPU에서 ~2K 시간 단계 이상: 예, 무조건입니다. 정확한 출력을 생성하는 한 줄의 변경이며 승리는 다음과 같이 커집니다. . 구체화되기를 원하는 시나리오는 없습니다. 대신 경로.
- ~512 타임스텝 미만, CPU 또는 비어텐션 아키텍처(CNN, Mamba와 같은 SSM): 관련이 없습니다. 능선 왼쪽에서는 고정 오버헤드가 전체 비용이며 주의는 병목 현상이 발생하지 않습니다.
- 위의 임계값은 측정이 아닌 민간 전설입니다 — 이는 일반 문헌에서 가져온 것이며 자체 모델과 카드의 교차는 10줄 벤치마크입니다. 라운드 숫자를 신뢰하기보다는 실행해 보세요.
결론
Flash Attention은 명확하고 정말 중요한 결과입니다. 메모리 계층 구조를 존중하고 소프트맥스를 다시 연결함으로써 정확한 주의를 계산합니다. 대신에 기억 , 그리고 IO 바운드가 그 이유를 정확하게 설명합니다. 이를 트레이딩 변환기에 채택하는 것은 정확한 산술에서 정확성 비용이 들지 않고 큰 메모리 승리를 거두는 한 줄 변경입니다.
그것이 하지 않는 일은 맨 위에 있는 질문에 대답하는 것입니다. "하루 종일 컨텍스트는 불가능합니다"를 "하루 종일 컨텍스트는 저렴합니다"로 변환합니다. 이는 결과가 아닌 실험 비용의 변화입니다. 무너지는 기억의 벽은 측정하라는 초대이며, 측정은 이것을 종이 요약에서 발견으로 바꾸는 것입니다.
참고자료
- Dao, T., Fu, D.Y., Ermon, S., Rudra, A., Re, C. "FlashAttention: IO 인식을 통한 빠르고 메모리 효율적인 정확한 주의." NeurIPS(2022). arXiv:2205.14135
- Dao, T. "FlashAttention-2: 더 나은 병렬성과 작업 분할을 통한 더 빠른 주의." ICLR(2024). arXiv:2307.08691
- Dao, T., Shah, J. "FlashAttention-3: 비동기 및 낮은 정밀도로 빠르고 정확한 주의." NeurIPS(2024). arXiv:2407.08608
- Vaswani, A., et al. "주의가 필요한 전부입니다." NeurIPS(2017).
- Milakov, M., Gimelshein, N. "소프트맥스에 대한 온라인 정규화 계산." arXiv:1805.02867 (2018).
Authors
Trading-systems engineer
Trading-systems engineer building bots since 2017: cross-exchange arbitrage (connected up to 30 venues), cointegration-based pairs arbitrage across spot and futures, scalping, news and sentiment-driven strategies, trend algorithms, and portfolio management and balancing algorithms. Also builds sub-millisecond order execution, big-data warehouses, backtesting engines, AI agents, and trading interfaces (incl. open-source profitmaker.cc). Stack: JS/TS, Python, Rust/Zig/Go, DevOps, backend, frontend, architecture.