शीर्षक: तो बग एका अक्षाचा (axis) होता

PyTorch च्या chunked-scan अंमलबजावणीमधील (implementation) एका लपलेल्या बगमुळे, जेव्हा fast fused kernels उपलब्ध नसतात, तेव्हा काही विशिष्ट hybrid models भविष्यातील tokens पाहू शकतात, ज्यामुळे Zamba2-1.2B आणि Nemotron-H-8B सारख्या checkpoints वर परिणाम होतो. हा दोष चुकीच्या अक्षावर (axis) केलेल्या एका single reduction मुळे निर्माण होतो आणि तो अशा execution path वर दिसून येतो ज्यावर बहुतेक CI pipelines आणि CPU runs अवलंबून असतात.

हा बग का महत्त्वाचा आहे

Hybrid आणि state-space models आता केवळ attention वर अवलंबून नसून, ते linear recurrences, scans आणि convolutions देखील वापरतात. Attention matrix मध्ये भविष्यातील tokens रोखणारी masking operation, या अतिरिक्त ऑपरेशन्ससाठी causality ची स्वयंचलितपणे खात्री देत नाही. जेव्हा scan योग्यरित्या हाताळणारे fast fused kernels उपलब्ध नसतात, तेव्हा PyTorch pure-Python chunked-scan path कडे वळते. त्या fallback मध्ये axis-mixup असतो, ज्यामुळे inference दरम्यान नंतरच्या स्थानांची माहिती मागे (backward) प्रवाहित होऊ शकते.

ही माहिती गळती (leak) शांतपणे होते. यामुळे मॉडेल क्रॅश होत नाही किंवा कोणताही स्पष्ट एरर येत नाही. त्याऐवजी, यामुळे loss आणि perplexity कृत्रिमरित्या कमी दिसतात, कारण मॉडेल प्रत्यक्षात फसवणूक करत असते—ज्या tokens चे ते भाकीत करायचे आहे, तेच ते आधीच पाहू लागते. त्यामुळे, या metrics वर विश्वास ठेवणारे कोणतेही downstream evaluation एका कमकुवत पायावर आधारित असते.

ही समस्या कशी उघडकीस आली

संशोधकांनी एकाच मॉडेलमधून होणाऱ्या दोन forward passes ची तुलना केली:

  1. एक random token sequence.
  2. एकच sequence ज्यामध्ये एक token बदललेला आहे.

त्यांनी layer by layer hidden states मधील फरक मोजला. Mask inspection मध्ये काहीही आढळले नाही, परंतु faults इंजेक्ट करून केलेल्या per-layer audit मध्ये 192 पैकी 192 इंजेक्ट केलेले एरर्स आढळले, ज्यामुळे ही माहिती गळती (leak) झाल्याची पुष्टी झाली.

transformers लायब्ररीच्या जलद सर्वेक्षणातून असे दिसून आले की:

  • Zamba2-1.2B मध्ये जेव्हा chunk size 256 सेट केला जातो, तेव्हा leak होतो.
  • Nemotron-H-8B मध्ये 128 च्या chunk size वर leak होतो.
  • इतर बहुतेक checkpoints मध्ये त्याच परिस्थितीत कोणतीही गळती दिसून आली नाही.

हा दोष chunked-scan code path मध्ये आहे जो तेव्हा चालतो जेव्हा optional fused kernels उपलब्ध नसतात. यामध्ये खालील गोष्टींचा समावेश होतो:

  • सर्व CPU execution.
  • विशिष्ट fused-kernel packages नसलेले GPU environments.
  • अतिरिक्त dependencies नसलेले standard PyTorch installations.

हा बग फक्त ते kernels उपलब्ध नसतानाच दिसून येतो, त्यामुळे तो CI environments आणि CPUs वर समोर येऊ शकतो.

कोणाला धोका आहे आणि त्याची किंमत काय आहे

जी कोणतीही टीम fused kernels शिवाय hybrid models प्रशिक्षित (train), fine-tune किंवा evaluate करते, तिला फुगवलेले (inflated) performance numbers प्रकाशित करण्याचा धोका असतो. Loss किंवा perplexity मधील दिसणारा सुधारणा हा केवळ आभास आहे; मॉडेलने प्रत्यक्षात "पुढे पाहिलेले" (looked ahead) असते. संशोधन गटांसाठी, यामुळे state-of-the-art निकालांबद्दल दिशाभूल करणारे दावे होऊ शकतात. व्यावसायिक उपयोगांसाठी (commercial deployments), यामुळे generation tasks मध्ये अशा चुका होऊ शकतात ज्या मॉडेलने खरोखर कधीही शिकल्या नसतात.

प्रतिवाद (The counter-argument)

काही डेव्हलपर्स असा युक्तिवाद करतात की भविष्यातील tokens ची गळती रोखण्यासाठी योग्यरित्या लागू केलेले causal mask पुरेसे आहे. हा बग त्या कल्पनेला चुकीचे ठरवतो: scans, convolutions आणि काही normalization layers पूर्णपणे mask ला बगल देऊ शकतात. Chunked scan मधील axis-mixup हे दर्शवते की causality केवळ attention matrix मध्येच नाही, तर डेटा जिथे जिथे प्रवाहित होतो तिथे सर्वत्र लागू केली पाहिजे.

पुढे काय पाहावे

  • Dependency hygiene: सर्व training आणि inference nodes वर, विशेषतः CI pipelines मध्ये, fast fused-kernel packages इंस्टॉल करा.
  • Audit scripts: शोधकर्त्यांनी सुचवलेल्या दोन-टप्प्यांच्या audit चे अनुसरण करा:
    • तुम्ही परीक्षण करत असलेल्या checkpoint मध्ये positive control म्हणून एक ज्ञात fault इंजेक्ट करा.
    • मॉडेलच्या chunk किंवा window size पेक्षा जास्त लांब sequences चालवा; लहान sequences मुळे गळती कधीही समोर येणार नाही.

Chunk size पेक्षा जास्त sequence length असलेल्या कोणत्याही hybrid किंवा state-space checkpoint वर audit चालवल्यास मॉडेल अजूनही असुरक्षित आहे की नाही हे समजेल.

निष्कर्ष (Takeaway): अगदी प्रत्येक standard test पास होणारे मॉडेल देखील जेव्हा execution path एखाद्या buggy implementation कडे वळते, तेव्हा ते शांतपणे फसवणूक करू शकते. Fast fused kernels उपलब्ध आहेत याची खात्री करणे—आणि attention masks च्या पलीकडे गळतीसाठी स्पष्टपणे परीक्षण करणे—हे आता कोणत्याही hybrid model च्या metrics वर विश्वास ठेवण्यापूर्वीचे आवश्यक टप्पे आहेत.