مرجع: گزارش فنی 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
از نظر ریاضی، حرکت رو به عقب را میتوان به شش مرحله تجزیه کرد:
- مقادیر میانی رو به جلو را دوباره محاسبه کنید.
- گرادیان جملهی درونجملهای
Aqk @ v_newرا محاسبه کن؛ - گرادیان حالت
dhرا در امتداد محور قطعه به عقب تکرار کنید. - نمایش WY را به عقب انتشار دهید؛
- توجه درونقطعهای را به عقب منتشر کنید؛
- کامسام معکوس را روی شیب دروازه انجام دهید.
مسیر اصلی 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(...) .
این هسته به ترتیب قطعات معکوس اجرا میشود و چهار نوع کار را با هم ترکیب میکند:
بازگشت رو به عقب
dhuگرادیان حالت بین تکهایdhرا حفظ میکند و سهم مسیر خروجی، مسیر بهروزرسانی دلتا و مسیر واپاشی تکه را در همان خراش VMEM جمعآوری میکند.الگوریتم WY به عقب از
v_new = u - w @ h،u = Akk @ (v * beta)،w = Akk @ (...)بهq, k, v, beta, gمییابد و سهم گرادیان را تاAkkبه دست میآورد.درون-رو به عقب
dAqkتولید شده توسط هسته dAv را مصرف میکند، به انتشار معکوس گرادیان ماتریس توجه درون-بخشی بهq, k, beta, gادامه میدهد و با نتایج WY جمع میشود.گیت رو به جلو یک کامسام محلی-قطعهای است، بنابراین پاس رو به عقب به یک کامسام معکوس نیاز دارد. پیادهسازی فعلی این را در همان هسته نگه میدارد و دیگر یک هسته کامسام جداگانه راهاندازی نمیکند.
محور تکهای از معنای شبکهای arbitrary برای حمل حالت بازگشتی رو به عقب استفاده میکند؛ محورهای سر و دستهای ابعاد موازی دارند.
۲.۴ موازیسازی زمینه اختیاری
وقتی context parallel فعال باشد، مسیر رو به عقب، ادغام گرادیان حالت متقاطع را بین هستههای dAv و فیوژن وارد میکند.
جریان این است:
-
chunk_gated_delta_rule_bwd_dhu_pre_process(...)فقط اولین سگمنت محلی واقعی را اسکن میکند، زیرا این سگمنتی است که میتواند از یک رتبه بالادستی، وضعیت دریافت کند و نتیجه زیر را تولید میکند:-
dS_ext: سهم خارجی این رتبه در گرادیان حالت ورودی؛ -
dM: حاصلضرب زنجیرهای ماتریس انتقال رو به عقب این رتبه.
-
-
dS_extوdMرا در امتداد آخرین بُعد بستهبندی کن و آنها را با یکall_gather_into_tensor(...)جایگزین کن. -
_merge_dht(...)با استفاده از فرادادههایpost_num_ranksوis_last_rankکه توسط تابع forward نگهداری میشوند، رتبههای پاییندستی را به ترتیب از دورترین به نزدیکترین رتبه ادغام میکند. - یک
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 میتواند در هستههای معکوس موجود ادغام شود یا خیر.
الگوی ارتباطی باید فشرده باقی بماند: خلاصههای سطح حالت را به جای تنسورهای کامل سطح توکن مبادله کنید.