# Bài 6 — Từ JEPA formulation tới một training step

[Bắt đầu](00-BAT-DAU-LOP-03.md) · Trước: [asymmetry](05-SELF-DISTILLATION-ASYMMETRY.md) · Tiếp: [masking](07-MASKING-THIET-KE-NHIEM-VU.md)

## 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](https://arxiv.org/html/2301.08243v1#S2).

Trong **I-JEPA/Audio-JEPA teacher–student recipe đang xét**, context encoder $f_\theta$, predictor $g_\phi$, target encoder $f_{\bar\theta}$. Tập context $C$, target $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.

```mermaid
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\in\mathbb R^{B\times1\times256\times128}$. Native patch 16 time × 16 mel → grid 16 × 8, $N=128$. Width encoder 768; predictor hidden 384, 6 blocks, outputs width 768. [train.py, phần tạo H/W và instantiate model](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/train.py), [VisionTransformer/Predictor](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/components/vision_transformer.py).

Chọn một context mask, một target mask, $r=0,5$ cho toy recipe trace: $|C|=|M|=64$. 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](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/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](https://arxiv.org/pdf/2507.02915v2) viết average squared L2 trên masked patch vectors, target EMA. Released [jepa.yaml](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/configs/model/jepa.yaml) chọn `loss_type: norm_mse`, `norm_pix_loss: true`. [Loss.forward và norm_mse_loss](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/components/loss.py) thực hiện:

$$a_{bj}=\frac{t_{bj}-\mu_{bj}}{\sqrt{v_{bj}+10^{-6}}},\quad
\mu_{bj}=\frac1d\sum_k t_{bjk},$$
$$\tilde t_{bj}=a_{bj}/\|a_{bj}\|,\quad \tilde p_{bj}=\hat t_{bj}/\|\hat 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]$, prediction $p=[-1,1]$. Mean target 2; centered target $[-1,1]$. Sau standardization rồi unit normalization, target direction là $[-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]$, cosine với normalized centered target là $1/\sqrt{10}$, loss $2-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](https://arxiv.org/html/2301.08243v1#S3) viết squared L2; official [train.py commit `52c1ae95…`](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/src/train.py) 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](research/jepa-ho-so-nguon.md).

## 4. Ai nhận gradient và ai được cập nhật?

Một step code-like, bỏ batching/AMP để thấy cơ chế:

~~~text
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=ac$, $\hat t=bh$, $t=\bar a x$, $L=\frac12(\hat t-\operatorname{sg}(t))^2$. Mid-training state $a=2,b=3,\bar a=3,5$ → prediction 6, target 7, loss 0,5.

$$\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,2$. EMA tau 0,9 then $\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](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/callbacks/MA_weight_update_callback.py).

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](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/optimizers/warmup_cosine.py), [CosineWDScheduler](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/optimizers/cosine_wd.py).

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](../06-GIAI-PHAU-PAPER.md).

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 $z_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.

<details>
<summary>Đáp án</summary>

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,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ế.

</details>

## Đà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.
