# JEPA — đối chiếu paper, code và checkpoint

**Ngày đọc source:** 11/10/2026. Các GitHub links dưới đây ghim commit bất biến. Line anchors được đọc từ raw file ở commit đó; dùng cả tên function/config key để định vị vì code renderer có thể hiện số dòng lệch.

## 1. I-JEPA: paper formula so với official implementation

Paper: [Assran et al., arXiv:2301.08243v3](https://arxiv.org/html/2301.08243). Code: [facebookresearch/ijepa @ 52c1ae95d05f743e000e8f10a1f3a79b10cff048](https://github.com/facebookresearch/ijepa/tree/52c1ae95d05f743e000e8f10a1f3a79b10cff048), commit 13/06/2023.

| Chủ đề | Paper | Code evidence tại pinned revision | Cách ghi đúng |
|---|---|---|---|
| Input/masking | §3: ảnh thành patches; 4 target blocks, scale .15–.20, AR .75–1.5; context scale .85–1, AR=1; overlap bị gỡ khỏi context. Appendix A.1 có recipe masking. | [MaskCollator init và multiblock config](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/src/masks/multiblock.py#L20-L46), [MaskCollator call](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/src/masks/multiblock.py#L112-L169), [V-H/14 config](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/configs/in1k_vith14_ep300.yaml#L18-L33) | Config ví dụ có patch 14, lưới 16×16, 1 context mask, 4 prediction masks. Geometry và input không giống random patch ratio của Audio-JEPA. |
| Target contextualization | §3 Targets nói target là output của target encoder trên full image; chỉ output sau encoder mới bị mask. Context branch chỉ xử lý context. | [train_step.forward_target](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/src/train.py#L295-L303): target_encoder(imgs) không truyền mask; LayerNorm từng token; rồi apply_masks. [VisionTransformer.forward](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/src/models/vision_transformer.py#L401-L425) chỉ áp mask khi masks được truyền. | Teacher thấy cả ảnh. Target là contextualized features của các vùng output được chọn, không phải output patch-local tính khi chỉ nhìn vùng target. |
| Predictor/vị trí | §3 Prediction dùng shared learned mask vector cộng positional embedding; dự đoán từng target block. | [VisionTransformerPredictor init/forward](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/src/models/vision_transformer.py#L220-L326): predictor projection, learned mask token, fixed 2D sin-cos position embedding, context+target concat, predictor blocks, norm, projection. | Vị trí là điều kiện dự đoán; không bỏ qua target position khi mô tả JEPA. |
| Loss normalization | §3 Loss ghi average squared-L2 trên target patches. | [train_step.loss_fn](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/src/train.py#L310-L313) gọi F.smooth_l1_loss(z,h), rồi AllReduce. Trong [forward_target](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/src/train.py#L295-L303), h qua F.layer_norm(h,(h.size(-1),)) trước khi gather. | Code khác paper: Smooth-L1/Huber trên target đã LayerNorm; mặc định PyTorch reduction là mean trên tensor rồi AllReduce. Nói rõ “paper viết squared L2; released train code dùng Smooth-L1 sau target feature LayerNorm”. |
| Gradient/EMA timing | §3: gradient tối ưu predictor + context encoder; target weights EMA. Appendix mô tả khởi tạo target giống context. | [train_step optimizer và momentum update](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/src/train.py#L321-L336): backward/optimizer step trước; sau đó torch.no_grad lấy momentum và cập nhật target params. | Ở code này EMA diễn ra sau optimizer update trong train_step. Không suy schedule của Audio-JEPA từ I-JEPA. |
| Optimization recipe | Appendix A.1: AdamW; global batch 2048; LR 1e-4→1e-3 trong 15 epochs rồi cosine tới 1e-6; WD .04→.4; EMA .996→1. | [init_opt](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/src/helper.py#L107-L156), [V-H/14 config](https://github.com/facebookresearch/ijepa/blob/52c1ae95d05f743e000e8f10a1f3a79b10cff048/configs/in1k_vith14_ep300.yaml#L41-L53) | Config cụ thể ghi start LR 2e-4, warmup 40 epochs, peak .001; khác Appendix paper ghi start 1e-4, warmup 15. WD endpoints và EMA endpoints khớp. Đây là chênh lệch paper/config, chưa kết luận recipe checkpoint cụ thể. |

## 2. Audio-JEPA: formulation v2 so với code release

Paper: [Tuncay et al., Audio-JEPA, arXiv:2507.02915v2](https://arxiv.org/abs/2507.02915v2), submitted 25/06/2025, last revised 26/09/2026, ICME 2025. Method evidence: PDF §§III–IV, pp. 3–4; Eq. 1–3 and Table I. Official code: [LudovicTuncay/Audio-JEPA @ ddd97ee9b88c00572f59bf2a00eb42120e6a4dea](https://github.com/LudovicTuncay/Audio-JEPA/tree/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea), commit 16/07/2026.

| Thành phần | Source-level evidence và function/config | Kết quả đọc được |
|---|---|---|
| Frontend | [MelSpecTransform.forward](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/data/components/mel_spec.py#L44-L69), [AudioSet data config](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/configs/data/audioset.yaml#L17-L30) | 10 giây, 32 kHz, 128 mel bins, 256 time bins. Code trừ mean waveform; gọi Kaldi fbank với Hanning/log filterbanks, low 20 Hz, high Nyquist; pad/truncate thành 256 time rows. |
| Grid/patch embedding | [train config wiring](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/train.py#L63-L96), [data config](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/configs/data/audioset.yaml#L17-L30), [PatchEmbed](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/components/vision_transformer.py#L45-L61) | input_size=(time=256, mel=128), patch=(16,16): 16 time rows ×8 mel columns =128 tokens. train.py calculates H/W and injects encoder/predictor dimensions; YAML omission of those runtime dimensions is intentional wiring, not a broken config. |
| Mask | [random block mask config](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/configs/masks/random_block.yaml#L1-L7), [MaskCollator call](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/masks/components/random_block.py#L43-L63) | Reads ratio range (.4,.6); samples one ratio uniformly on each collator call/batch, keeps floor(N·(1−ratio)), then draws a fresh token permutation per example. This is behavior of this code, not a JEPA-wide invariant. |
| Encoder and target | [JEPAModule.model_step](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/jepa_module.py#L82-L103), [VisionTransformer.forward](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/components/vision_transformer.py#L194-L287) | Context branch calls encoder with context_masks. Under torch.no_grad, target encoder receives spectrograms without masks and then prediction_masks gather target outputs. Teacher self-attention sees all input patches before positions are selected. |
| Predictor positions/shape | [VisionTransformerPredictor init](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/components/vision_transformer.py#L92-L148), [forward](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/components/vision_transformer.py#L149-L191), [model YAML](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/configs/model/jepa.yaml#L3-L19) | 12-layer, 768-dim encoder/target; learned mask token and fixed 2D sin-cos positions; predictor input 768→384, 6 transformer blocks, norm and projection 384→768. Context and target position vectors are selected separately. |
| Gradient paths | [JEPAModule.model_step](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/jepa_module.py#L87-L103), [configure_optimizers](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/jepa_module.py#L329-L385) | Loss graph includes context encoder and predictor. Target forward is under no-grad and target encoder is not in optimizer groups; a separate EMA callback mutates its parameters. |
| Paper loss vs code loss | Paper §III-B Eq. 2. [norm_mse_loss](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/components/loss.py#L6-L21), [Loss.forward](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/components/loss.py#L24-L50), [criterion config](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/configs/model/jepa.yaml#L21-L25) | Paper writes average squared L2 over masked vectors. Code config chooses norm_mse and norm_pix_loss=true. Loss.forward standardizes each target token over embedding coordinates using its target mean/variance + epsilon; norm_mse_loss L2-normalizes both vectors and computes per-token 2−2·dot, then mean. This is not the same scaling/normalization as raw squared L2. |
| Optimizer/LR | [optimizer and scheduler YAML](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/configs/model/jepa.yaml#L26-L41), [WarmupCosineScheduler init/get_lr](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/optimizers/warmup_cosine.py#L18-L54), [configure_optimizers](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/jepa_module.py#L358-L385) | AdamW YAML says lr=.0003, β=(.9,.95), WD=.05; LR scheduler says warmup 1,000 steps, start 1e-6, ref_lr=1e-3, final 0. Scheduler initializes group LR from start_lr and thereafter returns ref_lr schedule, so scheduled peak is 1e-3, not optimizer YAML lr .0003. Paper v2 states 1e-6→3e-4. |
| Weight decay | [optimizer + WD scheduler YAML](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/configs/model/jepa.yaml#L26-L46), [CosineWDScheduler get_wd/step](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/optimizers/cosine_wd.py#L16-L63), [configure_optimizers param groups](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/models/jepa_module.py#L336-L385) | AdamW field says .05, but WD scheduler has ref_wd=final_wd=1e-6 and sets every non-excluded group to that constant. Biases and 1-D parameters are excluded/zeroed. Thus current wired schedule intends effective WD 1e-6 on regular groups. Paper v2 says .05. This is a paper/config discrepancy, not checkpoint provenance. |
| EMA defaults/update | [MAWeightUpdate init](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/callbacks/MA_weight_update_callback.py#L24-L45), [on_train_batch_end](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/callbacks/MA_weight_update_callback.py#L47-L70), [update_weights](https://github.com/LudovicTuncay/Audio-JEPA/blob/ddd97ee9b88c00572f59bf2a00eb42120e6a4dea/src/callbacks/MA_weight_update_callback.py#L82-L92) | Defaults: initial_tau=.996, final_tau=1, method cos. Hook updates teacher at train-batch end using current tau, then computes next tau. Cosine denominator is len(train_dataloader)×max_epochs; numerator uses Lightning global_step. |

## 3. Audio-JEPA: EMA clock caveat

The public callback is batch-hooked, but the schedule reads global_step. In Lightning, global_step normally counts optimizer steps, while on_train_batch_end is dispatched per training batch. If gradient accumulation is enabled, the callback source can therefore update EMA on batches where student parameters did not just change, and tau may repeat because global_step did not advance. This is a **source-level conditional risk**, not evidence that the authors trained with accumulation.

The callback horizon, len(train_dataloader)×max_epochs, does not use estimated_stepping_batches or directly account for max_steps, limit_train_batches, early stopping, or accumulation. Meanwhile, LR/WD schedulers are constructed with trainer.estimated_stepping_batches in configure_optimizers. Unless training uses the simple one-batch/one-step, full-epoch regime, EMA and optimizer schedules can have different clocks. No run log or trainer launch config was used here to determine whether that happened in the released run.

## 4. Checkpoint provenance: what can and cannot be claimed

The public [Hugging Face Audio-JEPA model](https://huggingface.co/ltuncay/Audio-JEPA) lists JEPA.ckpt, README.md, config.json and an inference example. The [checkpoint file commit](https://huggingface.co/ltuncay/Audio-JEPA/commit/d430e4d32d27d22f1f0b1b5853711605129539ff) is dated 22/05/2025. The later [README/config commit](https://huggingface.co/ltuncay/Audio-JEPA/commit/c65d33bfdef48ccfada785493f3cc0db7409c06f) is dated 16/07/2026 and adds README/config/inference example. The official code snapshot above is also from 16/07/2026 and paper v2 was revised 26/09/2026.

Therefore, the card config documents the later card's architecture/loss, and the pinned repository shows how that code revision is wired. Neither proves that the May 2025 checkpoint used the July 2026 scheduler values or exact source revision. No checkpoint was downloaded, hashed, loaded, or matched to a training run manifest.

The card JSON has an axis-label inconsistency: it reports 256 time bins ×128 mel bins and 16×16 patches, but labels the grid as 8 time ×16 frequency and says 1.6 temporal positions/s. The stated input/data config and code wiring imply 16 temporal ×8 mel positions (1.6 temporal positions/s). Treat card grid_time/grid_freq labels as swapped; use source dimensions to describe the token grid.

## 5. Safe summary text for the chapter

> I-JEPA and Audio-JEPA provide examples of masked latent prediction with context and target encoders, a position-conditioned predictor, and an EMA-updated target. In the released code snapshots, the target sees the full input under no-gradient evaluation, while the context encoder sees visible tokens. Paper formulations and implementations need separate reporting: I-JEPA's official train code uses target feature LayerNorm plus Smooth-L1; Audio-JEPA's current code config uses target standardization plus normalized MSE, and its scheduler config does not match the paper's reported LR/weight-decay values. Those artifacts do not establish which exact recipe produced the public Audio-JEPA checkpoint.

Do not shorten this to “all JEPAs use EMA and MSE,” “the model's error is spoof probability,” or “the published checkpoint used the currently displayed scheduler.”
