ajaudio / studyAUDIO-JEPA · RESEARCH NOTES
9 phút đọc · Toàn văn
Mục lục bài · 9 mục

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 fθf_\theta, predictor gϕg_\phi, target encoder fθˉf_{\bar\theta}. Tập context CC, target MM; C∩M=∅C\cap M=\varnothing ở 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 .-> T
Xem 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 .-> T

Solid 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à X∈RB×1×256×128X\in\mathbb R^{B\times1\times256\times128}. Native patch 16 time × 16 mel → grid 16 × 8, N=128N=128. 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, r=0,5r=0,5 cho toy recipe trace: ∣C∣=∣M∣=64|C|=|M|=64. Không claim tất cả batches là 50%; config samples ratio 40–60%.

BướcGiá trị/shapeThứ cần giữ đúng
Patch projectionB × 128 × 768Conv kernel=stride=(16,16), không overlapping patches
Add positions; gather CB × 64 × 768Original grid indices được giữ; không renumber positions thành 1…64
Context blocks + final LNB × 64 × 768Transformer chỉ xử lý visible rows
Predictor projectionB × 64 × 384768→384 learned projection
Append learned mask tokens + target positionsB × 128 × 384Context 64 + target 64; mask token không có raw target patch
Predictor blocks + norm; select prediction rowsB × 64 × 384Outputs ở target suffix trong implementation
Predictor output projectionB × 64 × 768Width trở về target space
Teacher full forward + final LNB × 128 × 768Teacher attention đã dùng full input
Gather target M, no-gradB × 64 × 768Gather đúng cùng order với predictor
Criterion target processing/lossScalarLoss 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:

abj=tbj−μbjvbj+10−6,μbj=1d∑ktbjk,a_{bj}=\frac{t_{bj}-\mu_{bj}}{\sqrt{v_{bj}+10^{-6}}},\quad \mu_{bj}=\frac1d\sum_k t_{bjk},
t~bj=abj/∥abj∥,p~bj=t^bj/∥t^bj∥,\tilde t_{bj}=a_{bj}/\|a_{bj}\|,\quad \tilde p_{bj}=\hat t_{bj}/\|\hat t_{bj}\|,
Lrelease=1B∣M∣∑b,j∈M(2−2p~bj⊤t~bj).L_{release}=\frac1{B|M|}\sum_{b,j\in M}\left(2-2\tilde p_{bj}^{\top}\tilde t_{bj}\right).

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 t=[1,3]t=[1,3], prediction p=[−1,1]p=[-1,1]. Mean target 2; centered target [−1,1][-1,1]. Sau standardization rồi unit normalization, target direction là [−1,1]/2[-1,1]/\sqrt2, trùng normalized p; code-like loss 0. Raw squared L2 giữa p và t là 8. Nếu p=[1,2]p=[1,2], cosine với normalized centered target là 1/101/\sqrt{10}, loss 2−2/10≈1,3675442-2/\sqrt{10}\approx1,367544. 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. h=ach=ac, t^=bh\hat t=bh, t=aˉxt=\bar a x, L=12(t^−sg⁡(t))2L=\frac12(\hat t-\operatorname{sg}(t))^2. Mid-training state a=2,b=3,aˉ=3,5a=2,b=3,\bar a=3,5 → prediction 6, target 7, loss 0,5.

∂aL=(6−7)bc=−3,∂bL=(6−7)ac=−2,∂aˉL=0.\partial_aL=(6-7)bc=-3,\quad \partial_bL=(6-7)ac=-2,\quad \partial_{\bar a}L=0.

SGD eta 0,1 updates simultaneously using old gradients: a′=2,3,b′=3,2a'=2,3,b'=3,2. EMA tau 0,9 then aˉ′=0,9(3,5)+0,1(2,3)=3,38\bar a'=0,9(3,5)+0,1(2,3)=3,38. 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 zB−zSz_B-z_S. Bài 9 tách encoder transfer khỏi error scoring.

Bài tập tăng dần#

  1. Vì sao chọn M sau teacher full forward khác crop M rồi chạy teacher?
  2. Trong trace 50%, predictor nhận bao nhiêu rows và output bao nhiêu rows? Mask token width bao nhiêu?
  3. Scalar toy: nếu eta=0, tau=0,9, teacher có vẫn đổi không? Tính teacher mới.
  4. Code-like target standardization có áp lên prediction trong pinned Loss.forward không?
  5. “Có norm_mse trong HF config ⇒ checkpoint local đã train đúng config đó.” Thiếu provenance nào?
  6. Native và downstream cùng N=128. Nêu hai thông tin chưa giữ giống.
Đáp án
  1. 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.
  2. 64 context + 64 target =128 rows hidden384; output64×768. Learned mask token hidden384.
  3. Online a vẫn2, teacher từ3,5 về 0,9(3,5)+0,1(2)=3,350,9(3,5)+0,1(2)=3,35; EMA có thể đổi dù gradient update không đổi online. Điều đó khác tau=1 freeze teacher.
  4. Không; criterion chỉ center/standardize target, rồi L2-normalizes both. Encoder/predictor còn có architectural LN riêng.
  5. Exact checkpoint identity/hash, training config saved with artifact, commit/log/run provenance và overrides. Later card không chứng minh historical run.
  6. 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.

DỪNG LẠI & TỰ KIỂM TRA

Bạn đã giải thích được cơ chế trong bài?

↓ Bản Markdown nguyên gốcGiữ nguyên nội dung · Công thức, bảng và nguồn đầy đủ.Các chat bàn giao được mở trong Codex.

Gõ từ khóa để tìm bài học và đoạn liên quan.