skipLink.label

Quest 17 - DPO Loss Implementation

Quest 17: DPO Loss Implementation

hard 30 minutes

🎯 Learning Objectives

  • ✅ Understand DPO (Direct Preference Optimization) as a simplified alternative to RLHF
  • ✅ Implement the DPO loss function using log-space arithmetic
  • ✅ Apply numerical stability techniques to prevent overflow/underflow
  • ✅ Recognize when policy and reference models produce identical outputs

📖 Concept: DPO Loss

DPO (Direct Preference Optimization) เป็นวิธี alignment ที่ปฏิวัติวงการ — มันกำจัด reward model ที่ซับซ้อนออก แล้ว optimize โดยตรงบน preference data

แทนที่จะ train reward model แล้วใช้ RL (RLHF), DPO คำนวณ loss โดยตรงจาก chosen vs rejected responses:

L = -log(σ(β × (log π(y_w|x)/π_ref(y_w|x) - log π(y_l|x)/π_ref(y_l|x))))

ที่ไหน:

  • π = policy model (model ที่กำลัง train)
  • π_ref = reference model (model เดิม)
  • β = temperature parameter (ควบคุมความแรงของการ optimize)
  • y_w = chosen response, y_l = rejected response

คิดเหมือน martial arts: DPO เปรียบเหมือนการตัดสินใจทันทีโดยไม่ต้องคิดผ่าน reward model — เลือก technique ที่ดีกว่าโดยตรง


⚙️ How It Works

The DPO Loss Formula

1. Compute log-probability ratios:
ratio_chosen = log π(y_w|x) - log π_ref(y_w|x)
ratio_rejected = log π(y_l|x) - log π_ref(y_l|x)
2. Scale by beta:
scaled_diff = β × (ratio_chosen - ratio_rejected)
3. Apply sigmoid and log:
L = -log(σ(scaled_diff))

Why Log-Space Arithmetic?

Direct computation of probability ratios can overflow:

// DANGEROUS: may overflow
const prob = Math.exp(logProb); // exp(-1000) = 0
const ratio = prob_chosen / prob_rejected; // 0 / 0 = NaN
// SAFE: log-space
const logRatio = logProb_chosen - logProb_rejected; // always finite

The Sigmoid Function

function sigmoid(x) {
return 1 / (1 + Math.exp(-x));
}

When policy = reference, the loss equals log(2) ≈ 0.693 — this is the “no improvement” baseline.


💡 Example: Implementing DPO Loss

Step 1: Implement the core formula

function dpoLoss(policyLogps, refLogps, beta) {
// Log-space ratios
const ratioChosen = policyLogps.chosen - refLogps.chosen;
const ratioRejected = policyLogps.rejected - refLogps.rejected;
// Scaled difference
const scaledDiff = beta * (ratioChosen - ratioRejected);
// DPO loss: -log(σ(scaledDiff))
return -Math.log(1 / (1 + Math.exp(-scaledDiff)));
}

Step 2: Verify with known values

// When policy = reference → loss = log(2)
const loss = dpoLoss(
{ chosen: -1, rejected: -2 },
{ chosen: -1, rejected: -2 },
1.0
);
// loss ≈ 0.693 (log 2)

Step 3: Run the tests

Terminal window
node test.js
# Test: returns positive number
# Test: policy=reference gives ~log(2)
# Test: higher beta changes loss
# Test: handles large log-probs without overflow

⚠️ Common Mistakes

Mistake 1: Using exp() instead of log-space

Math.exp(logProb) for large negative logProbs gives 0 → NaN → Always work in log-space: subtract log-probs instead of dividing probs.

Mistake 2: Wrong formula sign

“L = log(σ(…))” instead of “L = -log(σ(…))” → The loss is NEGATIVE log-sigmoid. Forgetting the minus sign gives wrong gradients.

Mistake 3: Ignoring the beta parameter

“I’ll just hardcode beta = 1.0” → Beta controls how strongly you push away from reference. Different tasks need different beta values.

Mistake 4: Not returning a scalar

Returning an object instead of a single number → DPO loss must be a scalar number for backpropagation.


📝 Knowledge Check

📝 Knowledge Check

Q1:What is the DPO loss formula?

Q2:Why should you use log-space arithmetic instead of computing probabilities directly?

Q3:When the policy model matches the reference model exactly, what is the DPO loss?


🏋️ Quest: DPO Loss Implementation

Now it’s time to implement the DPO loss function!

  1. Download ไฟล์เริ่มต้นของ quest:

    Terminal window
    npx bluebeltdojo download quest-17-dpo-loss
    cd quest-17-dpo-loss
  2. เปิด problem.js ใน editor ของคุณพร้อมความช่วยเหลือของ AI

  3. Implement dpoLoss(policyLogps, refLogps, beta) — คำนวณ DPO loss

  4. สำคัญ: ใช้ log-space arithmetic เพื่อป้องกัน numerical overflow

  5. ตรวจสอบ solution ของคุณ:

    Terminal window
    node test.js
  6. When all tests pass, submit your solution:

    Terminal window
    npx bluebeltdojo submit

💡 Tip: DPO คือ algorithm ที่ทำให้ alignment ง่ายขึ้น — เข้าใจ loss function แล้วจะเข้าใจว่า model เรียนรู้จาก preferences อย่างไร


คำใบ้

  • อ่าน instructions ใน problem.js อย่างละเอียด
  • ใช้ log-space arithmetic เสมอ — อย่าใช้ exp() กับ log-probs
  • ตรวจสอบ: policy=reference ต้องได้ loss ≈ log(2) ≈ 0.693
  • ถ้าติดขัด ลองอ่าน “Common Mistakes” อีกครั้ง — อย่าดู solution โดยตรง