Bài 6 — Từ JEPA formulation tới một training step#
Bắt đầu · Trước: asymmetry · Tiếp: masking
Mục tiêu và tiền đề#
Bạn sẽ vẽ luồng tensor/gradient, chọn target positions sau teacher contextualization, tính một optimizer/EMA step và phân biệt paper equation, released code/config và checkpoint provenance. Cần patch shapes/attention ở lớp 1, loss và detach ở bài 2/5. Luồng chính là Audio-JEPA; I-JEPA được đối chiếu để hiểu thiết kế gốc. Không chạy pretraining ở bài này.
1. JEPA là ý tưởng predictive embedding, recipe phải đọc riêng#
Formulation khái quát: từ embedding của context x, predictor với conditioning z dự đoán embedding của target y. z có thể chỉ vị trí/transformation; nó không mặc định là latent random variable phải sampling. Một energy so compatibility trong embedding space chưa là density của audio. I-JEPA §2, Fig. 2.
Trong I-JEPA/Audio-JEPA teacher–student recipe đang xét, context encoder , predictor , target encoder . Tập context , target ; ở mask pairing này. Teacher nhận full input, sau đó chọn output ở M. Đây không là định nghĩa bắt buộc của mọi JEPA; bài 10 có recipes khác.
flowchart LR
X[Full spectrogram X] --> V[Patchify + positions + select C]
V --> E[Context encoder theta]
E --> P[Predictor phi + positions M]
X --> T[Full-input target encoder theta_bar]
T --> G[Select outputs at M; no gradient]
G --> N[Target normalization]
P --> L[Loss on M]
N --> L
L -. gradient .-> P
L -. gradient .-> E
E -. copy encoder weights by EMA .-> TXem mã sơ đồ
flowchart LR
X[Full spectrogram X] --> V[Patchify + positions + select C]
V --> E[Context encoder theta]
E --> P[Predictor phi + positions M]
X --> T[Full-input target encoder theta_bar]
T --> G[Select outputs at M; no gradient]
G --> N[Target normalization]
P --> L[Loss on M]
N --> L
L -. gradient .-> P
L -. gradient .-> E
E -. copy encoder weights by EMA .-> TSolid arrows là data, dashed gradient chỉ prediction/context paths; EMA arrow là weight update, không backprop. Không copy predictor vào target encoder.
2. Trace shape Audio-JEPA với một mask cụ thể#
Khai báo axes time trước, mel sau theo train/data wiring được đọc ở commit Audio-JEPA ddd97ee9b88c00572f59bf2a00eb42120e6a4dea. Input là . Native patch 16 time × 16 mel → grid 16 × 8, . Width encoder 768; predictor hidden 384, 6 blocks, outputs width 768. train.py, phần tạo H/W và instantiate model, VisionTransformer/Predictor.
Chọn một context mask, một target mask, cho toy recipe trace: . Không claim tất cả batches là 50%; config samples ratio 40–60%.
| Bước | Giá trị/shape | Thứ cần giữ đúng |
|---|---|---|
| Patch projection | B × 128 × 768 | Conv kernel=stride=(16,16), không overlapping patches |
| Add positions; gather C | B × 64 × 768 | Original grid indices được giữ; không renumber positions thành 1…64 |
| Context blocks + final LN | B × 64 × 768 | Transformer chỉ xử lý visible rows |
| Predictor projection | B × 64 × 384 | 768→384 learned projection |
| Append learned mask tokens + target positions | B × 128 × 384 | Context 64 + target 64; mask token không có raw target patch |
| Predictor blocks + norm; select prediction rows | B × 64 × 384 | Outputs ở target suffix trong implementation |
| Predictor output projection | B × 64 × 768 | Width trở về target space |
| Teacher full forward + final LN | B × 128 × 768 | Teacher attention đã dùng full input |
| Gather target M, no-grad | B × 64 × 768 | Gather đúng cùng order với predictor |
| Criterion target processing/loss | Scalar | Loss theo recipe code dưới đây |
Nếu dùng nhiều context/target masks, implementation repeats/gathers batches; không tự gọi shape vẫn là B như case một mask. Position/output order mismatch làm loss so nhầm đáp án dù shapes giống nhau. model_step trong jepa_module.py.
Teacher patch j là output đã contextualize, không embedding của crop j chạy riêng. Encoder final LayerNorm khác criterion target standardization, cần trace cả hai.
3. Loss paper và code không hoàn toàn giống nhau#
Audio-JEPA v2, §III.B Eq. 2, §IV viết average squared L2 trên masked patch vectors, target EMA. Released jepa.yaml chọn loss_type: norm_mse, norm_pix_loss: true. Loss.forward và norm_mse_loss thực hiện:
Variance dùng target.var(dim=-1), tức sample/correction 1 theo PyTorch default; không đúng cùng variance divisor với parameter-less LayerNorm. Unit normalization thực dùng F.normalize epsilon floor cho norm nhỏ; công thức chia norm phía trên giả định nonzero/nondegenerate, không thay exact zero behavior. Chỉ target được center/standardize trong criterion; cả target và prediction được L2-normalize. Không viết thành center cả hai nếu code không làm vậy.
Toy: target raw , prediction . Mean target 2; centered target . Sau standardization rồi unit normalization, target direction là , trùng normalized p; code-like loss 0. Raw squared L2 giữa p và t là 8. Nếu , cosine với normalized centered target là , loss . Ví dụ này giải thích geometry của loss, không tensor đo từ checkpoint.
I-JEPA cũng có distinction: paper §3 viết squared L2; official train.py commit 52c1ae95… teacher outputs được F.layer_norm rồi F.smooth_l1_loss với mean reduction. I-JEPA code không đơn giản là raw MSE giống công thức trong map. Dossier ghi đầy đủ paper/code đối chiếu.
4. Ai nhận gradient và ai được cập nhật?#
Một step code-like, bỏ batching/AMP để thấy cơ chế:
X, C, M = sample_and_mask()
H_C = context_encoder_theta(X, C)
prediction = predictor_phi(H_C, C, M)
with no_gradient:
target = target_encoder_theta_bar(X)[:, M, :]
loss = criterion(prediction, target)
backward(loss) # theta and phi only
optimizer_step(theta, phi)
EMA_update(theta_bar, theta)Mask sampling/index selection không có trainable parameter gradient. Fixed positional tables không optimizer gradient; learned mask token, predictor projections/blocks/norm parameters và context encoder parameters nằm trên prediction path. Teacher no-grad; target architecture copy có weights đổi bằng EMA. Optimizer groups ở code chỉ encoder/predictor; biases/1D parameters có weight-decay exclusions, không suy mọi parameter cùng decay.
Numerical backprop toy, deliberately scalar/no norm/no attention, để kiểm chain rule: context input c=1, full teacher input x=2. , , , . Mid-training state → prediction 6, target 7, loss 0,5.
SGD eta 0,1 updates simultaneously using old gradients: . EMA tau 0,9 then . Next full prediction 7,36, target 6,76, residual 0,60, loss 0,18. Teacher initialization in actual recipe starts as encoder copy; unequal weights ở toy là một state đã đi qua training, không initialization rule.
Nếu update b rồi recompute gradient a trong cùng calculation, đó là thuật toán khác. Nếu target gets gradient descent thay EMA, cũng là thuật toán khác. Numerical script kiểm toy/finite differences với target cached; nó không chạy torch/model.
5. EMA schedule và optimizer schedule cần đọc lifecycle#
Audio-JEPA code gọi EMA qua on_train_batch_end; callback updates weights bằng current tau rồi advances tau theo cosine, default 0,996→1. Schedule denominator dùng số batches/epoch × max_epochs, còn numerator dùng global_step. MAWeightUpdate callback.
Trong case automatic optimization, một batch = một optimizer update không bị skipped, có thể đọc trace optimizer→EMA như trên. Accumulation, AMP skipped steps hoặc max_steps override có thể làm quan hệ batch/global_step khác; chưa chạy runtime để audit các edge cases. Không tự biến hook batch-end thành theorem “mỗi successful optimizer step đúng một EMA”. Buffers và normalization running stats cũng cần kiểm khi architecture có chúng; pinned callback iterates named parameters.
Config AdamW khai báo LR 3e−4/WD 0,05, nhưng LR scheduler reference là 1e−3 và WD scheduler reference/final 1e−6. Scheduler writes effective values vào parameter groups. Đây khác recipe prose v2 §IV.C. WarmupCosineScheduler, CosineWDScheduler.
HF card/config revision ngày 16/07/2026 là documentation evidence; checkpoint file được phát hành trước đó. Không có embedded training config/log/hash của bản local bạn đã dùng được audit ở đây. Do vậy kết luận đúng là paper và released recipe khác nhau, chưa là “checkpoint của paper mình chắc chắn train bằng effective values này”. Mismatch loss/geometry không phủ định downstream results; nó giới hạn precision của lời giải thích pretraining.
6. Native geometry và downstream geometry#
Native code axes time256 × mel128, patch16×16 → 16 time × 8 mel, 128 tokens. HF config ghi grid_time 8/grid_freq 16 có axis discrepancy với chính dimensions và code; dùng explicit shape/code cho trace, ghi discrepancy trong dossier.
Paper SOICT §3 reports downstream time256 × mel128, patch 8 time × 32 mel → 32 × 4, 128 tokens. Patch weights interpolate, positional embeddings regenerate. Token count equal không giữ physical support equal: native hop khoảng 39 ms/time frame, downstream 10 ms. Same 128 rows không có nghĩa cùng nhiệm vụ audio. Đây là reported method, chưa chạy tensor audit. Giải phẫu paper.
Paper dùng context encoder; upstream Audio-JEPA §III.C evaluates target encoder. Context và EMA weights không mặc định bằng nhau. Downstream không có predictor/target; score . Bài 9 tách encoder transfer khỏi error scoring.
Bài tập tăng dần#
- Vì sao chọn M sau teacher full forward khác crop M rồi chạy teacher?
- Trong trace 50%, predictor nhận bao nhiêu rows và output bao nhiêu rows? Mask token width bao nhiêu?
- Scalar toy: nếu eta=0, tau=0,9, teacher có vẫn đổi không? Tính teacher mới.
- Code-like target standardization có áp lên prediction trong pinned Loss.forward không?
- “Có norm_mse trong HF config ⇒ checkpoint local đã train đúng config đó.” Thiếu provenance nào?
- Native và downstream cùng N=128. Nêu hai thông tin chưa giữ giống.
Đáp án
- Full teacher self-attention contextualizes target với cả clip; crop-only teacher không thấy cùng context, nên output target khác.
- 64 context + 64 target =128 rows hidden384; output64×768. Learned mask token hidden384.
- Online a vẫn2, teacher từ3,5 về ; EMA có thể đổi dù gradient update không đổi online. Điều đó khác tau=1 freeze teacher.
- Không; criterion chỉ center/standardize target, rồi L2-normalizes both. Encoder/predictor còn có architectural LN riêng.
- Exact checkpoint identity/hash, training config saved with artifact, commit/log/run provenance và overrides. Later card không chứng minh historical run.
- Grid orientation/time–mel partition và temporal support/hop; còn frequency support/front-end khác. Equal token count là shape fact hạn chế.
Đào sâu tự chọn#
Tự trace nhiều M blocks I-JEPA: target blocks có thể overlap nhau, context loại union targets; predictor một block per call (vectorization có thể batch các calls). Loss weighting của patch xuất hiện trong hai blocks khác một loss trên union. Đọc normalization/reduction/EMA endpoints từ source đã pin trước khi viết reproduction recipe. Không thay những điều unknown bằng scalar toy.