پرش به محتوای اصلی
پرش به محتوای مقاله

«انعطاف در انواع داده»؛ هدف اصلی چارچوب آزمایشی hijax

·۲۲ تیر ۱۴۰۵۲۵ دقیقه مطالعه۲ بازدید
راهنما
تعریف انواع داده جدید در JAX با کتابخانه hijax
تعریف انواع داده جدید در JAX با کتابخانه hijax
اشتراک‌گذاری
واقعاً چه چیز جدید است؟

معرفی مفهومی به نام hi types که برخلاف پایتری‌ها، در نمایش داخلی JAX دارای هویت واحد هستند و اجازه می‌دهند نوع مماس (tangent) با نوع اصلی متفاوت باشد.

تصور کنید می‌خواهید در یک مدل یادگیری عمیق، داده‌ها را به‌صورت فشرده ذخیره کنید اما هم‌زمان نیاز دارید که تمام محاسبات ریاضی و مشتقات مدل به‌طور دقیق و بدون خطا باقی بماند. تا پیش از این، توسعه‌دهندگان JAX برای رسیدن به این هدف مجبور بودند بین ساختارهای سخت‌گیرانه یا مدل‌های منعطف اما مبهمی به نام pytree یکی را انتخاب کنند. پیاده‌سازی یک آرایه کوانتیده معمولاً کاربر را مجبور می‌کرد تا بین یک ساختار صلب و یک pytree منعطف اما نامشخص، یکی را برگزیند.

طبق مستندات فنی این پروژه، کتابخانه آزمایشی hijax این تضاد را برطرف کرده است. این ابزار به برنامه‌نویسان اجازه می‌دهد «انواع سطح بالا» (hi types) تعریف کنند که به‌عنوان مقادیر واحد و شناسنامه‌دار در نمایش‌های داخلی JAX (jaxprs) وجود دارند.

تا ۱۲ ژوئیه ۲۰۲۶، اکثر توسعه‌دهندگان برای دسته‌بندی آرایه‌ها از پایتری (pytree) — که شبیه به یک جعبه ابزار است که وسایل مختلف را کنار هم می‌گذارد اما هویت هر وسیله را به تنهایی نمی‌بیند — استفاده می‌کردند. مشکل این بود که پایتری‌ها شفاف هستند؛ یعنی JAX فقط مجموعه‌ای از برگ‌های آرایه‌ای مستقل را می‌بیند. بنابراین، اعمال قواعد ثابت (invariants) دشوار بود؛ مثلاً نمی‌شد به‌راحتی تضمین کرد که یک ضریب مقیاس (scale factor) با ابعاد یک مقدار کوانتیده مطابقت دارد یا خیر. همچنین تعریف رفتارهای سفارشی برای مشتق‌گیری (differentiation) و دسته‌بندی (batching) در این ساختار ممکن نبود. در حالی که پایتری صرفاً یک ظرف است، یک hi type هویت مستقل خود را در jaxprs دارد، مفهوم خاص خود را از دسته‌بندی تحت vmap دارد و می‌تواند اطلاعات مربوط به توزیع داده (sharding) را در نوع داده خود برای حالت توزیع صریح حمل کند.

نکته قابل توجه این است که hi typeها در jaxprs به‌عنوان یک مقدار واحد از یک نوع واحد ظاهر می‌شوند، نه مجموعه‌ای از برگ‌ها. آن‌ها اجازه می‌دهند قواعد داخلی حفظ شوند تا کاربران تنها از طریق عملیات‌های تعریف‌شده و ثابت، آن‌ها را تولید یا مصرف کنند. حیاتی‌ترین نکته این است که «نوع مماس» (tangent type) آن‌ها می‌تواند با ساختار اصلی (primal) متفاوت باشد؛ یعنی مشتقات لزوماً نباید «همان پایتری، اما برای مماس‌ها» باشند.

همان‌طور که در بحث‌های گذشته‌ی ما درباره‌ی بهینه‌سازی حافظه در مدل‌های زبانی اشاره کردیم، مدیریت دقیق نحوه نمایش داده‌ها در حافظه، کلید افزایش سرعت استنتاج است. در hijax، این انواع داده در jaxprs به‌عنوان یک مقدار واحد ظاهر می‌شوند و کاربران تنها از طریق عملیات‌های تعریف‌شده با آن‌ها تعامل دارند. این موضوع مانع از پراکندگی محاسبات در لایه‌های پایین‌تر می‌شود.

سازوکار انواع سطح بالا

در hijax، توسعه‌دهندگان با ارث‌بری از کلاس HiType یک هویت جدید تعریف می‌کنند. از آن‌جهت که این نوع باید قابل هش (hashable) و قابل مقایسه برای برابری باشد، توسعه‌دهندگان معمولاً از frozen dataclasses برای این منظور استفاده می‌کنند. این فرآیند در سه مرحله اصلی رخ می‌دهد:

  • Lowering (پایین‌افکنی): شما متد lo_ty را پیاده‌سازی می‌کنید تا مشخص کنید کدام آرایه‌های استاندارد JAX (lojax) این نوع را تشکیل می‌دهند. این کار به JAX اجازه می‌دهد انواع پایین‌افکنی‌شده را بدون نیاز به داشتن یک مقدار فیزیکی در دست، محاسبه کند. همچنین متدهای lower_val و raise_val را برای تبدیل بین شیء سطح بالا و لیست اجزای آرایه‌ای آن پیاده می‌کنید. برای مثال، پیاده‌سازی QArrayTy از lo_ty استفاده می‌کند تا لیستی شامل یک ShapedArray برای مقادیر int8 و یک ShapedArray دیگر برای مقیاس float32 برگرداند.
  • Registration (ثبت): شما از register_hitype استفاده می‌کنید تا یک کلاس مقدار پایتونی (مانند یک dataclass) را با تعریف hi type مرتبط کنید. این فراخوانی شامل یک lambda است که نوع داده (مثلاً خواندن شکل و توزیع) را از هر مقدار داده شده محاسبه می‌کند، مشابه روش کار jax.typeof برای آرایه‌ها. پس از ثبت، jax.typeof مستقیماً روی اشیاء سفارشی کار می‌کند و تبدیل‌های JAX آن‌ها را در هر جایی که یک مقدار انتظار می‌رود، می‌پذیرند.
  • Primitives (پریمیتیوها): شما زیرکلاس‌های VJPHiPrimitive را می‌نویسید. این‌ها تنها روش‌های مجاز و sanctioned برای تولید یا مصرف نوع جدید هستند. با تعریف in_avals و out_aval خاص، این پریمیتیوها تضمین می‌کنند که قواعد داخلی هرگز توسط کدهای خارجی نقض نشوند. این پریمیتیوها متد expand را فراهم می‌کنند، جایی که نوع سطح بالا در نهایت برای اجرا به آرایه‌های lojax سازنده خود تجزیه می‌شود.

حل چالش کوانتش

در یک آرایه‌ی کوانتیده (Quantization) — که شامل یک مقدار qvalue (از نوع int8) و یک مقیاس scale (از نوع f32) است — کوانتش در امتداد آخرین محور اتفاق می‌افتد و برای هر ردیف یک مقیاس وجود دارد. در یک پایتری استاندارد، نوع مماس (که برای گرادیان‌ها استفاده می‌شود) مجبور است با انواع برگ‌ها مطابقت داشته باشد. از آن‌جا که آرایه‌های int8 مماس‌های بدیهی float0 دارند، یک آرایه کوانتیده مبتنی بر پایتری نمی‌تواند به‌طور مؤثر گرادیان‌ها را منتقل کند.

اما با hijax، شما می‌توانید to_tangent_aval را تعریف کنید تا مشخص شود مماسِ یک آرایه کوانتیده، در واقع یک آرایه پیوسته float32 است. این کار ساختار اصلی را از ساختار مماس جدا می‌کند (decouple). این جداسازی اجازه می‌دهد گرادیان‌ها از طریق یک «تخمین‌گر مستقیم» (Straight-Through Estimator یا STE) جریان یابند. با پیاده‌سازی قواعد خاص vjp_fwd و vjp_bwd_retval روی پریمیتیوهای Quantize و Dequantize‌، سیستم گام گسسته کوانتش را در مسیر بازگشتی (backward pass) به‌صورت یک تابع همانی (identity function) در نظر می‌گیرد.

برای مثال، پریمیتیو Quantize تضمین می‌کند که x_aval.dtype در هنگام مقداردهی اولیه float32 باشد. متد expand آن منطق واقعی کوانتش را پیاده می‌کند: محاسبه مقیاس به صورت jnp.max(jnp.abs(x), axis=-1) / 127.0 و سپس گرد کردن ورودی مقیاس‌بندی‌شده به int8. پریمیتیو متناظر Dequantize صرفاً عملیات معکوس را انجام می‌دهد: qx.qvalue.astype('float32') * qx.scale[..., None].

عملکرد و اجرا

انواع سفارشی امکان اجرای عملیاتی را می‌دهند که بسیار سریع‌تر از تبدیل دائمی داده‌های کوانتیده به حالت عادی (dequantizing) برای هر مرحله است. نمونه بارز آن پریمیتیو MatmulQ است که ضرب یک آرایه متراکم f32 (به نام x) در یک ماتریس وزنی کوانتیده (qw) را مدیریت می‌کند.

جزئیات پیاده‌سازی MatmulQ:

  • ادغام مقیاس‌ها (Folding Scales): به‌جای تبدیل کامل ماتریس وزنی به حالت متراکم، متد expand مقیاس‌های هر ردیف را مستقیماً در عملوند متراکم ادغام می‌کند: (x * qw.scale) @ qw.qvalue.astype(jnp.float32). این روش از این حقیقت بهره می‌برد که مقیاس‌ها در امتداد محور انقباض به‌صورت ردیفی اعمال می‌شوند.
  • بررسی نوع (Type Checking): این پریمیتیو در زمان ساخت بررسی می‌کند که محورهای انقباض توزیع‌نشده (unsharded) باشند و در صورت عدم تطابق ابعاد یا مشخصات توزیع، خطای TypeError صادر می‌کند. به‌طور خاص بررسی می‌کند که x_aval.shape[1] == q_aval.shape[0] باشد و اطمینان حاصل می‌کند که هیچ‌کدام از محورهای انقباض sharded نباشند.
  • بهره‌وری: انجام ضرب ماتریسی سنگین مستقیماً روی محموله int8 از سربار حافظه ناشی از ایجاد یک کپی کامل و متراکم (dequantized copy) جلوگیری می‌کند. خروجی یک ShapedArray با ابعاد حاصل و نوع داده float32 است.

به دلیل شناسایی این انواع در jaxprs، کد ردیابی‌شده (traced code) پاکیزه می‌ماند. یک عملیات کوانتیده به‌عنوان یک فراخوانی واحد به یک پریمیتیو سطح بالا (مثلاً call_hi_primitive[_prim=Quantize[]{}]) ظاهر می‌شود، نه مجموعه‌ای از دست‌کاری‌های پراکنده آرایه‌ای. این موضوع از خطاهای «نشت ردیاب» (leaked tracer) جلوگیری می‌کند. برای مثال، اگر کاربر سعی کند درون یک تابع jit به qx.qvalue دسترسی پیدا کند، JAX خطای AttributeError می‌دهد زیرا شیء در آن لحظه در واقع یک DynamicJaxprTracer است (به‌طور دقیق: DynamicJaxprTracer has no attribute qvalue).

دسترسی مستقیم به ویژگی‌ها (attributes) تنها در متد expand مجاز است، جایی که JAX قبلاً تصمیم به استفاده از اجزای lojax گرفته است. تلاش برای فراخوانی یک سازنده (constructor) روی آرایه‌های ردیابی‌شده به‌جای استفاده از یک پریمیتیو، منجر به خطای TypeError: No constant handler for type: <class 'jax._src.interpreters.partial_eval.DynamicJaxprTracer'> می‌شود.

تبدیل‌های پیشرفته: vmap و scan

مدیریت دسته‌ها (batches) در hijax نیازمند یک MappingSpec است. از آن‌جا که JAX نمی‌تواند حدس بزند یک نوع سفارشی چگونه باید نگاشت شود، توسعه‌دهنده باید تعریف کند که رتبه‌ها چگونه از طریق متدهای inc_rank و dec_rank افزایش یا کاهش یابند. این متدها اندازه محور و یک spec را می‌گیرند تا به‌ترتیب نوع عنصر و نوع دسته‌بندی‌شده را برگردانند.

برای QArrayTy در حالت vmap (بردارسازی)، یک QArraySpec یک زیرکلاس ساده از MappingSpec است که داده داخلی ندارد. یک دسته از آرایه‌های q8[2,3] که روی یک محور پیشرو جدید پشته شده‌اند، تبدیل به q8[n,2,3] می‌شوند. متد dec_rank محور نگاشت‌شده را حذف کرده و مشخصات توزیع (sharding spec) را با برش دادن پارتیشن توزیع به‌روز می‌کند.

در سمت پریمیتیو، قواعد دسته‌بندی (batch rules) باید برای مدیریت ترکیباتی از آرگومان‌های دسته‌بندی‌شده و نشده پیاده‌سازی شوند. برای مثال، quantize_batch بررسی می‌کند که آیا بعد ورودی d برابر با None است یا خیر. اگر نباشد، محور را به موقعیت ۰ منتقل کرده، کوانتش را اعمال می‌کند و نتیجه را با یک QArraySpec() برمی‌گرداند.

در مورد jax.lax.scan (حلقه‌ها)، سیستم بر leading_axis_spec تکیه دارد. این اجازه می‌دهد حلقه در هر گام یک برش (slice) از نوع سطح بالا را مصرف کند. چون scan همیشه روی محور پیشرو حرکت می‌کند، اگر تمام مقادیر اسکن‌شده از نوع hi type باشند، توسعه‌دهنده باید length را به‌طور صریح ارائه دهد، زیرا هیچ آرگومان آرایه‌ای برای استنتاج اندازه توسط JAX وجود ندارد. بنابراین یک عملیات sum_dequantized می‌تواند روی پشته‌ای از آرایه‌های q8 پیمایش کرده و مجموع مقادیر تبدیل‌شده آن‌ها را جمع کند.

توزیع داده و محاسبات پراکنده

در حالت توزیع صریح JAX، اطلاعات توزیع (sharding) بخشی از نوع داده هستند. hijax اجازه می‌دهد داده‌های توزیع، مانند یک فیلد NamedSharding را مستقیماً روی HiType ثبت کنید.

  • انتشار (Propagation): پریمیتیوها توزیع را از ورودی به خروجی از طریق قواعد تایپینگ (typing rules) منتقل می‌کنند. برای QArrayTy مقدار qvalue توزیع نوع را حمل می‌کند، در حالی که مقیاس (scale) با حذف آخرین محور از spec با استفاده از self.sharding.update(spec=jax.P(*self.sharding.spec[:-1])) به‌دست می‌آید.
  • تعامل با Lojax: متد lo_ty آرایه‌های سازنده را با توزیع‌های مربوطه مهر می‌زند تا JAX بتواند پارتیشن‌بندی را در سراسر مش (mesh) ردیابی کند. هنگام چاپ، این‌ها به‌عنوان نشانگرهای @ ظاهر می‌شوند (مثلاً q8[8@i,3]).
  • shard_map: برای عبور از مرز دیدهای جهانی (global views) به دیدهای هر-دستگاه (per-device views)، hijax از HiPspec استفاده می‌کند. این یک معادل MappingSpec برای مشخصات پارتیشن است. متد to_lo آن، یک spec سطح بالا را به مشخصات پارتیشن jax.P مجزا برای هر جزء پایین‌افکنی‌شده تبدیل می‌کند. برای QArrayP متد to_lo مقدار (self.spec, self.spec) را برمی‌گرداند زیرا qvalue و scale باید با هم توزیع شوند.
  • تقارن (Symmetry): چون کوانتش ردیفی صرف‌نظر از رتبه اعمال می‌شود، کوانتش تکه‌به‌تکه (shard-by-shard) در یک shard_map دقیقاً با کوانتش جهانی مطابقت دارد. این موضوع با بررسی اینکه qrows.qvalue و qrows.scale دقیقاً با کوانتش جهانی یکی هستند، تایید می‌شود.

تطبیق‌پذیری: از آرایه‌های رتبه-۱ تا تاپل‌ها

این چارچوب فراتر از کوانتش کاربرد دارد. یک مورد استفاده، آرایه‌های رتبه-۱ (Rank-1 arrays) است که در آن یک ماتریس $m \times n$ به‌صورت حاصل‌ضرب خارجی دو بردار: col (با ابعاد f32[m]) و row (با ابعاد f32[n]) نمایش داده می‌شود.

ویژگی‌های آرایه رتبه-۱:

  • جلوگیری از تجسم (Avoiding Materialization): پریمیتیو MatmulR1 ضرب را به‌صورت jnp.outer(x @ r1.col, r1.row) انجام می‌دهد تا تضمین شود که حاصل‌ضرب متراکم کامل $m \times n$ هرگز در حافظه ساخته (materialize) نشود.
  • فضای مماس (Tangent Space): برخلاف آرایه‌های کوانتیده، آرایه‌های رتبه-۱ تحت عمل جمع بسته نیستند (جمع دو ماتریس رتبه-۱ عموماً رتبه-۲ است). بنابراین، نوع مماس برای Rank1Ty به‌صورت یک ShapedArray متراکم از f32 تعریف شده است. اختلالات از روی منیفولد خارج می‌شوند، به این معنی که مماس‌ها و کوتانس‌ها (cotangents) باید متراکم باشند.
  • دسترسی پریمیتیو: برای حفظ مرز انتزاع، از پریمیتیو Factors برای بازیابی اجزای col و row در طول قواعد ردیابی‌شده استفاده می‌شود، نه دسترسی مستقیم به ویژگی‌ها. پریمیتیو Outer به‌عنوان سازنده عمل کرده و دو بردار را گرفته و یک شیء Rank1 برمی‌گرداند.

علاوه بر این، TupTy اجازه ایجاد کانتینرهای عمومی را می‌دهد. برخلاف ساختار ثابت یک آرایه کوانتیده، یک hi-tuple می‌تواند هر ترکیبی از انواع آرایه و سایر انواع سطح بالا، از جمله تاپل‌های تو در تو، را در خود جای دهد.

پیاده‌سازی تاپل و نگاشت:

  • تفویض (Delegation): TupTy متدهای lo_ty ،lower_val و to_tangent_aval را به انواع سازنده‌اش تفویض می‌کند. این کلاس از itertools.islice در raise_val استفاده می‌کند تا لیست تخت‌شده‌ی آرایه‌های lojax را دوباره به انواع اجزای سطح بالای مربوطه تقسیم کند.
  • اندیس‌گذاری پویا: پریمیتیوهایی مانند GetTupElt از یک پارامتر استاتیک self.idx که از طریق params ارسال شده، برای دسترسی به عناصر خاص استفاده می‌کنند. قاعده batch برای GetTupElt باید هر دو حالت ورودی‌های دسته‌بندی‌شده و نشده را مدیریت کند، زیرا ورودی‌های in_dims می‌توانند None باشند.
  • دسته‌بندی منعطف: TupSpec برای هر جزء یک ورودی محور حمل می‌کند. این اجازه می‌دهد هر عنصر یک تاپل در یک عملیات vmap روی محورهای مختلف نگاشت شود (یا اصلاً نشود)، همان‌طور که در تابع swap مشاهده می‌شود که در آن عنصر اول در ورودی و عنصر دوم در خروجی نگاشت شده است. این کار اجازه ایجاد امضاهای پیچیده‌ای مانند Tup{float32[3],float32[]} -> Tup{float32[],float32[3]} را می‌دهد.

این انعطاف‌پذیری ساختاری به این معنی است که توسعه‌دهندگان اکنون می‌توانند انواع داده‌های متناسب با دامنه تخصصی خود را طراحی کنند که به‌طور کامل در اکوسیستم تبدیل‌های JAX — شامل jit ،grad ،vmap ،scan و shard_map — مشارکت کنند، بدون اینکه کارایی لایه‌ی پایین‌افکنی XLA را فدا کنند.

چرا این موضوع مهم است؟

این ابزار با فراهم کردن امکان تعریف انواع داده‌ای که هم‌زمان با autodiff و vmap سازگارند، تخصص توسعه‌دهندگان را در پیاده‌سازی عملیات‌های بهینه (مانند کوانتش) افزایش می‌دهد. اعتبار این رویکرد در این است که اجازه می‌دهد مدل‌ها بدون از دست دادن دقت ریاضی، از سرعت سخت‌افزاری داده‌های کم‌دقت بهره ببرند.

تأثیر برای ایران

این ابزار برای پژوهشگران هوش مصنوعی در ایران که با محدودیت سخت‌افزاری (GPU/TPU) مواجه‌اند، راهکاری حیاتی برای اجرای مدل‌های بزرگتر از طریق کوانتش بهینه فراهم می‌کند.

·نگاه ما
تحریریه دات‌هوش

hijax در واقع شکاف میان انتزاع سطح بالای پایتون و اجرای سخت‌گیرانه XLA را می‌بندد. با انتقال هویت نوع داده از سطح پایتون به درون jaxprs، JAX دیگر نیازی ندارد هر بار ساختارهای پیچیده را به تکه‌های کوچک تجزیه کند. این تغییر پارادایم، مسیر را برای پیاده‌سازی انواع داده‌های سخت‌افزاری بهینه‌تر (مانند فرمت‌های جدید FP8 یا INT8) بدون تغییر در موتور اصلی JAX هموار می‌کند.

منابع

این گزارش با خط‌لولهٔ خودکار دات‌هوش از منابع معتبر جهانی تدوین و زیر نظر تحریریه منتشر شده است. روش کار ما

گفتگو

پنج‌شنبه‌های هوش‌محور

بسته‌ی هفتگی دات‌هوش

۵ خبر، ۲ ابزار، ۱ پرامپت در هر شماره. به‌زودی راه‌اندازی می‌شود — هر پنج‌شنبه صبح.

خبر کلیدی
ابزار کاربردی
پرامپت حرفه‌ای
تحلیل پژوهش
به‌زودی
زاویه‌ی ایرانی
به‌زودی
تمرین این هفته
به‌زودی

راهنماهای دات‌هوش

راهنماهای کاربردیِ دات‌هوش برای کار با هوش مصنوعی — از همین‌جا شروع کنید:

دات‌هوش

راهنمای فارسی هوش مصنوعی — با نگاه به ایران

اخبار روزانه، معرفی ابزارها و مدل‌ها، و آموزشِ کار با هوش مصنوعی؛ همیشه با این پرسش که از ایران چه چیزی کار می‌کند و چه چیزی نه.