KDA Chunk Backward — طراحی هسته TPU/Pallas

مرجع: گزارش فنی Kimi Linear - https://github.com/MoonshotAI/Kimi-Linear/blob/master/tech_report.pdf

کد:

  • tokamax/_src/ops/experimental/kda/base.py — آرگومان عمومی KDA و قرارداد حالت بازگشتی
  • tokamax/_src/ops/experimental/kda/reference.py — پیاده‌سازی خالص ارجاع بازگشتی JAX
  • tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu.py — آداپتور VJP سفارشی توکامکس
  • tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_kernel.py — نقاط ورود عمومی سطح پایین هسته
  • tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_types.py — قرارداد KdaResiduals تایپ شده
  • tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_fwd_kernel.py — ساخت ماتریس‌های باقیمانده رو به جلو و حالت‌های ذخیره شده
  • tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_bwd_kernel.py — ارکستراتور عقب مانده و هسته های پالاس
  • tokamax/_src/ops/experimental/kda/common.py — کمک‌کننده‌ی گیت مشترک cumsum، بازگشت حالت و mini-batch
  • tokamax/_src/ops/experimental/kda/cp_utils.py — ادغام گرادیان موازی-زمینه
  • tokamax/_src/ops/experimental/kda/utils.py — کمک‌کننده‌های رو به عقب برای ترازبندی، عدم ترازبندی و نرمال‌سازی L2

۱. هدف و جریان داده در پاس رو به عقب

KDA توالی را در بخش‌هایی از گام‌های زمانی BT پردازش می‌کند، که به طور پیش‌فرض BT=64 است.

برای هر بخش، خروجی رو به جلو مجموع دو بخش است:

\mathbf{o}
= \underbrace{s\,(\mathbf{q}\odot 2^{\mathbf{g} })\,\mathbf{h} }_{\text{inter: reads historical state} }
+ \underbrace{\mathrm{tril}(\mathbf{A}_{qk})\,\mathbf{v}_{\text{new} } }_{\text{intra: reads earlier tokens within the chunk} }

به طور شهودی، خروجی یک پرس و جو از دو مسیر حاصل می‌شود:

  • مسیر بین رشته‌ای: اطلاعاتی را از وضعیت تاریخی h که هنگام ورود به این بخش از داده وجود دارد، می‌خواند؛
  • مسیر درون: اطلاعات را از نتایج به‌روز شده توکن‌های قبلی درون این بخش می‌خواند.

ورودی Tokamax، PallasMosaicTpuKimiDeltaAttentionVjp._fwd(...) است که باقیمانده‌های رو به جلو و کتانژانت‌های خروجی تایپ شده را به chunk_kda_bwd_custom(...) ارسال می‌کند. هماهنگ‌کننده، گرادیان خروجی do و گرادیان حالت نهایی اختیاری dht را دریافت می‌کند و موارد زیر را تولید می‌کند:

  • dq, dk, dv : گرادیان‌ها نسبت به پرس‌وجو/کلید/مقدار؛
  • db : گرادیان نسبت به beta ؛
  • dg : گرادیان نسبت به ورودی gate عمومی. تابع cumsum معکوس هسته ادغام‌شده ابتدا گرادیان گیت فعال‌شده به ازای هر توکن را قبل از cumsum تولید می‌کند؛ وقتی use_gate_in_kernel=True ، kda_gate_bwd(...) آن گرادیان را بیشتر به گیت عمومی خام نگاشت می‌کند؛
  • dh0 : گرادیان نسبت به حالت اولیه، فقط زمانی برگردانده می‌شود که فراخواننده initial_state ارائه دهد؛
  • dA, dbias : گرادیان‌ها نسبت به a_log و delta_time_bias اختیاری، زمانی که فعال‌سازی گیت در داخل هسته محاسبه می‌شود.

dht یک ورودی اختیاری است که گرادیان حالت نهایی را نشان می‌دهد. این ورودی فقط زمانی غیرصفر است که حالت نهایی همچنان توسط یک تابع زیان پایین‌دست مورد استفاده قرار گیرد؛ وقتی یک لایه توجه معمولی فقط خروجی o را مصرف می‌کند، dht می‌تواند None باشد و به عنوان صفر در نظر گرفته می‌شود.

۱.۱ مقادیر ذخیره شده رو به جلو و مرز محاسبه مجدد رو به عقب

مسیر برگشت، تمام محاسبات رو به جلو را از ابتدا انجام نمی‌دهد. پیاده‌سازی فعلی، کمیت‌های مورد نیاز مسیر برگشت را به سه دسته تقسیم می‌کند.

مانده‌های ذخیره شده رو به جلو:

  • Aqk : ماتریس توجه کلید-پرس‌وجوی درون‌قطعه‌ای، شکل [H, B, T, BT] ؛
  • Akk : ماتریس معکوس WY، یعنی A_kk^{-1} ، به شکل [H, B, T, BT] ؛
  • h : حالت پنهان در ابتدای هر بخش، شکل [H, B, NT, K, V] ؛ تابع forward آن را در KdaResiduals هنگام rematerialize_for_backward=False حفظ می‌کند.

محاسبه مجدد رو به عقب:

  • w : وزن موثر کلید/پاک کردن WY؛
  • qg : پرس‌وجوی دروازه‌دار؛
  • kg : کلید قفل‌دار؛
  • v_new : مقدار اصلاح‌شده با واحد WY.

محاسبه مجدد اختیاری:

  • وقتی use_gate_in_kernel=True ، هر دو مسیر آماده‌سازی رو به عقب، در حال حاضر دروازه پس از cumsum را از g_org ، a_log و delta_time_bias اختیاری از طریق kda_gate_chunk_cumsum ، از جمله مسیر saved- h که در آن باقیمانده رو به جلو در حال حاضر شامل g_cumsum نیز می‌شود، دوباره محاسبه می‌کنند.
  • وقتی rematerialize_for_backward=True ، تابع backward مسیر saved- h را طی نمی‌کند، بلکه در عوض، حالت forward را برای بازیابی h و v_new دوباره اجرا می‌کند.

این یک بازسازی دستی در VJP سفارشی Mosaic است، نه فراخوانی jax.checkpoint یا jax.remat .

۱.۲ نمایش WY و معناشناسی گیت

به‌روزرسانی وضعیت قاعده دلتای درون‌قطعه‌ای ذاتاً سریالی است: هر توکن قدرت نوشتن خود را با beta کنترل می‌کند و توکن بعدی به وضعیتی که توسط توکن قبلی به‌روزرسانی می‌شود بستگی دارد.

نمایش WY این به‌روزرسانی سریالی را در یک ماتریس پایین-مثلثی A_kk کدگذاری می‌کند و از ماتریس معکوس ذخیره‌شده رو به جلو A = A_kk^{-1} برای بازنویسی محاسبات درون-بخشی در متمول‌های موازی استفاده می‌کند:

\mathbf{u}=\mathbf{A}(\mathbf{v}\odot\boldsymbol\beta)
\mathbf{w}=\mathbf{A}(\mathbf{k}\odot\boldsymbol\beta\odot 2^{\mathbf{g} })
\mathbf{v}_{\text{new} }=\mathbf{u}-\mathbf{w}\,\mathbf{h}

Akk معکوس ذخیره شده (I+L)^{-1} است. ماتریس اکیداً پایین مثلثی L از قبل به beta سطری وابسته است و beta نیز به طور صریح در مقادیر و سمت راست کلید-درب‌دار مورد استفاده برای ساخت u و w ضرب می‌شود. بنابراین، به عقب باید هر دو مسیری را که beta از طریق آنها بر نتیجه تأثیر می‌گذارد، در نظر گرفت.

در معادلات زیر، g نشان‌دهنده‌ی گیت تجمعی محلی-قطعه‌ای در فضای log2 است. این یک کمیت داخلی است، نه آرگومان خام گیت عمومی. بنابراین، تمام واپاشی‌ها به صورت 2^(...) نوشته می‌شوند و در هسته از طریق exp2 محاسبه می‌شوند:

  • qg = q * 2^g : پرس و جو، وضعیت تاریخیِ از شروعِ هر بخش تا موقعیت فعلی را می‌خواند؛
  • kg = k * 2^(g_C - g) : پس از نوشتن کلید، قبل از اینکه بتواند وارد حالت بخش بعدی شود، باید به انتهای بخش واپاشی کند.
  • 2^g_C : واپاشی کل حالت عبوری از یک تکه، که در آن g_C گیت تجمعی در انتهای تکه است.

رابطه‌ی بین‌قطعه‌ای متناظر به صورت زیر است:

\mathbf{h}^{[t+1]}
=2^{\mathbf{g}_C}\odot\mathbf{h}^{[t]}
+\mathrm{kg}^\top\mathbf{v}_{\text{new} }

۱.۳ رابط برنامه‌نویسی کاربردی (API) و مرجع متغیر

نقطه ورودی: tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu.py::PallasMosaicTpuKimiDeltaAttentionVjp._fwd(...) و به دنبال آن tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_kernel.py::chunk_kda_bwd_custom(...) .

همه تانسورهای مسیر اصلی از طرح‌بندی head-first [H, B, T, ...] استفاده می‌کنند. API عمومی KDA و backend پالاس از قبل از این طرح‌بندی استفاده می‌کنند؛ هیچ ترانهاده ورودی در مرز custom-VJP وجود ندارد. backend به جای تانسورهای اصلی بازپخش شده که توسط قرارداد عمومی VJP توکامکس ارائه می‌شوند، از کپی‌های هم‌تراز و به صورت اختیاری نرمال‌سازی شده L2 ذخیره شده در KdaResiduals استفاده می‌کند.

ورودی‌ها و مقادیر ذخیره‌شده‌ی رو به جلو:

متغیر شکل معنی
query, key /باقیمانده q, k [H, B, T, K] پرس و جو / کلید
value عمومی / ارزش باقیمانده v [H, B, T, V] ارزش
beta [H, B, T] قدرت نوشتن دلتا-روال به ازای هر توکن
gate عمومی [H, B, T, K] ورودی دروازه به ازای هر توکن قبل از cumsum محلی تکه‌ای
g_cumsum [H, B, T, K] / None گیت پس-کامسام داخلی در فضای log2؛ در صورت حذف، دوباره محاسبه می‌شود
g_org [H, B, T, K] / None گیت خام را هنگام اجرای فعال‌سازی در هسته حفظ می‌کند
a_log [H] / None پارامتر فعال‌سازی گیت
delta_time_bias [H*K] / None بایاس فعال‌سازی گیت اختیاری
Aqk [H, B, T, BT] ماتریس توجه درون بخشی ذخیره شده رو به جلو
Akk [H, B, T, BT] ماتریس معکوس WY با ذخیره رو به جلو
h [H, B, NT, K, V] حالت پنهان در شروع هر قطعه، ذخیره شده در مسیر سریع
do [H, B, T, V] گرادیان خروجی
initial_state [B, N, H, K, V] / None وضعیت اولیه تحت قرارداد عمومی KDA
dht [B, N, H, K, V] / None کتانژانت حالت نهایی تحت قرارداد عمومی KDA

کمیت‌های واسطه‌ای معکوس:

متغیر شکل معنی
u [H, B, T, V] مقدار مؤثر WY، فقط به صورت گذرا در داخل هسته محاسبه مجدد وجود دارد
w [H, B, T, K] وزن کلید/پاک کردن موثر WY
qg, kg [H, B, T, K] پرس‌وجو/کلید دروازه‌دار
v_new [H, B, T, V] ارزش پس از WY و اصلاح وضعیت تاریخی
dAqk [H, B, T, BT] گرادیان ماتریس توجه درون، از هسته dAv

خروجی‌ها:

متغیر شکل معنی
dq, dk, dg [H, B, T, K] گرادیان‌های ورودی-درگاه پرس‌وجو/کلید/عمومی
dv [H, B, T, V] گرادیان مقدار
db [H, B, T] گرادیان beta
dh0 [B, N, H, K, V] / None گرادیان حالت اولیه، مطابق با شکل ورودی عمومی
dA, dbias [H] / [H*K] / None گرادیان‌های a_log / delta_time_bias

محدودیت‌های استاتیک:

  • پیکربندی Mosaic ارائه شده BT=chunk_size=64 را ارائه می‌دهد.
  • T آماده‌شده باید بر BT بخش‌پذیر باشد؛ ورودی‌های با طول متغیر قبل از ساخت باقیمانده تراز می‌شوند؛
  • مسیر اصلی از گیت log2 استفاده می‌کند، یعنی use_exp2=True ;
  • اجرای غیر CP از قرارداد K<=256 تحویلی بک‌اند، شامل K/V غیر منطبق با ۱۲۸، پشتیبانی می‌کند؛ CP در حال حاضر ایجاب می‌کند که K و V هر دو مضربی از ۱۲۸ باشند.

هماهنگ‌کننده‌ی رو به عقب می‌تواند با وارد کردن یک محور N=1 ، یک کتانژانت حالت نهایی چهاربعدی داخلی را نرمال‌سازی کند، اما این یک مسیر سازگاری پیاده‌سازی است. قرارداد عمومی Tokamax KDA حالت‌های بازگشتی را در فرم پنج‌بعدی [B, N, H, K, V] می‌پذیرد و برمی‌گرداند. اجرای Pallas با طول ثابت به N=1 نیاز دارد.


۲. زنجیره فراخوانی مبتنی بر کد و ساختار هسته

زنجیره تماس فعلی Tokamax به شرح زیر است:

PallasMosaicTpuKimiDeltaAttentionVjp._fwd
└─ chunk_kda_bwd_custom
   ├─ unpack KdaResiduals
   ├─ align do with the retained cu_seqlens/aligned_cu_seqlens
   ├─ restore derived CP metadata retained by forward
   ├─ Stage 0: recover forward intermediate quantities
   │  ├─ if use_gate_in_kernel: kda_gate_chunk_cumsum(...)
   │  ├─ saved-state path rematerialize_for_backward=False:
   │  │  └─ fused_recompute_w_u_vnew_from_h_pallas(...)
   │  │     → w, qg, kg, v_new
   │  └─ low-memory path rematerialize_for_backward=True:
   │     ├─ _recompute_w_u_fwd(...)
   │     │  → w, u, qg, kg
   │     └─ chunk_gated_delta_rule_fwd_h(...)
   │        → h, v_new
   ├─ Stage 1: chunk_kda_bwd_dAv_kernel(...)
   │  → dAqk, dv
   ├─ optional context parallel:
   │  ├─ chunk_gated_delta_rule_bwd_dhu_pre_process(...)
   │  │  → dS_ext, dM
   │  ├─ all_gather_into_tensor(...)
   │  ├─ _merge_dht(...)
   │  └─ construct the dht used by the fusion kernel
   ├─ Stage 2-5: _fused_dhu_wy_intra_cumsum_pallas_jit(...)
   │  → dq, dk, dv, db, dg, dh0
   ├─ optional gate backward:
   │  └─ kda_gate_bwd(...)
   │     → dg, dA, dbias
   ├─ optional l2norm_bwd(...) for dq/dk
   ├─ _unalign_output(...) for varlen dq/dk/dv/dg/db
   └─ restore dh0 shape and input dtypes

از نظر ریاضی، حرکت رو به عقب را می‌توان به شش مرحله تجزیه کرد:

  1. مقادیر میانی رو به جلو را دوباره محاسبه کنید.
  2. گرادیان جمله‌ی درون‌جمله‌ای Aqk @ v_new را محاسبه کن؛
  3. گرادیان حالت dh را در امتداد محور قطعه به عقب تکرار کنید.
  4. نمایش WY را به عقب انتشار دهید؛
  5. توجه درون‌قطعه‌ای را به عقب منتشر کنید؛
  6. کامسام معکوس را روی شیب دروازه انجام دهید.

مسیر اصلی saved- h این شش مرحله را بر روی سه هسته پالاس نگاشت می‌کند:

recompute fusion        → w, qg, kg, v_new
dAv                     → dAqk, dv
dhu/WY/intra/cumsum     → dq, dk, dv, db, dg, dh0

مسیر کم حافظه انتخاب شده توسط rematerialize_for_backward=True علاوه بر این _recompute_w_u_fwd و chunk_gated_delta_rule_fwd_h را برای اجرای مجدد حالت رو به جلو فراخوانی می‌کند. آماده‌سازی موازی-متن، گیت رو به عقب، نرمال‌سازی L2 رو به عقب و عدم هم‌ترازی varlen عملیات اختیاری در اطراف هسته‌های اصلی هستند.

۲.۱ محاسبه مجدد هسته فیوژن

نقطه ورود: fused_recompute_w_u_vnew_from_h_pallas(...) .

این هسته h ذخیره شده و Akk ذخیره شده رو به جلو را می‌خواند و موارد زیر را به صورت موازی در هر قطعه دوباره محاسبه می‌کند:

u     = Akk @ (v * beta)
w     = Akk @ (k * beta * 2^g)
v_new = u - w @ h
qg    = q * 2^g
kg    = k * 2^(g_C - g)

u فقط درون هسته وجود دارد و پس از محاسبه v_new در HBM نوشته نمی‌شود. از آنجایی که h از قبل ذخیره شده است، v_new دیگر به تکرار حالت رو به جلو بین تکه‌ای وابسته نیست، بنابراین هر تکه می‌تواند به طور مستقل زمان‌بندی شود.

مسیر کم‌حافظه زمانی استفاده می‌شود که h حفظ نشده باشد. ابتدا w, u, qg, kg از طریق _recompute_w_u_fwd(...) بازیابی می‌کند، سپس chunk_gated_delta_rule_fwd_h(...) را برای اجرای مجدد حالت و بدست آوردن h, v_new فراخوانی می‌کند. این مسیر وابستگی سریالی بین تکه‌ها را دوباره معرفی می‌کند و محاسبه را با حافظه باقیمانده پایین‌تر رو به جلو عوض می‌کند.

هسته dAv 2.2

نقطه ورود: chunk_kda_bwd_dAv_kernel(...) .

این هسته فقط خروجی داخلی را مدیریت می‌کند:

\mathrm{tril}(\mathbf{A}_{qk})\,\mathbf{v}_{\text{new} }

پس‌رانش مربوطه عبارت است از:

d\mathbf{A}_{qk}
=s\,\mathrm{tril}(d\mathbf{o}\,\mathbf{v}_{\text{new} }^\top)
d\mathbf{v}_{\text{new} }
=\mathrm{tril}(\mathbf{A}_{qk})^\top d\mathbf{o}

در کد، scale فقط یک بار، هنگام تولید dAqk ، ضرب می‌شود. هسته ادغام بعدی مستقیماً dAqk از قبل مقیاس‌بندی شده را مصرف می‌کند و از مقیاس‌بندی مکرر جلوگیری می‌کند.

Aqk ذخیره شده از قبل شامل scale است، بنابراین شاخه value-gradient Aqk^T @ do بدون عامل دیگری استفاده می‌کند. تانسوری به نام dAqk قبل از fused intra backward مقیاس‌بندی می‌شود زیرا آن هسته رابطه query-key مقیاس‌بندی نشده زیرین را مستقیماً متمایز می‌کند. اگرچه chunk_kda_bwd_dAv_kernel q و k در امضای عمومی خود حفظ می‌کند، اما لانچر و هسته فعلی فقط از پارامترهای v ، Aqk ، do ، scale و tiling استفاده می‌کنند.

2.3 dhu/WY/intra/cumsum Fusion Kernel

نقطه ورود: _fused_dhu_wy_intra_cumsum_pallas_jit(...) .

این هسته به ترتیب قطعات معکوس اجرا می‌شود و چهار نوع کار را با هم ترکیب می‌کند:

  1. بازگشت رو به عقب dhu گرادیان حالت بین تکه‌ای dh را حفظ می‌کند و سهم مسیر خروجی، مسیر به‌روزرسانی دلتا و مسیر واپاشی تکه را در همان خراش VMEM جمع‌آوری می‌کند.

  2. الگوریتم WY به عقب از v_new = u - w @ h ، u = Akk @ (v * beta) ، w = Akk @ (...) به q, k, v, beta, g می‌یابد و سهم گرادیان را تا Akk به دست می‌آورد.

  3. درون-رو به عقب dAqk تولید شده توسط هسته dAv را مصرف می‌کند، به انتشار معکوس گرادیان ماتریس توجه درون-بخشی به q, k, beta, g ادامه می‌دهد و با نتایج WY جمع می‌شود.

  4. گیت رو به جلو یک کامسام محلی-قطعه‌ای است، بنابراین پاس رو به عقب به یک کامسام معکوس نیاز دارد. پیاده‌سازی فعلی این را در همان هسته نگه می‌دارد و دیگر یک هسته کامسام جداگانه راه‌اندازی نمی‌کند.

محور تکه‌ای از معنای شبکه‌ای arbitrary برای حمل حالت بازگشتی رو به عقب استفاده می‌کند؛ محورهای سر و دسته‌ای ابعاد موازی دارند.

۲.۴ موازی‌سازی زمینه اختیاری

وقتی context parallel فعال باشد، مسیر رو به عقب، ادغام گرادیان حالت متقاطع را بین هسته‌های dAv و فیوژن وارد می‌کند.

جریان این است:

  1. chunk_gated_delta_rule_bwd_dhu_pre_process(...) فقط اولین سگمنت محلی واقعی را اسکن می‌کند، زیرا این سگمنتی است که می‌تواند از یک رتبه بالادستی، وضعیت دریافت کند و نتیجه زیر را تولید می‌کند:
    • dS_ext : سهم خارجی این رتبه در گرادیان حالت ورودی؛
    • dM : حاصلضرب زنجیره‌ای ماتریس انتقال رو به عقب این رتبه.
  2. dS_ext و dM را در امتداد آخرین بُعد بسته‌بندی کن و آنها را با یک all_gather_into_tensor(...) جایگزین کن.
  3. _merge_dht(...) با استفاده از فراداده‌های post_num_ranks و is_last_rank که توسط تابع forward نگهداری می‌شوند، رتبه‌های پایین‌دستی را به ترتیب از دورترین به نزدیکترین رتبه ادغام می‌کند.
  4. یک dht [B,N,H,K,V] بسازید و گرادیان حالت ادغام‌شده‌ی هر عنصر دسته‌ای را قبل از ورود به هسته‌ی پس‌رونده‌ی ادغام‌شده، در آخرین جایگاه قطعه‌ی محلی حقیقی آن قرار دهید.

جهت رو به عقب CP از رتبه پایین‌دستی به رتبه بالادستی می‌رود که توسط segment_ids تراز شده و ابرداده رتبه بازیابی شده در ContextParallelMetadata تعیین می‌شود. B>1 به طور مستقل برای هر عنصر دسته‌ای مدیریت می‌شود. قرارداد CP یک initial_state خارجی را نمی‌پذیرد یا یک گرادیان حالت اولیه را برنمی‌گرداند.

۲.۵ دروازه اختیاری رو به عقب

وقتی use_gate_in_kernel=True ، chunk_kda_bwd_custom(...) تابع kda_gate_bwd(...) را بعد از هسته به صورت برعکس فراخوانی می‌کند.

در این مرحله، هسته فیوژن قبلاً cumsum معکوس محلی-قطعه‌ای را اعمال کرده است، بنابراین dg آن گرادیان نسبت به گیت per-token فعال شده قبل از cumsum است. سپس kda_gate_bwd(...) فعال‌سازی را متمایز می‌کند و این گرادیان را به gate عمومی خام، a_log و delta_time_bias اختیاری نگاشت می‌کند و به ترتیب dg ، dA و dbias را برمی‌گرداند.


۳. سازگاری TPU

هدف از این انتخاب‌های پیاده‌سازی، همسو کردن هرچه بیشتر مسیر رو به عقب با مدل اجرای MXU/VPU/HBM مربوط به TPU است.

۳.۱ گیت log2 و exp2

گیت تجمعی مصرف شده توسط هسته‌های درون و حالت-بازگشتی در فضای log2 نمایش داده می‌شود، بنابراین واپاشی‌ها به طور یکنواخت به صورت 2^g نوشته می‌شوند. ورودی گیت عمومی توسط مرحله گیت/کامسام رو به جلو به این نمایش تبدیل می‌شود؛ نباید آن را با مقدار تجمعی داخلی اشتباه گرفت.

این کار دو فایده دارد:

  • واپاشی‌های درون‌قطعه‌ای و بین‌قطعه‌ای را می‌توان از طریق جمع و تفریق در توان ترکیب کرد؛
  • هسته از سخت‌افزار exp2 استفاده می‌کند و از تغییر مکرر پایه exp جلوگیری می‌کند.

۳.۲ ذخیره‌سازی fp32 و bf16

ورودی‌ها معمولاً bf16 هستند، اما matmul ها از fp32 accumulation استفاده می‌کنند.

در داخل هسته اصلی فیوژن، ورودی‌های fp32، jax.lax.Precision.HIGHEST را انتخاب می‌کنند، در حالی که ورودی‌های bf16 از دقت نقطه پیش‌فرض با انباشت fp32 استفاده می‌کنند. نقطه جمع معکوس به صراحت از jax.lax.Precision.HIGHEST برای هر دو نوع ورودی d استفاده می‌کند. این انتخاب‌ها برای dg اهمیت دارند، جایی که سهم‌های تجمعی گیت تقریباً می‌توانند خنثی شوند.

۳.۳ VMEM Scratch واسطه‌های بین مرحله‌ای را حفظ می‌کند

هسته فیوژن، حالت معکوس dh را در یک خراش VMEM [MB, K, V] قرار می‌دهد و آن را به ترتیب معکوس به‌روزرسانی می‌کند.

در عین حال، dv_new ، گرادیان‌های میانی WY، گرادیان‌های درون میانی و نتایج محلی مورد نیاز cumsum معکوس، همگی تا حد امکان در VMEM در گردش باقی می‌مانند. این امر از نوشتن و خواندن مکرر HBM بین مراحل منطقی dhu، WY، درون و cumsum جلوگیری می‌کند.

۳.۴ کنترل‌های مینی‌بچ DMA دانه‌بندی

یک برنامه پالاس، کاشی‌های head/chunk MB را به طور همزمان پردازش می‌کند.

MB از یک ردپای تخمینی VMEM انتخاب می‌شود، اما پیاده‌سازی دقیق آن بسته به هسته متفاوت است. chunk_kda_bwd_dAv_kernel از کمک‌کننده مشترک estimate_mini_batch(...) استفاده می‌کند. هسته محاسبه مجدد saved- h ، هسته اصلی فیوژن و پیش‌پردازش CP در حال حاضر از روش‌های اکتشافی محلی VMEM-budget با محدودیت‌های خاص هسته و تنظیمات تقسیم‌پذیری/هم‌ترازی استفاده می‌کنند.

  • خیلی کوچک منجر به دانه‌بندی ناکافی DMA و استفاده کم از پهنای باند HBM می‌شود.
  • اگر خیلی بزرگ باشد، فشار VMEM و رجیستر و سربار ناشی از باز شدن استاتیک افزایش می‌یابد.

بنابراین هر لانچر، مساحت اشغال شده توسط کاشی‌های خود را تخمین می‌زند و یک MB با توجه به الزامات شبکه، تقسیم‌پذیری و ترازبندی ابعاد جزئی TPU خود انتخاب می‌کند. لانچر محاسبه مجدد saved h می‌تواند تعداد قطعات مسطح شده خود را زمانی که هیچ مقسوم علیه مناسبی در دسترس نیست، افزایش دهد. لانچرهای اصلی فیوژن و CP، MB تا زمانی که تعداد سرها را تقسیم کند، کاهش می‌دهند.

۳.۵ جابجایی داده‌های BlockSpec و CP DMA صریح

سه هسته‌ی اصلی saved h از شبکه‌های Pallas و اشیاء BlockSpec برای توصیف کاشی‌های HBM استفاده می‌کنند. لانچرهای آنها حاوی یک حلقه‌ی کپی ناهمزمانِ دست‌نویسِ صریح نیستند.

پیش‌پردازش CP متفاوت است: _chunk_gated_delta_rule_bwd_dhu_pre_process_kernel پیاده‌سازی فعال CP است و به صراحت از ورودی‌های pltpu.make_async_copy ، سمافورهای DMA و ورودی‌های double-buffered-VMEM استفاده می‌کند. این یک اسکن معکوس تک برنامه‌ای روی گروه‌های هد و تکه‌ها است که انتقال ورودی بعدی را با کار ماتریس فعلی همپوشانی می‌دهد و خلاصه dS_ext/dM هر گروه هد را به صورت غیرهمزمان می‌نویسد.

۳.۶ اصول کاشی‌کاری و لایه‌گذاری

کاشی‌کاری TPU ابعاد انتهایی خاصی را با مرزهای کاشی سخت‌افزاری هم‌تراز می‌کند. اگرچه عناصر لایه‌گذاری‌شده از نظر معنایی صفر هستند، اما همچنان پهنای باند واقعی را در سطح HBM و DMA مصرف می‌کنند.

پیاده‌سازی فعلی از یک اصل عملی پیروی می‌کند: تنها زمانی یک چیدمان فشرده‌شده‌ی پیچیده‌تر معرفی کنید که ترافیک HBM که تغییر چیدمان می‌تواند حذف کند، از هزینه‌ی اضافی تغییر شکل/جمع‌آوری/کپی بیشتر شود.

بنابراین:

  • غیر CP به صورت معکوس از موارد K/V غیر هم‌تراز تحویل داده شده بدون اعمال محدودیت خط CP پشتیبانی می‌کند، در حالی که پیش‌پردازش CP مستلزم آن است که K و V مضربی از ۱۲۸ باشند؛
  • [BT, BT] ماتریس‌های کوچکی مانند Aqk و Akk ، فاصله‌گذاری ثابت ناشی از هم‌ترازی سخت‌افزاری را می‌پذیرند؛
  • ابعاد دنباله‌دار اسکالر مانند beta/db از طرح‌بندی‌های صریح تک‌بعدی یا دوبعدی انتخاب‌شده توسط هر لانچر استفاده می‌کنند.

به عبارت دیگر، خودِ لایه‌بندی لزوماً چیز بدی نیست که باید حذف شود. تا زمانی که هزینه ایندکس‌گذاری و تبدیل طرح‌بندی که با حذف لایه‌بندی ایجاد می‌شود بیشتر باشد، حفظ طرح‌بندی ساده در واقع سریع‌تر است.

۳.۷ نرمال‌سازی زیربلوک در داخل رو به عقب

معکوس درون شامل عبارات واپاشی مانند 2^(g_r - g_j) است. از آنجایی که g یک کمیت تجمعی است، تفریق مستقیم و سپس توان‌بندی ممکن است باعث ایجاد مشکلات مربوط به محدوده عددی شود.

این پیاده‌سازی، قطعه را به زیربلوک‌های کوچک‌تر تقسیم می‌کند و در هر زیربلوک یک گیت مرجع برای نرمال‌سازی انتخاب می‌کند. به این ترتیب، عبارت توان نسبت به نقطه مرجع به دو قسمت تقسیم می‌شود و مقادیر میانی در محدوده قابل کنترل‌تری نگه داشته می‌شوند.

همین ساختار زیربلوک، سازماندهی بلوک‌های مورب و غیر مورب را در تعداد ثابتی از متمول‌های دسته‌ای نیز راحت می‌کند و از نوشتن یک حلقه درجه دوم روی زیربلوک‌ها جلوگیری می‌کند.

۳.۸ شکل‌های ثابت نگهدارنده‌ی وارلن

هسته‌های TPU به شکل‌های ایستا نیاز دارند. توالی‌های با طول متغیر، یک هسته ناهموار جداگانه را فعال نمی‌کنند، اما همچنان از تانسورهای ثابت [B, T] و یک شبکه ثابت (H // MB, B, NT) استفاده می‌کنند.

آداپتور cu_seqlens مشتق می‌کند، هر بخش منطقی را با BT تراز می‌کند، و هر دو فراداده اصلی و تراز شده را در KdaResiduals حفظ می‌کند. ابتدا do را به عقب تراز می‌کند، سپس از segment_ids تراز شده و نگاشت تکه‌ای حفظ شده استفاده می‌کند:

  • تکه‌هایی با شناسه قطعه ۰ در حال پر کردن هستند؛
  • آخرین بخش هر قطعه حقیقی، تابع بازگشتی رو به عقب را با dht آغاز می‌کند؛
  • اولین بخش از هر بخش حقیقی، dh0 را خروجی می‌دهد.

فقط برای مثال، B=1, T=12, BT=4 را در نظر بگیرید که دو دنباله با طول‌های ۸ و ۴ را در خود جای داده‌اند. آداپتور ارائه شده همچنان به BT=64 نیاز دارد؛ مقدار کمتر فقط باعث فشرده‌تر شدن مثال فراداده می‌شود.

token segment_ids = [1,1,1,1, 1,1,1,1, 2,2,2,2]
chunk_seg_ids     = [   1,        1,        2   ]
                       chunk0     chunk1    chunk2

برای اجرای با طول ثابت، _fused_dhu_wy_intra_cumsum_pallas_jit یک تانسور شناسه قطعه [B,T] با تمام یک‌ها را ترکیب می‌کند، بنابراین هر عنصر دسته‌ای به عنوان یک قطعه در نظر گرفته می‌شود. پس از هسته به عقب، گرادیان‌های توکن با طول متغیر با طرح اصلی [B,T_original] تراز نمی‌شوند.


۴. تحلیل عملکرد فیوژن

اندازه‌گیری‌های این بخش به عنوان شواهد بهینه‌سازی تاریخی حفظ می‌شوند. آن‌ها قبل از آداپتور فعلی Tokamax، قرارداد هم‌ترازی/باقیمانده متمرکز، انتقال فراداده CP فعلی و برخی تغییرات نامگذاری در سطح راه‌اندازی، وجود داشته‌اند. آن‌ها باید قبل از استفاده به عنوان آستانه پذیرش یا خط مبنای رگرسیون، در کامیت فعلی دوباره اجرا شوند.

پیکربندی اندازه‌گیری تاریخی: TPU v6e، bf16، varlen seq_lens=[1800,1500,2000,1200,800] ، H=16, B=1, T=8192, K=V=128, BT=64 . پس از قطعه‌بندی، قطر اصلی 16 × 133 است.

۴.۱ چرا فیوژن مسیر اصلی بهینه‌سازی است

بیشتر مراحل KDA معکوس، شدت حسابی پایینی دارند و وقتی به صورت جداگانه تقسیم می‌شوند، به راحتی به پهنای باند HBM محدود می‌شوند.

اگر هر مرحله ریاضی به یک هسته مستقل تبدیل شود، تعداد زیادی از تانسورهای میانی به HBM رفت و برگشت می‌کنند، به عنوان مثال:

  • w, qg, kg, v_new ؛
  • dAqk, dv ؛
  • dh, dv_new ؛
  • گرادیان‌های WY و درون‌میانجی؛
  • ورودی و خروجی کامسام معکوس.

پیاده‌سازی فعلی، این مراحل را در سه هسته پالاس فشرده می‌کند. مزیت اصلی کاهش عملیات ریاضی نیست، بلکه کاهش خواندن و نوشتن در مرزهای مراحل است.

۴.۲ اندازه‌گیری‌های تاریخی مسیر اصلی سه‌گانه

مسیر اصلی از سه بخش تشکیل شده است:

اندازه‌گیری: TPU v6e ، bf16، varlen seq_lens=[1800,1500,2000,1200,800] (N=5 بخش)، H=16، B=1، T=8192، K=V=128، BT=64؛ پس از قطعه‌بندی dim پیشرو = 2128 = 16×133 (NT=133 = 128 بخش پایه + 5 padding بخش). پهنای باند HBM برابر با 1.6 ترابایت بر ثانیه در نظر گرفته شده است.

هسته زمان اجرا (میکرو ثانیه) اچ‌بی‌ام (MB) حد پایین وزن مخصوص (میکرو ثانیه) استفاده از وزن بدن
fused_recompute_w_u_vnew_from_h_pallas ۲۸۵.۰ ۴۱۸.۹ ۲۶۱.۸۳ ۹۱.۹٪
chunk_kda_bwd_dAv_kernel ۱۴۸.۳ ۲۰۹.۲ ۱۳۰.۷۴ ۸۸.۲٪
_fused_dhu_wy_intra_cumsum_pallas_jit ۹۰۰.۱ ۹۵۱.۸ ۵۹۴.۹ ۶۶.۱٪

از نظر تاریخی، دو هسته اول نشان دادند که یک ساختار محاسباتی ساده و الگوی خواندن/نوشتن مستقیم می‌تواند از v6e HBM به خوبی برای این حجم کاری استفاده کند.

در آن اندازه‌گیری‌ها، آخرین هسته فیوژن هدف اصلی بهینه‌سازی بود. اگرچه تعداد زیادی از رفت و برگشت‌های HBM را حذف می‌کند، اما از نظر داخلی شامل بازگشت حالت، چندین عملیات ماتریس درون قطعه‌ای، کاهش گرادیان گیت و جمع معکوس است، بنابراین سربار محاسباتی، خط لوله و طرح‌بندی محلی را با هم مخلوط می‌کند.

اندازه‌گیری: TPU v7x ، bf16، varlen seq_lens=[1800,1500,2000,1200,800] (N=5 segment)، H=16، B=1، T=8192، K=V=128، BT=64؛ پس از قطعه‌بندی dim پیشرو = 2128 = 16×133 (NT=133 = 128 قطعه پایه + 5 padding قطعه). پهنای باند HBM برابر با 3.69 ترابایت بر ثانیه در نظر گرفته شده است.

هسته زمان اجرا (میکرو ثانیه) اچ‌بی‌ام (MB) حد پایین وزن مخصوص (میکرو ثانیه) استفاده از وزن بدن
fused_recompute_w_u_vnew_from_h_pallas ۱۸۱.۶۷۵ ۴۱۸.۹ ۱۱۳.۵۲ ۶۲.۵٪
chunk_kda_bwd_dAv_kernel ۱۰۰.۲۳۷ ۲۰۹.۲ ۵۶.۶۹ ۵۶.۶٪
_fused_dhu_wy_intra_cumsum_pallas_jit ۷۸۴.۲۹۹ ۹۵۱.۸ ۲۵۷.۹۴ ۳۲.۹٪

۴.۳ ویژگی‌های عملکرد Context Parallel

اندازه‌گیری: TPU v6e-4 (cp_size=4، موازی‌سازی SPMD چهار هسته‌ای)، bf16، تک‌قطعه per_rank_T=8192 → سراسری T=32768، H=16، K=V=128، BT=64؛ کم‌نور شدن پیشرو به ازای هر رتبه پس از قطعه‌بندی = 2128 = 16×133.

هسته‌های تکی (مبنای یکسان برای padding):

هسته زمان اجرا (میکرو ثانیه) اچ‌بی‌ام (MB) حد پایین وزن مخصوص (میکرو ثانیه) استفاده از وزن بدن
fused_recompute_w_u_vnew_from_h_pallas ۲۷۳.۴ ۴۰۶.۳ ۲۵۳.۹۵ ۹۲.۹٪
chunk_kda_bwd_dAv_kernel ۱۴۶.۸ ۲۰۲.۹ ۱۲۶.۸۱ ۸۶.۴٪
_fused_dhu_wy_intra_cumsum_pallas_jit ۸۷۲.۷ ۹۱۵.۱ ۵۷۱.۹۷ ۶۵.۵٪
chunk_gated_delta_rule_bwd_dhu_pre_process CP ۵۹۱.۲ ۲۳۸.۸ ۱۴۹.۲۶ ۲۵.۲٪

اندازه‌گیری: TPU v7x-4 (cp_size=4، موازی‌سازی SPMD چهار هسته‌ای)، bf16، تک‌قطعه per_rank_T=8192 → سراسری T=32768، H=16، K=V=128، BT=64؛ کم‌نور شدن پیشرو به ازای هر رتبه پس از قطعه‌بندی = 2128 = 16×133.

هسته زمان اجرا (میکرو ثانیه) اچ‌بی‌ام (MB) حد پایین وزن مخصوص (میکرو ثانیه) استفاده از وزن بدن
fused_recompute_w_u_vnew_from_h_pallas ۱۸۱ ۴۰۶.۳ ۱۱۰.۱۱ ۶۰.۸٪
chunk_kda_bwd_dAv_kernel ۱۰۰ ۲۰۲.۹ ۵۴.۹۹ ۵۵.۰٪
_fused_dhu_wy_intra_cumsum_pallas_jit ۷۸۵ عدد ۹۱۵.۱ ۲۴۷.۹۹ ۳۱.۶٪
chunk_gated_delta_rule_bwd_dhu_pre_process CP ۵۱۱ ۲۳۸.۸ ۶۴.۷۲ ۱۲.۷٪

مسیر CP اندازه‌گیری شده از همان سه هسته اصلی مسیر غیر CP استفاده کرد؛ کار اضافی آن از ادغام گرادیان حالت بین رتبه‌ای حاصل شد.

در نمایه تاریخی، هزینه اصلی CP خودِ مجموعه نبود، بلکه پیش‌فرآیند محلی قبل از ارتباط بود:

  • پیش‌پردازش نیاز به جمع‌آوری گرادیان حالت ورودی و ماتریس انتقال رو به عقب در امتداد توالی قطعات این رتبه دارد؛
  • این مرحله ماهیت اسکن/کاهش دارد و موازی‌سازی کامل آن مانند یک متمول معمولی برای هر تکه دشوار است.
  • خودِ all_gather نسبتاً ارزان است، زیرا آنچه مبادله می‌شود گرادیان حالت فشرده و ماتریس انتقال است، نه تانسورهای توکن توالی کامل.

بنابراین، تمرکز بهینه‌سازی مسیر CP باید روی پیش‌پردازش باشد:

  • خواندن‌های مکرر ورودی کل قطعه را کاهش دهد؛
  • بهبود موازی‌سازی زنجیره محصول/اسکن؛
  • یا بخشی از منطق ادغام حالت را به یک مرحله عقب مانده موجود منتقل/ترکیب کنید.

۴.۴ سازگاری عددی

پیاده‌سازی فعلی compute_intra_backward(...) به صراحت dg_acc + dg_intra را به dtype مرجع ورودی تبدیل می‌کند و سپس قبل از compute_reverse_cumsum_dg(...) به fp32 برمی‌گرداند. این مرز dtype بخشی از رفتار عددی تحویل داده شده است و باید هنگام بازسازی فیوژن حفظ یا عمداً دوباره اعتبارسنجی شود. مقایسه‌ها در برابر یک مرجع فیوژن نشده یا XLA باید از تلرانس‌های مناسب برای این نقطه برش استفاده کنند.

۴.۵ ماتریس اعتبارسنجی مجدد برای پیاده‌سازی فعلی

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

  • مسیر saved- h ( rematerialize_for_backward=False ) و مسیر full-rematerialization ( rematerialize_for_backward=True
  • ورودی‌های با طول ثابت و متغیر، شامل چندین بخش و B>1 ؛
  • یک initial_state ، output_final_state=True و یک کتانژانت حالت نهایی غیر صفر ارائه شده است؛
  • گیت‌های از پیش محاسبه‌شده و use_gate_in_kernel=True ، با و بدون delta_time_bias ؛
  • use_qk_l2norm=True ؛
  • اشکال K/V غیر هم‌تراز با CP که توسط قرارداد backend مجاز هستند؛
  • اجرای CP در شکل‌های K/V با تراز ۱۲۸ و اندازه‌های CP چندگانه پشتیبانی‌شده؛
  • ورودی‌های bf16 و مسیر fp32 پشتیبانی‌شده.

اعتبارسنجی مجدد عملکرد باید گزارش دقیقی از کامیت، تولید و توپولوژی TPU، نسخه‌های نرم‌افزار، شکل‌ها، توزیع varlen، اندازه CP، سیاست گرم شدن، تعداد تکرار و اینکه آیا هر عدد مربوط به هر دستگاه یا تأخیر انتها به انتها است، ارائه دهد. تخمین‌های پهنای باند در سطح هسته باید با یک اندازه‌گیری معکوس انتها به انتها همراه باشد تا سربارهای راه‌اندازی، جمعی و آداپتور قابل مشاهده باشند.


۵. دستورالعمل‌های بهینه‌سازی آینده

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

۵.۱ بهبود عملکرد نسخه ۷

در اندازه‌گیری‌های تاریخی، هسته‌های اصلی در v7x سریع‌تر اجرا شدند اما میزان استفاده از پهنای باند تخمینی کمتری نسبت به v6e داشتند. یک نمایه فعلی ابتدا باید این نتیجه را تأیید کند؛ اگر این نتیجه همچنان درست باشد، کار بعدی باید بر روی تطبیق بهتر هسته‌های ادغام‌شده با مدل اجرای v7 متمرکز شود:

  • اندازه کاشی‌ها، انتخاب مینی بچ و فشار VMEM برای نسخه ۷ را دوباره بررسی کنید؛
  • کاهش سربار طرح‌بندی محلی و ترافیک لایه‌بندی در جایی که به وضوح در پروفایل‌ها قابل مشاهده است؛
  • مشخص کنید کدام بخش‌های هسته اصلی فیوژن به جای اینکه محدود به پهنای باند باشند، محدود به محاسبات یا خط لوله هستند.

هدف تغییر تجزیه ریاضی نیست، بلکه تنظیم مجدد مسیر ترکیبی موجود است تا با پهنای باند و قابلیت محاسباتی بالاتر v7، مقیاس‌پذیری بهتری داشته باشد.

۵.۲ بهینه‌سازی موازی متن

مسیر CP قبل از تبادل رتبه‌بندی متقابل، یک پیش‌پردازش محلی اضافه می‌کند. نمایه تاریخی نشان می‌دهد که این کار اسکن/کاهش محلی، گلوگاه بزرگ‌تری نسبت به خود ارتباط جمعی بوده است؛ نمایه فعلی باید قبل از شروع کار بهینه‌سازی، تعادل را تأیید کند.

بنابراین، بهینه‌سازی CP در آینده باید موارد زیر را در اولویت قرار دهد:

  • کاهش خواندن‌های مکرر در پیش‌پردازش؛
  • بهبود موازی‌سازی حاصلضرب زنجیره ماتریس انتقال؛
  • بررسی اینکه آیا بخشی از آماده‌سازی گرادیان حالت CP می‌تواند در هسته‌های معکوس موجود ادغام شود یا خیر.

الگوی ارتباطی باید فشرده باقی بماند: خلاصه‌های سطح حالت را به جای تنسورهای کامل سطح توکن مبادله کنید.