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

Bài 2 — Tách target, context, loss và gradient#

Bắt đầu · Trước: representation · Tiếp: speech units và contrastive

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

Bạn sẽ đọc một SSL objective như một quy trình tạo supervision, tính bốn loss nhỏ và chỉ ra thứ được khuyến khích/không được bảo đảm. Cần log/softmax/gradient ở xác suất lớp 1. MM là tập vị trí chịu prediction loss, BB là batch size, dd là embedding width; đừng dùng cùng ký hiệu MM cho cả số block và mask positions trong một derivation.

Trực giác: thầy có thể yêu cầu bạn chép lại âm thanh, đoán mã cluster, chọn đoạn đúng giữa các đoạn nhiễu, hoặc dự đoán một vector teacher. Tất cả đều “không có human label”, nhưng câu hỏi và đáp án rất khác. SSL là nguồn supervision; loss là quy tắc thưởng/phạt; JEPA là họ predictive embedding designs. Chúng không nằm trên cùng một trục.

1. Đọc sáu câu trước tên model#

Câu hỏiMột lựa chọn có thể cóVì sao quan trọng
Target được tạo từ đâu?Input, random quantizer, k-means, teacher, current encoderTarget cố định hay có thể trôi cùng student
Target chứa gì?Signal, hard ID, distribution, continuous vector“Continuous input” không nói target type
Student thấy gì?Past, visible patches, masked full-length sequence, second viewXác định thông tin có sẵn
Loss tính trên đâu?Masked tokens, mọi token, pooled view, future offsetsSố denominator và đơn vị supervision
Gradient qua đâu?Encoder, predictor, quantizer; detach targetĐường cập nhật thực tế
Constraint/update gì?Negatives, fixed labels, EMA, normalization, varianceNghiệm hằng có bị phạt hay dynamics đổi gì

Sau đó mới thêm data, architecture, budget và downstream output. I-JEPA §2 phân biệt joint embedding, input-space prediction và predictive embedding; data2vec §3 là ví dụ latent regression có trước Audio-JEPA. Bảng model đầy đủ đặt riêng ở hồ sơ recipe, để bài giảng không biến thành danh sách tên.

2. Input reconstruction: dự đoán số nào?#

Giả sử target patch xj∈Rpx_j\in\mathbb R^p là pp giá trị log-mel đã patchify. Một convention rõ ràng:

Lrec=1B∣M∣p∑b=1B∑j∈M∑k=1p(x^bjk−xbjk)2.L_{rec}=\frac{1}{B|M|p}\sum_{b=1}^B\sum_{j\in M}\sum_{k=1}^p(\hat x_{bjk}-x_{bjk})^2.

Nếu mỗi sample có mask size khác nhau, cần nói mean theo sample trước hay gộp mọi phần tử; hai cách cho weights khác. Loss không có đơn vị vật lý phổ quát: nếu input log-mel đã normalize, đơn vị cũng đổi.

Toy đã giải: một patch có target [1,3][1,3], prediction [2,1][2,1]. Tổng squared error 1+4=51+4=5; mean theo hai scalar là 2,5; squared vector distance là 5. Với L=12∑k(x^k−xk)2L=\frac12\sum_k(\hat x_k-x_k)^2, derivative theo prediction là [1,−2][1,-2]. Với mean hai scalar, derivative cũng [1,−2][1,-2] ở toy p=2p=2, nhưng đây là trùng hợp coefficient 1/p=1/21/p=1/2; không phải quy tắc chung.

Conditional-mean bridge: với fixed target random variable TT có finite second moment và context C=cC=c, đặt μ=E[T∣c]\mu=E[T\mid c]. Khai triển

E[∥T−a∥2∣c]=E[∥T−μ∥2∣c]+∥a−μ∥2.E[\|T-a\|^2\mid c]=E[\|T-\mu\|^2\mid c]+\|a-\mu\|^2.

Cross-term bằng 0 vì E[T−μ∣c]=0E[T-\mu\mid c]=0. Tối ưu output squared loss là a=μa=\mu. Nhưng proof nói về output prediction với fixed distribution, không nói mọi hidden feature của encoder bị xóa, và không áp trực tiếp khi target encoder cũng học. AudioMAE §3 cho một reconstruction recipe cụ thể; bài 4 sẽ phân biệt decoder và encoder.

3. Hard units và soft distributions#

Hard target c∈{1,…,K}c\in\{1,\ldots,K\} là một index, không phải embedding dimension. Classifier trả logits z∈RKz\in\mathbb R^K, pk=ezk/∑lezlp_k=e^{z_k}/\sum_l e^{z_l}; CE là −log⁡pc-\log p_c.

Toy: p=[0,6;0,3;0,1]p=[0,6;0,3;0,1] (dấu chấm phẩy tách phần tử), target class 2. CE =−ln⁡0,3≈1,203973=-\ln0,3\approx1,203973 nats. Derivative theo logits là p−e2=[0,6;−0,7;0,1]p-e_2=[0,6;-0,7;0,1]. Đó là push probability về ID được giao; không chứng minh ID là phoneme chuẩn.

Soft target q=[0,5;0,4;0,1]q=[0,5;0,4;0,1] giữ uncertainty. Soft CE H(q,p)=−∑qkln⁡pk≈0,967260H(q,p)=-\sum q_k\ln p_k\approx0,967260; KL:

DKL(q∥p)=∑qkln⁡(qk/pk)=H(q,p)−H(q)≈0,023912.D_{KL}(q\|p)=\sum q_k\ln(q_k/p_k)=H(q,p)-H(q)\approx0,023912.

Nếu qq được detach/fixed, hai loss có cùng gradient p−q=[0,1;−0,1;0]p-q=[0,1;-0,1;0] theo logits. Chúng có giá trị khác nhau, vì entropy H(q)≈0,943348H(q)\approx0,943348. Nếu cho gradient qua qq, nhận xét “chỉ khác constant” không còn đúng. Soft assignment có thể phản ánh channel thay vì articulation; softness không tự tạo relevance.

Nguồn cho hard-unit recipe: HuBERT §§2.1–2.2; soft CE/KL và derivative ở đây là dẫn xuất biên soạn, không một cấu hình HuBERT.

4. InfoNCE: chọn đúng positive giữa các candidate#

Cho query q∈Rdq\in\mathbb R^d, keys ki∈Rdk_i\in\mathbb R^d, positive index ++, similarity sis_i, temperature τ>0\tau>0. Candidate probability

pi=esi/τ∑lesl/τ,LNCE=−ln⁡p+,∂L∂si=pi−1[i=+]τ.p_i=\frac{e^{s_i/\tau}}{\sum_l e^{s_l/\tau}},\quad L_{NCE}=-\ln p_+,\quad \frac{\partial L}{\partial s_i}=\frac{p_i-\mathbf1[i=+]}{\tau}.

Toy: scores [2,0,0][2,0,0], positive thứ nhất, τ=1\tau=1. p+=e2/(e2+2)≈0,786986p_+=e^2/(e^2+2)\approx0,786986, loss ln⁡(1+2e−2)≈0,239545\ln(1+2e^{-2})\approx0,239545. Gradients theo ba scores là [−0,213014;0,106507;0,106507][-0,213014;0,106507;0,106507]. Giảm temperature làm phân phối sắc hơn tại cùng scores; gradient có cả factor 1/τ1/\tau và xác suất đổi, nên không nói gradient luôn tăng cùng một tỷ lệ.

Nếu mọi key/query cho cùng score, với ba candidates loss là ln⁡3≈1,098612\ln3\approx1,098612, không bằng 0. Điều này cho thấy collapse không đạt giá trị tối ưu khi dữ liệu/năng lực cho phép phân biệt positives. Nó chưa chứng minh mọi parameterization tránh stationary collapsed state hoặc học đúng invariance. CPC §2.2, Eq. 4 dùng contrastive future prediction; wav2vec2 §2.3 dùng masked-context discrimination.

Hai failure cases. Positive cùng recording channel, negatives từ channels khác: query có thể thắng bằng channel. Hai đoạn đều genuine, cùng phoneme và speaker nhưng khác utterance bị coi negative: loss vẫn đẩy chúng xa, dù downstream có thể muốn gộp. “False negative” cần nói notion equivalence của task: instance discrimination chủ ý phân biệt instances; nó không hứa invariance đúng forensics.

5. Continuous feature matching: vector target cũng là một thiết kế#

Teacher trả tj∈Rdt_j\in\mathbb R^d, predictor trả t^j\hat t_j. Với target detach, mean scalar MSE có gradient

L=1B∣M∣d∑b,j,k(t^bjk−sg⁡(tbjk))2,∇t^L=2(t^−t)B∣M∣d.L=\frac1{B|M|d}\sum_{b,j,k}(\hat t_{bjk}-\operatorname{sg}(t_{bjk}))^2,\quad \nabla_{\hat t}L=\frac{2(\hat t-t)}{B|M|d}.

sg⁡\operatorname{sg} giữ nguyên giá trị forward, derivative qua nhánh đó bằng 0. Nếu target có norm lớn gấp 10, raw error có thể lớn gấp 100 khi cả prediction/target scale cùng nhau. So errors giữa checkpoints mà không khóa scale có thể so đơn vị khác nhau.

L2 normalization v~=v/∥v∥\tilde v=v/\|v\| (với norm khác 0) cho

∥p~−t~∥2=2−2cos⁡(p,t).\|\tilde p-\tilde t\|^2=2-2\cos(p,t).

Toy p=[3,4]p=[3,4], t=[6,8]t=[6,8]: raw squared distance 25, normalized distance 0. t=[0,5]t=[0,5] cho cosine 4/54/5, normalized distance 0,4. Layer normalization/standardization khác unit-norm normalization: một bên center/scale theo feature axis, bên kia giữ hướng và đặt norm 1. Bài 6 sẽ đọc cả hai bước trong released Audio-JEPA config/code.

Smooth L1 với residual r=t^−tr=\hat t-t và threshold β>0\beta>0:

ℓβ(r)={r2/(2β),∣r∣≤β,∣r∣−β/2,∣r∣>β.\ell_\beta(r)=\begin{cases}r^2/(2\beta),&|r|\le\beta,\\|r|-\beta/2,&|r|>\beta.\end{cases}

Derivative là r/βr/\beta ở vùng nhỏ, sign(r)(r) ở vùng lớn. Với β=1\beta=1, r=0,5r=0,5 cho loss 0,125, gradient 0,5; r=2r=2 cho loss 1,5, gradient 1. Đây là convention Smooth L1 dùng trong data2vec §3.4; không đồng nhất nó với raw MSE hoặc mọi định nghĩa Huber có hệ số khác.

6. Từ loss tới forensic claim cần thêm một cầu#

Một target learned có thể match bằng vector hằng nếu recipe cho phép. Fixed informative targets không thể tự trôi, nhưng context có thể thiếu thông tin để predict chúng: khi đó student tối ưu có thể trả conditional average. Negatives tạo discrimination, nhưng discrimination của channel chưa là detection của synthesis. Chống collapse và forensic relevance là hai câu hỏi riêng.

Paper SOICT dùng weighted CE trên bona fide/spoof sau khi lấy encoder pretrained. Latent loss ở pretraining và CE ở downstream là hai giai đoạn khác nhau; không suy CE thấp là JEPA prediction tốt. Chương 06.

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

  1. Masked spectrogram input nhưng target là k-means ID: target type và loss family là gì?
  2. Tính CE khi pc=0,25p_c=0,25. Nếu prediction bằng soft target, KL bằng gì và soft CE có bằng 0 không?
  3. Ba candidate scores đều bằng nhau. InfoNCE bằng bao nhiêu? Nó có xác lập học cue đúng task không?
  4. p=[1,0]p=[1,0], t=[2,0]t=[2,0]: raw squared distance và normalized distance? Biến nào bị normalized loss bỏ?
  5. Hai paper đều nói “MSE”, một paper sum theo d=768d=768, một paper mean theo d. Với cùng residuals loss/gradient khác theo factor nào? Có thể so learning rates trực tiếp không?
Đáp án
  1. Hard discrete target; thường CE nếu predict ID. Masking là student-view/task design, không quyết định target là input reconstruction.
  2. −ln⁡0,25=ln⁡4≈1,386294-\ln0,25=\ln4\approx1,386294. Khi p=qp=q, KL 0, soft CE bằng H(q)H(q); không bằng 0 trừ q là point mass.
  3. ln⁡3\ln3. Không: sampling/task alignment và parameter dynamics vẫn chưa được kiểm.
  4. Raw 1, normalized 0; độ dài bị bỏ khỏi so sánh loss. Điều này không tự chứng minh raw encoder norm không có cue.
  5. Factor 768 nếu các axes/positions khác giống nhau. Update còn phụ thuộc optimizer, LR, scaling và normalization; không so riêng LR để kết luận bước học tương đương.

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

Suy ra softmax CE gradient từ L=−zc+log⁡∑ezkL=-z_c+\log\sum e^{z_k}. Khi teacher cùng học, phân biệt derivative mà implementation chọn với derivative của hàm số nếu recompute mọi nhánh; bài 5–6 sẽ dùng cached target để kiểm finite differences đúng computational path. Không lấy một InfoNCE MI bound làm phép đo forensic information của checkpoint mà chưa kiểm assumptions của bound và sampling.

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.