Aralıklı Kontrol Noktası ve Aktifleştirme Yeniden Hesaplama
Type: Build
Languages: Python (with numpy, optional torch)
Prerequisites: Phase 10 Lesson 04 (Pre-Training Mini-GPT), Phase 10 Lesson 05 (Scaling & Distributed)
Time: ~70 minutes
Sorun
Bir transformatör eğitimi, her katman için geriye ayrılmış her operasyonun girişlerini saklar: dikkat girişleri, Q/K/V projeksiyonları, softmax çıkışı, FFN girişleri, norm çıkışları ve kalan akım. Gizli boyutlu bir katman için d, dizi uzunluğu L, partiBBu , emir üzerine .12 B L * dkatman başına yüzen.
- Evet .
d=8192, L=8192, B=1Bu, BF16'da katman başına 800 MB'dir. 64 katlı bir model 51 GB aktivasyonlara sahiptir.L^2(bkz: baş başına) ve tensor paralel kısmi kopyaları oluşturmadan önce.
İki taraflı fatur: BF16 ağırlıkları artı optimizer durumu 80GB'ye uygun olabilir, ancak etkinleştirmeler sizi öteye itebilir. Gradient kontrol noktası (aka activation recalculation) standart düzeltme. Çoğu etkinleştirmeyi bırakın; geriye dönerken ileriyi tekrar yapın.
Naifce yapıldığında, kontrol noktası, adım başına yaklaşık %33 daha fazla ileri geçiş FLOP maliyetindedir. İyi yapıldı Korthikanti et al.'ın "akıllı seçimi" başına seçici kontrol noktası 5x hafıza tasarruf edersiniz.
Anlaşım
Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geriye Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri Geri
output = layer(input)Geriye dönmek istiyor .grad_inputve grad_paramsOnları hesaplamak için:
input(Bilgilemek içingrad_params = input.T @ grad_outputDüzsel katmanlar için)- bazı aktive derivatifler arası (ReLU/GELU/softmax'ın derivatifleri aktive değerine bağlıdır)
Ön geçit otomatik olarak otograd grafikinde depolanır.tensor.retain_grad()ve girişine ihtiyacı olan her operasyon bir referans tutar.
Tam Kontrol Noktası Saçma
Ağı ikiye böl .NÖnceki bölümler. Önceki bölümler sırasında, her bölüm için sadece giriş depolayın. Geriye geçişler gerektiğinde, segmentin önceki geçişini yeniden çalıştırın, sonra farklılaştırın.
Örnek: 32 katmanlı transformatör, her katman 1 katmanlı 32 bölüme ayrılmıştır.
- Hatıra: 32 katman giriş (küçük) vs 32 * (katman başına etkinleştirme hacmi) (çok büyük).
- Ekstra hesaplama: Segmente başına 1 ekstra ileri, yani %33 daha fazla ileri FLOP toplamı (geriye doğru 2x ileri olduğu için, tam adım 1 + 1 + 2 = 4 birim yerine 1 + 2 = 3 olur).
Bu Chen et al. 2016 tarihli orijinal tarifi: her bir kontrol noktası sqrt(L)L=64, 8 kontrol noktası.
Seçimsel Kontrol Noktası (Korthikanti 2022)
Tüm etkinleştirmeler aynı maliyetli değil.BLLheadsFFN gizli etkinliği BL*4dUzun sekanslarda softmax hakimdir.
Seçimsel kontrol noktası, ucuz depolama aktivasyonlarını (lineer projeksiyonlar, kalıntılar) tutar ve sadece pahalı olanları (özen) yeniden hesaplar.
Megatron-Core bunu "seçici" etkinleştirme yeniden hesaplama olarak uyguluyor.
Çıkarım
Yeniden hesaplama alternatifleri: ileri ve geriye doğru CPU RAM'a devreye aktarma. PCIe bant genişliği gerektirir; boş bant genişliği yeniden maddeleşme maliyetinden fazla olduğunda yararlıdır. Karışık stratejiler yaygın: bazı katmanları kontrol et, diğerlerini boşalt.
FSDP2 birinci sınıf bir seçenek olarak yükten çıkartır. GPU hafıza boğazında bulunduklarında yükten çıkartır.
Ücret Modelini Yeniden Hesapla
Her adımda saf bir kontrol noktası ile FLOPs .kkatmanları L- ...
flops_fwd_normal = L * f_layer
flops_bwd_normal = 2 * L * f_layer
flops_total_normal = 3 * L * f_layer
flops_fwd_ckpt = L * f_layer
flops_recompute = L * f_layer # one extra forward per layer in the segment
flops_bwd_ckpt = 2 * L * f_layer
flops_total_ckpt = 4 * L * f_layer
overhead = 4 / 3 - 1 = 0.33 = 33%Seçimsel kontrol noktası ile sadece dikkat çekirdeğini yeniden hesaplarsınız, tüm katmanı değil:
flops_recompute_selective = L * f_attention ~= L * f_layer * 0.15
overhead_selective = (3 + 0.15) / 3 - 1 = 0.05 = 5%Hatıra Kaydetme Modülü
Katman başına etkinleştirme hacmi: A- Evet .Lkatmanlar, toplam aktivasyon hafızası: L * A- Evet .
Tam kontrol noktası (sektör boyutu 1): sadece depolama L input_volume(~L 1/10 AStandart bir transformatör için).9 L A * 1/10- Evet .
Kontrol noktası her zaman .kkatmanlar: depolama L/k * AEk olarak .k-1aktif segment içindeki katmanların değeri.
- Evet .
k = sqrt(L), bellek ve yeniden hesaplama maliyeti hem ölçeklesqrt(L)En iyi fiyat değişikliği.
Kontrol Noktasına Ne Zaman Gitmemek
- Bir boru hattının en iç katmanları uçuşta zaten.
- Eğlence hesabına hükmeden ilk ve son katmanlar (transformatörlerde nadirdir).
- FlashAttention'ı kullanan dikkat çekirdekleri Flash zaten softmax hızını yeniden hesaplar, bu yüzden ek katman seviyesindeki kontrol işaretlemeyi üstte biraz ekler.
Uygulama Şekilleri
- Function wrapper:Bir bölümü içine sarın
torch.utils.checkpoint.checkpoint(fn, input)Sadece PyTorch mağazaları .input, geriye dönüp her şeyi yeniden hesaplar.
- Decorator-based:Etiketlemenin kontrol noktası olarak yapılması gereken katmanlar; eğitmen, hangi bölümlerin toplanıp sarılacağına konfigürasyon zamanında karar verir.
- Manual explicit recompute:Sıradan bir alışkanlık olarak geriye geçmeyi kendin yaz.
recompute_forwardÖncekiyi depolanan giriş ile çiftleştirir.
Üçü de aynı fonksiyonel sonuç verir.
TP / PP / FP8 ile etkileşim
- Tensor parallel:Kontrol noktası girişleri yeniden hesaplama sırasında toplanmalı veya yeniden dağıtılmalıdır; iletişim maliyetini karşılamak.
- Pipeline parallel:Tipik bir örnektir. Her boru hattının aşamasının ileriye doğru kontrol edilmesi böylece geri sıra mikrobatçlar aktifleşme belleğini yeniden kullanabilmektedir.
- FP8 recompute:amax tarihleri yeniden hesaplama sırasında güncellenmiş orijinal ileri veya FP8 ölçek sürüşleri ile eşleşmelidir.
Yapın
Adım 1: Bölümlerle Oyuncak Model
pythonimport numpy as np
def linear_forward(x, w, b):
return x @ w + b
def relu(x):
return np.maximum(x, 0)
def layer_forward(x, w1, b1, w2, b2):
h = relu(linear_forward(x, w1, b1))
return linear_forward(h, w2, b2)
def model_forward(x, params):
activations = [x]
h = x
for w1, b1, w2, b2 in params:
h = layer_forward(h, w1, b1, w2, b2)
activations.append(h)
return h, activationsİkinci Adım: Geriye Alışmak İçin Tüm Aktivasyonlara İhtiyaç Var
pythondef model_backward(grad_output, activations, params):
grads = [None] * len(params)
g = grad_output
for i in range(len(params) - 1, -1, -1):
w1, b1, w2, b2 = params[i]
x_in = activations[i]
h_pre = linear_forward(x_in, w1, b1)
h = relu(h_pre)
gh = g @ w2.T
gw2 = h.T @ g
gb2 = g.sum(axis=0)
g_pre = gh * (h_pre > 0)
gx = g_pre @ w1.T
gw1 = x_in.T @ g_pre
gb1 = g_pre.sum(axis=0)
grads[i] = (gw1, gb1, gw2, gb2)
g = gx
return g, gradsAdım 3: Kontrol Noktası-Her-k hafıza
pythondef model_forward_checkpointed(x, params, k=4):
saved_inputs = [x]
h = x
for i, (w1, b1, w2, b2) in enumerate(params):
h = layer_forward(h, w1, b1, w2, b2)
if (i + 1) % k == 0:
saved_inputs.append(h)
return h, saved_inputs
def model_backward_checkpointed(grad_output, saved_inputs, params, k=4):
grads = [None] * len(params)
g = grad_output
segments = [(j * k, min((j + 1) * k, len(params))) for j in range(len(saved_inputs))]
for seg_idx in range(len(saved_inputs) - 1, -1, -1):
start, end = segments[seg_idx]
if start >= end:
continue
x_in = saved_inputs[seg_idx]
_, seg_acts = model_forward(x_in, params[start:end])
g, seg_grads = model_backward(g, seg_acts, params[start:end])
for j, gr in enumerate(seg_grads):
grads[start + j] = gr
return g, gradsDördüncü Adım: Maliyet modeli
pythondef checkpoint_cost(n_layers, segment_size, flops_per_layer=1.0):
fwd = n_layers * flops_per_layer
recompute = n_layers * flops_per_layer
bwd = 2 * n_layers * flops_per_layer
return {
"fwd": fwd,
"recompute": recompute,
"bwd": bwd,
"total": fwd + recompute + bwd,
"overhead_vs_no_ckpt": (fwd + recompute + bwd) / (fwd + bwd) - 1.0,
}
def selective_checkpoint_cost(n_layers, attention_fraction=0.15,
flops_per_layer=1.0):
fwd = n_layers * flops_per_layer
recompute = n_layers * attention_fraction * flops_per_layer
bwd = 2 * n_layers * flops_per_layer
return {
"fwd": fwd,
"recompute": recompute,
"bwd": bwd,
"total": fwd + recompute + bwd,
"overhead_vs_no_ckpt": (fwd + recompute + bwd) / (fwd + bwd) - 1.0,
}Adım 5: Hatıra Tahminici
pythondef activation_memory_mb(n_layers, hidden=8192, seq=8192,
batch=1, bytes_per_value=2):
per_layer = 12 * batch * seq * hidden * bytes_per_value
return n_layers * per_layer / 1e6
def memory_after_checkpoint(n_layers, segment_size, hidden=8192,
seq=8192, batch=1, bytes_per_value=2):
n_seg = max(1, n_layers // segment_size)
saved = (n_seg + segment_size) * 1 * batch * seq * hidden * bytes_per_value
return saved / 1e6Adım 6: Optimal Bölüm Boyutu
pythondef optimal_segment(n_layers):
return int(round(np.sqrt(n_layers)))Adım 7: Seçimçi Kontrol Noktası Kararı
pythondef should_recompute(layer_type, activation_bytes, recompute_flops_ratio):
if layer_type == "attention" and activation_bytes > 100 * 1e6:
return True
if layer_type == "ffn" and activation_bytes > 500 * 1e6:
return recompute_flops_ratio < 0.1
return FalseKullan
- torch.utils.checkpoint- Evet .
from torch.utils.checkpoint import checkpointPyTorch'daki kanonik ambalaj. Bir fonksiyonu sarar; sadece girişleri saklar, geriye doğru yeniden hesaplar. - Megatron-Core activation recomputation: destekler
selective- Evet .fullveblock2024+ sınır eğitiminde standart. - FSDP2 offload- Evet .
module.to_empty(device="cpu")- Evet .offload_policyFSDP2'de yeniden hesaplama yerine CPU'ya etkinleştirmelerini kısaltır. - DeepSpeed ZeRO-Offload: Optimizer durumları ve etkinleştirmeleri için CPU yükü çıkartmak, kontrol noktasını tamamlamak.
Gönder
Bu ders bize çok yararlı .outputs/prompt-activation-recompute-policy.md model yapılandırmasını (katmanlar, gizli, seq, parti) ve mevcut GPU belleğini alan ve katman başına yeniden hesaplama politikasını (hiçbir / seçici / tam / yüklenme) yayınlayan bir istek.
Egzersizler
- Doğru olduğunu kontrol et.
model_forward+model_backward(tam aktivasyon) vsmodel_forward_checkpointed+model_backward_checkpointedParametre gradiyenti makinenin hassasiyetine eşit olmalıdır.
- Tarama bölümü boyutu
k1 ' denL- FLOP'u ve hafızayı çiz.
- Seçimsel kontrol işaretlemeyi uygulayın: dikkat modülünün girişini, ancak aralarını değil saklayın. 32 katlı bir model için FLOP üst üstlük vs tam katman kontrol işaretlemesini seq=8192'de ölçün.
- Çıkarma ekleyin. Segment girişlerini simülasyonlu bir "CPU tamponu"na (ayrı bir liste) kaydetin. "PCIe bant genişliği" byte/zaman olarak ölçün ve çıkarma ve yeniden hesaplama arasındaki kesinti noktasını bulun.
- Gerçek PyTorch transformatörünü , içinde ve dışında bir referans göster .
torch.utils.checkpoint. hafıza ölçümleri (dentorch.cuda.max_memory_allocated) ve adım zaman.
Anahtar Terimler
| Term | What people say | What it actually means |
|---|---|---|
| Gradient checkpointing | "Save memory by redoing forward" | Store segment inputs only; recompute intermediates during backward to get gradient-support tensors |
| Activation recomputation | "Same as checkpointing" | The HPC-flavored name for the same technique |
| Segment size (k) | "How many layers per checkpoint" | Number of layers whose intermediates are dropped and rematerialized together |
| Selective checkpointing | "Korthikanti's trick" | Recompute only expensive-to-store activations (attention softmax); keep cheap ones |
| Full checkpointing | "The naive version" | Recompute every layer's intermediates in every segment |
| Block checkpointing | "Coarse-grained" | Checkpoint whole transformer blocks; largest granularity |
| FLOP overhead | "The compute tax" | Extra FLOPs per step = (recompute FLOPs) / (fwd + bwd FLOPs); 33% naive, 5% selective |
| Activation offload | "Ship to CPU" | Move activations to CPU RAM across forward->backward; alternative to recompute |
| sqrt-L rule | "The classical optimum" | For uniform-cost layers, optimal checkpoint spacing is sqrt(L) layers |
| Attention-softmax volume | "The O(L^2) problem" | L^2 heads batch floats; dominates activation memory at long contexts |
Daha Fazla Okumak
- Chen et al., 2016 -- "Training Deep Nets with Sublinear Memory Cost"- ...diğerleri kontrol etmek için resmileştirilen orijinal kağıt.
- Korthikanti et al., 2022 -- "Reducing Activation Recomputation in Large Transformer Models"-- Seçkin etkinleştirme yeniden hesaplama ve resmi maliyet analizi
- Pudipeddi et al., 2020 -- "Training Large Neural Networks with Constant Memory using a New Execution Algorithm"-- ters modunda yeniden maddeleşme yoluyla alternatif sabit hafıza yaklaşımı
- Ren et al., 2021 -- "ZeRO-Offload: Democratizing Billion-Scale Model Training"-- Ölçüsünde aktifleştirme yükü
- PyTorch torch.utils.checkpoint docs-- Standart API
- Megatron Bridge activation recomputation documentation-- Seçkin, tam ve blok modları
This free lesson is part of the AI Engineering from Scratch curriculum. Read the full explanation, run the lesson code, and verify the result in the interactive reader or from the repository source.
Browse the complete course catalog or open this lesson on GitHub.