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

Bài 5 — Self-distillation, stop-gradient và constraints#

Bắt đầu · Trước: continuous targets · Tiếp: một training step JEPA

Mục tiêu và tiền đề#

Bạn sẽ trace một term có stop-gradient, giải thích encoder shared vẫn học trong SimSiam, tính EMA và variance/covariance penalties. Cần chain rule/gradient từ lớp 1. Trực giác: đáp án teacher thay chậm, chặn đạo hàm teacher, thêm predictor, giữ spread của batch là bốn can thiệp khác nhau. Chúng có thể phối hợp nhưng không thay thế nhau.

1. BYOL và SimSiam: học match hai views#

BYOL online path gồm encoder fθf_\theta, projector gθg_\theta, predictor qϕq_\phi; target path gồm encoder/projector EMA fθˉ,gθˉf_{\bar\theta},g_{\bar\theta}, không predictor. Hai augmentation views được dùng hai chiều; một term match normalized prediction với detach target projection, tương đương 2−2cos⁡2-2\cos. Optimizer cập nhật online/predictor, EMA cập nhật target. Representation downstream lấy encoder, không mặc định lấy projection. BYOL §3.1, Eq. 1–3.

SimSiam dùng encoder/projector shared weights cho hai views, predictor và symmetric negative cosine:

L=12D(p1,sg⁡(z2))+12D(p2,sg⁡(z1)),D(p,z)=−cos⁡(p,z).L=\tfrac12D(p_1,\operatorname{sg}(z_2))+\tfrac12D(p_2,\operatorname{sg}(z_1)),\quad D(p,z)=-\cos(p,z).

Không EMA teacher. Ở term đầu, target side z2z_2 không gradient; ở term sau, view 2 prediction path vẫn gradient. Shared weights tổng hợp gradients của prediction paths. SimSiam §3, Algorithm 1 và Eq. 4.

Đây là view matching, khác context→masked-region prediction ở bài 6. “Non-contrastive” chỉ nói không dùng contrastive negatives ở loss ấy; không cho biết target continuous/discrete, positions hay teacher.

2. Stop-gradient: giá trị đúng nhưng đạo hàm bị chọn khác#

Định nghĩa sg⁡(u)=u\operatorname{sg}(u)=u trong forward, ∂sg⁡(u)/∂u=0\partial\operatorname{sg}(u)/\partial u=0 trong backward. Một toy scalar raw matching, không loss cosine SimSiam thật:

L(a)=12(ax1−sg⁡(ax2))2,x1=1, x2=2, a=1.L(a)=\tfrac12(a x_1-\operatorname{sg}(a x_2))^2,\quad x_1=1,\ x_2=2,\ a=1.

Forward prediction 1, target 2, residual −1, loss 0,5. Trong step này cache target t=2t=2; derivative qua prediction là

∂L∂a=(ax1−t)x1=−1.\frac{\partial L}{\partial a}=(a x_1-t)x_1=-1.

SGD learning rate 0,1 cho a′=1,1a'=1,1. Nếu bỏ detach và differentiating cả hai branches, derivative là (ax1−ax2)(x1−x2)=+1(ax_1-ax_2)(x_1-x_2)=+1, cho a′=0,9a'=0,9. Cùng forward value, khác gradient rule.

Giới hạn đáng hiểu: recompute cả target bằng a mới thì full forward loss a2/2a^2/2 tăng ở update detach. Điều đó không mâu thuẫn optimizer: nó giảm loss đối với target đã cache trong step, không tối ưu đồng thời mọi parameter dependence của forward expression. Finite difference để kiểm backward detach phải giữ target fixed; nếu recompute target khi perturb a, bạn đang kiểm derivative của hàm khác.

Toy trên cũng cho thấy stop-gradient tự nó không đảm bảo ổn định: có thể làm a tăng theo rule ấy. BYOL/SimSiam recipe có vector normalization, predictor, architecture/optimizer/data; không lấy một scalar toy để dự báo chúng collapse hay explode.

3. EMA: weights chậm, không phải gradient target#

Sau optimizer step, một convention:

θˉt+1=τtθˉt+(1−τt)θt+1,0≤τt≤1.\bar\theta_{t+1}=\tau_t\bar\theta_t+(1-\tau_t)\theta_{t+1},\quad 0\le\tau_t\le1.

Numerical toy: old teacher scalar 10, online sau optimizer 14, τ=0,9\tau=0,9 → teacher mới 0,9(10)+0,1(14)=10,40,9(10)+0,1(14)=10,4. Nếu teacher nhận gradient descent trực tiếp, update sẽ phụ thuộc residual/Jacobian/LR; EMA chỉ phụ thuộc weights và decay. Predictor không được EMA vào target encoder nếu architecture target chỉ copy encoder.

Với fixed tau:

θˉt=τtθˉ0+(1−τ)∑i=1tτt−iθi.\bar\theta_t=\tau^t\bar\theta_0+(1-\tau)\sum_{i=1}^t\tau^{t-i}\theta_i.

Hệ số initial state còn τt\tau^t. Half-life của hệ số cũ là k1/2=ln⁡(0,5)/ln⁡τk_{1/2}=\ln(0,5)/\ln\tau cho 0<τ<10<\tau<1. τ=0,9\tau=0,9 → 6,5788 steps; 0,996 → 172,9400 steps. Đơn vị là updates theo recipe, không seconds/epochs. Với schedule, hệ số trở thành tích ∏τi\prod\tau_i; không dùng fixed-tau half-life như hằng của cả run. Teacher gradients bị chặn trong step nhưng weights của nó phụ thuộc lịch sử online updates; hai khái niệm không mâu thuẫn. BYOL §3.1 là nguồn update rule, phép khai triển/half-life ở đây do biên soạn.

4. Vì sao “có EMA + predictor” chưa là proof?#

Một unconstrained model có fθ(x)=vf_\theta(x)=v và teacher fθˉ(x)=vf_{\bar\theta}(x)=v cho mọi x; predictor trả v. Raw matching loss 0. Với unit-norm loss, chọn fixed nonzero v rồi match normalized v cũng cho 0. Nếu online/teacher weights ở fixed point giống nhau, EMA giữ nguyên. Stop-gradient không tạo penalty mới lên nghiệm hằng.

Điều này chỉ chứng minh nghiệm hằng có thể tồn tại trong setup đã giả định. Nó không chứng minh optimizer của BYOL/I-JEPA thực sẽ hội tụ về nghiệm ấy. Paper BYOL §3.2 phân biệt undesirable equilibria với empirical dynamics; SimSiam §4 ablates stop-gradient trong recipe của họ. Không nâng empirical ablation sang universal theorem.

Predictor tách representation task khỏi prediction mapping; mismatch capacity hoặc update speed có thể đổi dynamics. “Thêm predictor” không tự giới hạn nó dùng positions/shortcut, cũng không làm target informative. Kiến trúc và constraints cần đọc cùng data/views.

5. VICReg: phạt ba điều khác nhau#

Cho projected batches Z,Z′∈Rn×dZ,Z'\in\mathbb R^{n\times d}, mỗi row là một sample. Mean feature μj\mu_j, sample covariance C(Z)=Zc⊤Zc/(n−1)C(Z)=Z_c^\top Z_c/(n-1), n>1n>1:

s(Z,Z′)=1n∑i∥zi−zi′∥2,s(Z,Z')=\frac1n\sum_i\|z_i-z'_i\|^2,
v(Z)=1d∑jmax⁡(0,γ−Var⁡(Z:j)+ϵ),v(Z)=\frac1d\sum_j\max(0,\gamma-\sqrt{\operatorname{Var}(Z_{:j})+\epsilon}),
c(Z)=1d∑j≠kC(Z)jk2.c(Z)=\frac1d\sum_{j\ne k}C(Z)_{jk}^2.

Loss λs+μ[v(Z)+v(Z′)]+ν[c(Z)+c(Z′)]\lambda s+\mu[v(Z)+v(Z')]+\nu[c(Z)+c(Z')]. Invariance aligns paired views; variance hinge discourages tiny spread; covariance reduces duplicated variation. Gradient chạy hai branches/projector, không cần EMA/stop-gradient. VICReg §4.1, Eq. 1–6.

Toy đã giải: Z=Z′=[(−1,−1),(+1,+1)]Z=Z'=[(-1,-1),(+1,+1)], dùng sample variance, γ=1\gamma=1, ϵ=0\epsilon=0 chỉ cho số học. s=0s=0, variance từng chiều 2 → v=0v=0. Covariance C=[2222]C=\begin{bmatrix}2&2\\2&2\end{bmatrix} → c=(22+22)/2=4c=(2^2+2^2)/2=4. Spread có nhưng hai chiều trùng nhau; covariance penalty nhận ra redundancy.

Nếu mọi rows là [7,7][7,7], s=0, c=0, v=1 mỗi branch khi epsilon 0; loss còn 2μ2\mu. Với epsilon thật, v=max⁡(0,γ−ϵ)\max(0,\gamma-\sqrt\epsilon) chứ không đúng 1. Covariance penalty alone không phạt constant rows. Hơn nữa variance penalty dương không đồng nghĩa exact constant point có nonzero gradient: variance derivative bằng 0 tại exact constants với smooth epsilon. “Nghiệm bị phạt” khác “mọi dynamics thoát được điểm đó”.

High variance/full rank vẫn có thể là channel hoặc random noise; batch stats ở projector cũng chưa chứng minh encoder giữ cue task. Bài 8 sẽ dùng phản ví dụ ở mức token/utterance.

6. Barlow Twins: cross-correlation identity, không cùng loss VICReg#

Với hai batches đã center theo sample, cross-correlation:

Rjk=∑iZijZik′∑iZij2∑iZ′ik2,R_{jk}=\frac{\sum_i Z_{ij}Z'_{ik}}{\sqrt{\sum_i Z_{ij}^2}\sqrt{\sum_i {Z'}_{ik}^2}},
LBT=∑j(1−Rjj)2+λ∑j≠kRjk2.L_{BT}=\sum_j(1-R_{jj})^2+\lambda\sum_{j\ne k}R_{jk}^2.

Hai branches shared, gradients qua cả hai; không target EMA/stop-gradient/predictor trong core recipe. Barlow Twins §2.1, Eq. 1–2.

Toy Z=Z′=[(−1,−1),(1,1)]Z=Z'=[(-1,-1),(1,1)] cho R toàn 1: diagonal loss 0, off-diagonal 2λ2\lambda. Nếu constant dimensions, denominator formula là 0; implementation epsilon/normalization cần đọc trước khi mô tả value/gradient. Không gán variance hinge của VICReg cho Barlow Twins. Zero correlation cũng chưa là statistical independence cho nonlinear distributions.

7. Neo vào audio forensics và bài tập#

View augmentation định nghĩa điều cần match. Nếu nó xóa cue synthesis hợp lệ cho task, match-view objective có thể giảm nhạy với cue đó. Đây là suy luận từ task information, không một empirical conclusion về Audio-JEPA. Audio-JEPA main masking prediction không giống BYOL augmentation recipe; EMA chung không làm hai methods đồng nhất.

  1. Vì sao encoder shared trong SimSiam vẫn nhận gradient từ view 2?
  2. Tính target EMA nếu old 4, online new 8, tau 0,75. Predictor params có được copy nếu teacher chỉ là encoder không?
  3. Với stop-gradient toy, finite difference cần giữ gì cố định?
  4. Nếu thêm covariance penalty vào matching-only objective, constant solution có bị phạt không? Variance term làm khác gì?
  5. High-rank noise và high-rank forensic cue đều thỏa spread: constraint cho biết điều gì còn thiếu?
Đáp án
  1. View 2 là prediction branch ở symmetric term thứ hai; detach chỉ chặn một computational path, không freeze shared encoder.
  2. Teacher 5; predictor không EMA vào encoder nếu recipe không có correspondence ấy.
  3. Cache target giá trị tại step. Recompute target dưới perturbed params là kiểm đạo hàm của objective khác.
  4. Constant centered covariance bằng zero nên c=0. Variance hinge phạt spread thấp, nhưng exact-state gradient/dynamics cần xét riêng.
  5. Thiếu relevance, accessibility, stability và evidence head reliance. Chống một geometric degeneracy không xác lập usefulness.

Đào sâu tự chọn#

Chứng minh normalized cosine gradient ∇pcos⁡(p,t)=t~/∥p∥−cos⁡(p,t)p/∥p∥2\nabla_p\cos(p,t)=\tilde t/\|p\|-\cos(p,t)p/\|p\|^2 với norms nonzero. Gradient tiếp tuyến với sphere, cho thấy normalized loss không trực tiếp thưởng tăng norm như raw dot product. EMA thực có parameters/buffers, hooks và accumulation: cần đọc code trước khi xác định “mỗi step”, bài 6 làm điều đó cho released recipe.

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.