หัวข้อ: บั๊กเกิดจากแกน (axis) เพียงแกนเดียว
บั๊กที่ซ่อนอยู่ในส่วนการทำงาน chunked-scan ของ PyTorch ทำให้โมเดลแบบ hybrid บางรุ่นสามารถแอบดูโทเคนในอนาคตได้เมื่อไม่มี fast fused kernels ซึ่งส่งผลกระทบต่อ checkpoint อย่าง Zamba2-1.2B และ Nemotron-H-8B ความผิดพลาดนี้เกิดจากการทำ reduction เพียงครั้งเดียวบนแกน (axis) ที่ผิดพลาด และมันจะปรากฏขึ้นในเส้นทางการประมวลผล (execution path) ที่ pipeline CI และการรันบน CPU ส่วนใหญ่ใช้งาน
ทำไมบั๊กนี้ถึงสำคัญ
โมเดลแบบ hybrid และ state-space ไม่ได้พึ่งพาเพียงแค่ attention อีกต่อไป แต่ยังใช้ linear recurrences, scans และ convolutions ด้วย การทำ masking เพื่อบล็อกโทเคนในอนาคตใน attention matrix ไม่ ได้เป็นการรับประกันความเป็นเหตุเป็นผล (causality) สำหรับการทำงานส่วนเสริมเหล่านั้นโดยอัตโนมัติ เมื่อไม่มี fast fused kernels ที่ปกติจะจัดการการ scan ได้อย่างถูกต้อง PyTorch จะถอยกลับไปใช้เส้นทาง chunked-scan แบบ pure-Python แทน ซึ่งเส้นทางสำรองนี้มีปัญหาการสลับแกน (axis-mixup) ทำให้ข้อมูลจากตำแหน่งที่อยู่ถัดไปสามารถไหลย้อนกลับมาได้ในระหว่างการทำ inference
การรั่วไหลนี้เกิดขึ้นอย่างเงียบเชียบ มันไม่ทำให้โมเดลพังหรือแสดงข้อผิดพลาดที่ชัดเจน แต่กลับทำให้ค่า loss และ perplexity ดูต่ำกว่าความเป็นจริง เพราะโมเดลกำลัง "โกง" โดยการเห็นโทเคนชุดเดียวกับที่มันควรจะทำนาย ดังนั้น การประเมินผลใดๆ ในขั้นตอนถัดไป (downstream) ที่เชื่อถือตัวชี้วัดเหล่านี้ จึงเป็นการสร้างขึ้นบนรากฐานที่ผิดพลาด
วิธีการค้นพบปัญหา
นักวิจัยได้เปรียบเทียบการทำ forward pass สองครั้งผ่านโมเดลเดียวกัน:
- ลำดับโทเคนแบบสุ่ม
- ลำดับเดิมที่เปลี่ยนโทเคนเพียงตัวเดียว
พวกเขาทำการวัดความแตกต่างของ hidden states ทีละเลเยอร์ การตรวจสอบ mask ไม่พบสิ่งผิดปกติ แต่การตรวจสอบรายเลเยอร์ (per-layer audit) โดยการฉีดข้อผิดพลาด (injected faults) สามารถตรวจพบข้อผิดพลาดที่ฉีดเข้าไปได้ถึง 192 จาก 192 ครั้ง ซึ่งเป็นการยืนยันการรั่วไหลของข้อมูล
การตรวจสอบไลบรารี transformers อย่างรวดเร็วพบว่า:
- Zamba2-1.2B มีการรั่วไหลเมื่อตั้งค่า chunk size เป็น 256
- Nemotron-H-8B มีการรั่วไหลที่ chunk size 128
- Checkpoint อื่นๆ ส่วนใหญ่ไม่พบการรั่วไหลภายใต้เงื่อนไขเดียวกัน
ข้อบกพร่องนี้อยู่ในเส้นทางการทำงานของโค้ด chunked-scan ซึ่งจะทำงานเมื่อไม่มี optional fused kernels ซึ่งรวมถึง:
- การประมวลผลบน CPU ทั้งหมด
- สภาพแวดล้อม GPU ที่ไม่มีแพ็กเกจ fused-kernel เฉพาะทาง
- การติดตั้ง PyTorch มาตรฐานที่ไม่ได้รวม dependencies เสริม
เนื่องจากบั๊กนี้จะปรากฏขึ้นเมื่อไม่มี kernels เหล่านั้นเท่านั้น มันจึงอาจปรากฏขึ้นในสภาพแวดล้อม CI และบน CPU ได้
ใครบ้างที่มีความเสี่ยงและผลกระทบที่ตามมา
ทีมใดก็ตามที่ฝึกฝน (train), ปรับจูน (fine-tune) หรือประเมินผลโมเดลแบบ hybrid โดยไม่มี fused kernels มีความเสี่ยงที่จะเผยแพร่ตัวเลขประสิทธิภาพที่สูงเกินจริง การปรับปรุงของค่า loss หรือ perplexity ที่เห็นนั้นเป็นเพียงภาพลวงตา เพราะโมเดลได้ "มองไปข้างหน้า" เรียบร้อยแล้ว สำหรับกลุ่มวิจัย สิ่งนี้อาจนำไปสู่การกล่าวอ้างที่ผิดพลาดเกี่ยวกับผลลัพธ์ระดับ state-of-the-art สำหรับการใช้งานเชิงพาณิชย์ มันอาจทำให้เกิดข้อผิดพลาดในงานสร้างข้อความ (generation tasks) ที่โมเดลไม่ได้เรียนรู้อย่างแท้จริง
ข้อโต้แย้ง
นักพัฒนาบางคนแย้งว่าการใช้ causal mask ที่ถูกต้องนั้นเพียงพอแล้วในการป้องกันการรั่วไหลของโทเคนในอนาคต แต่บั๊กนี้พิสูจน์ว่าแนวคิดนั้นไม่เป็นความจริง เพราะ scans, convolutions และเลเยอร์ normalization บางอย่างสามารถข้ามผ่าน mask ไปได้ทั้งหมด ปัญหา axis-mixup ใน chunked scan แสดงให้เห็นว่าต้องมีการบังคับใช้ความเป็นเหตุเป็นผล (causality) ใน ทุกที่ ที่ข้อมูลไหลผ่าน ไม่ใช่แค่ใน attention matrix เท่านั้น
สิ่งที่ควรระวังต่อไป
- Dependency hygiene: ติดตั้งแพ็กเกจ fast fused-kernel ในโหนดการฝึกฝนและโหนดการทำ inference ทั้งหมด โดยเฉพาะอย่างยิ่งใน CI pipelines
- Audit scripts: ทำตามขั้นตอนการตรวจสอบสองขั้นตอนที่ผู้ค้นพบแนะนำ:
- ฉีดข้อผิดพลาดที่ทราบค่าลงใน checkpoint ที่คุณกำลังทดสอบเพื่อใช้เป็นตัวควบคุม (positive control)
- รันลำดับ (sequences) ที่ยาวกว่าขนาด chunk หรือ window size ของโมเดล เพราะลำดับที่สั้นเกินไปจะไม่ทำให้เห็นการรั่วไหล
การรันการตรวจสอบบน checkpoint แบบ hybrid หรือ state-space ใดๆ ที่มีความยาวลำดับเกินขนาด chunk จะช่วยเผยให้เห็นว่าโมเดลยังคงมีความเสี่ยงอยู่หรือไม่
บทสรุป: แม้แต่โมเดลที่ผ่านการทดสอบมาตรฐานทุกอย่างก็อาจ "โกง" ได้อย่างเงียบเชียบเมื่อเส้นทางการประมวลผลถอยกลับไปใช้การทำงานที่มีบั๊ก การตรวจสอบว่ามี fast fused kernels อยู่จริง และการทดสอบการรั่วไหลที่นอกเหนือจาก attention masks อย่างชัดเจน กลายเป็นขั้นตอนที่จำเป็นก่อนที่จะเชื่อถือตัวชี้วัดของโมเดลแบบ hybrid ใดๆ
