LongStraw ประกาศว่าเทคนิค branch-replay ของตนสามารถประมวลผลตำแหน่งโทเคน (token positions) ได้ถึง 2.1 ล้านตำแหน่งสำหรับการทำ post-training แบบ reinforcement-learning (RL) โดยใช้ GPU H20 เพียงแปดตัว ซึ่งช่วยลดค่าใช้จ่ายด้านฮาร์ดแวร์ลงถึงหนึ่งระดับ (an order of magnitude) คำกล่าวอ้างนี้มีความสำคัญเนื่องจากการฝึกโมเดลที่มีบริบทระยะยาว (long-context models) ตามปกติแล้วต้องใช้ GPU ประสิทธิภาพสูงจำนวนหลายสิบตัว ซึ่งเป็นอุปสรรคสำหรับห้องปฏิบัติการวิจัยและสตาร์ทอัพส่วนใหญ่
ทำไม long-context RL ถึงมีราคาแพง
การทำ fine-tuning โมเดลภาษาขนาดใหญ่โดยใช้ RL มักจะมีการรัน rollout เพื่อสร้างคำตอบทางเลือกที่หลากหลายสำหรับ prompt เดียวกัน การรัน rollout แต่ละครั้งต้องผ่านกระบวนการ back-propagation ดังนั้นต้นทุนการคำนวณจึงเพิ่มขึ้นตามจำนวนตำแหน่งโทเคนทั้งหมดที่ประมวลผล ในปัจจุบัน pipeline ที่ตั้งเป้าบริบทระดับหนึ่งล้านโทเคนมักต้องใช้ GPU ตั้งแต่ 64 ถึง 128 ตัวเพื่อให้เสร็จสิ้นภายในระยะเวลาที่เหมาะสม ต้นทุนของฮาร์ดแวร์ดังกล่าว ประกอบกับค่าไฟฟ้าและระบบหล่อเย็นที่จำเป็น เป็นข้อจำกัดในการขยายความยาวของบริบท (context length) ของผู้ใช้งาน
เทคนิค branch replay ช่วยลดภาระงานได้อย่างไร
แนวทางของ LongStraw ขึ้นอยู่กับการสังเกตสองประการเกี่ยวกับการสร้างข้อความของ transformer:
- Prompt และส่วนเริ่มต้นของคำตอบจะเหมือนกันในทุกๆ rollout
- เฉพาะส่วนท้ายที่แตกต่างกัน (divergent tail) ของแต่ละคำตอบเท่านั้นที่จำเป็นต้องใช้การคำนวณใหม่
ระบบจะสร้าง execution stack ที่รับรู้ถึงสถาปัตยกรรม (architecture-aware) เพื่อบันทึกค่า activations สำหรับส่วน prefix ที่ใช้ร่วมกัน เมื่อมีการสำรวจ branch ใหม่ ระบบจะทำการ replay ส่วน prefix ที่เก็บไว้ใน cache แทนที่จะคำนวณใหม่ จากนั้นจึงรัน backward pass เฉพาะในส่วนที่เพิ่มเข้ามาใหม่เท่านั้น ในทางปฏิบัติ หมายความว่า backward pass จะเข้าถึงตำแหน่งโทเคนน้อยลงมาก ซึ่งช่วยลดการคำนวณดิบ (raw compute) ลงได้ถึง 8 ถึง 16 เท่า
ผลกระทบในทันที
- ประมวลผลตำแหน่งโทเคนได้ 2.1 ล้านตำแหน่งบน GPU H20 เพียงแปดตัว ซึ่งเป็นงบประมาณฮาร์ดแวร์ที่ตามปกติแล้วจะรองรับภาระงานได้เพียงเศษเสี้ยวของจำนวนนี้เท่านั้น
- มุ่งเป้าไปที่คอขวดของ long-context RL โดยตรง ซึ่งต้นทุนด้านหน่วยความจำและการคำนวณจะพุ่งสูงขึ้นอย่างรวดเร็วเมื่อบริบทขยายใหญ่ขึ้น
- ห้องปฏิบัติการสามารถปรับเปลี่ยนการจัดสรร GPU ได้: ฮาร์ดแวร์ชุดเดิมซึ่งโดยหลักแล้วเป็นตัวเร่งความเร็วสำหรับการอนุมาน (inference accelerator) อาจนำมาใช้สำหรับการฝึก (training) ได้ในขณะนี้ แม้ว่าผลลัพธ์อาจแตกต่างกันเมื่อใช้การ์ดรุ่นอื่นก็ตาม
คำถามที่ยังไม่มีคำตอบและข้อจำกัด
การประกาศนี้ไม่ได้ระบุตัวเลขความเร็วในการฝึกและเส้นโค้งการลู่เข้า (convergence curves) ดังนั้นเราจึงไม่ทราบว่าการลดการคำนวณลงนั้นจะช่วยลดเวลาที่ใช้จริง (wall-clock time) หรือเพียงแค่ทำให้การใช้งาน GPU (GPU occupancy) ต่ำลงเท่านั้น วิธีการนี้ถูกอธิบายไว้สำหรับการสุ่มแบบ autoregressive แต่พฤติกรรมเมื่อใช้กับกลยุทธ์แบบ non-autoregressive หรือแบบผสม (hybrid) ยังไม่ได้รับการทดสอบ เนื่องจาก H20 เป็นตัวเร่งความเร็วสำหรับการอนุมานเป็นหลัก ประสิทธิภาพบนการ์ดสำหรับฝึกที่นิยมมากกว่า เช่น H100 หรือ B200 อาจแตกต่างออกไป
การทดสอบประสิทธิภาพ (benchmarks) จากหน่วยงานอิสระยังไม่ได้ยืนยันตัวเลขของ LongStraw หากไม่มีการตรวจสอบจากบุคคลที่สาม ชุมชนควรพิจารณาผลลัพธ์เหล่านี้ว่าเป็นสิ่งที่น่ามีความหวังแต่ยังเป็นเพียงข้อมูลเบื้องต้นเท่านั้น
สิ่งที่จะเกิดขึ้นตามมา
หากแนวคิด branch-replay สามารถขยายไปสู่ขั้นตอนการทำ fine-tuning ด้วย RL แบบอื่นๆ เช่น Direct Preference Optimization (DPO) หรือ Proximal Policy Optimization (PPO) อุปสรรคด้านต้นทุนสำหรับโมเดลที่มีบริบทระยะยาวก็อาจหมดไป
สิ่งที่ควรจับตามอง
- ความพยายามในการทำซ้ำ (replication) โดยบุคคลที่สามบนสถาปัตยกรรม GPU ที่หลากหลาย
- ข้อมูลอัปเดตจาก LongStraw เกี่ยวกับปริมาณงานในการฝึก (training throughput) และคุณภาพของโมเดลขั้นสุดท้ายเมื่อเทียบกับ pipeline มาตรฐาน
