KDA চাঙ্ক ব্যাকওয়ার্ড — TPU/প্যালাস কার্নেল ডিজাইন

সূত্র: কিমি লিনিয়ার প্রযুক্তি প্রতিবেদন — 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 — টোকাম্যাক্স কাস্টম-ভিজেপি অ্যাডাপ্টার
  • 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 — শেয়ার্ড গেট কামসাম, স্টেট রিকারেন্স, এবং মিনি-ব্যাচ হেল্পার
  • 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 ইনপুটের সাপেক্ষে গ্রেডিয়েন্ট। ফিউজড কার্নেলের রিভার্স কামসাম প্রথমে কামসামের আগে প্রতি-টোকেন সক্রিয় গেটের গ্রেডিয়েন্ট তৈরি করে; যখন 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] ; যখন rematerialize_for_backward=False , তখন ফরওয়ার্ড এটিকে KdaResiduals এ ধরে রাখে।

পশ্চাৎমুখী পুনঃগণনা:

  • w : WY কার্যকরী কী / ইরেজ ওয়েট;
  • qg : গেটেড কোয়েরি;
  • kg : গেটেড কী;
  • v_new : WY-সংশোধিত মান।

ঐচ্ছিক পুনঃগণনা:

  • যখন use_gate_in_kernel=True , তখন উভয় ব্যাকওয়ার্ড প্রিপারেশন পাথই বর্তমানে kda_gate_chunk_cumsum এর মাধ্যমে g_org , a_log এবং ঐচ্ছিক delta_time_bias থেকে পোস্ট-কামসাম গেটটি পুনরায় গণনা করে, যার মধ্যে সেভড- h পাথটিও অন্তর্ভুক্ত, যেখানে ফরোয়ার্ড রেসিড্যুয়ালে বর্তমানে g_cumsum ও থাকে।
  • যখন rematerialize_for_backward=True , তখন ব্যাকওয়ার্ড প্রক্রিয়াটি সংরক্ষিত h পথটি অনুসরণ করে না, বরং h এবং v_new পুনরুদ্ধার করার জন্য ফরোয়ার্ড স্টেট রিকারেন্সটি পুনরায় চালায়।

এটি মোজাইক কাস্টম ভিজেপি-এর মধ্যেকার ম্যানুয়াল রিম্যাটেরিয়ালাইজেশন, 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 উপর নির্ভরশীল, এবং uw গঠন করতে ব্যবহৃত মান এবং গেটেড-কী-র ডান-পার্শ্বেও beta স্পষ্টভাবে গুণ করা থাকে। সুতরাং, beta যে উভয় পথেই ফলাফলকে প্রভাবিত করে, ব্যাকওয়ার্ড পদ্ধতিকে অবশ্যই সেই উভয় পথই বিবেচনা করতে হবে।

নীচের সমীকরণগুলিতে, g লগ২ স্পেসে চাঙ্ক-লোকাল কিউমুলেটিভ গেটকে বোঝায়। এটি একটি অভ্যন্তরীণ রাশি, কাঁচা পাবলিক গেট আর্গুমেন্ট নয়। অতএব সমস্ত ক্ষয়কে 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} }

১.৩ এপিআই এবং ভেরিয়েবল রেফারেন্স

এন্ট্রি পয়েন্ট: 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(...)

সমস্ত মেইন-পাথ টেনসর হেড-ফার্স্ট [H, B, T, ...] লেআউট ব্যবহার করে। পাবলিক KDA API এবং Pallas ব্যাকএন্ড ইতিমধ্যেই এই লেআউটটি ব্যবহার করে; কাস্টম-VJP সীমানায় কোনো ইনপুট ট্রান্সপোজ নেই। ব্যাকওয়ার্ড প্রক্রিয়াটি Tokamax-এর জেনেরিক VJP কন্ট্রাক্ট দ্বারা সরবরাহকৃত রিপ্লে করা মূল টেনসরের পরিবর্তে KdaResiduals এ সংরক্ষিত অ্যালাইনড এবং ঐচ্ছিকভাবে L2-নরম্যালাইজড কপিগুলো গ্রহণ করে।

ইনপুট এবং অগ্রবর্তী-সংরক্ষিত মানসমূহ:

পরিবর্তনশীল আকৃতি অর্থ
পাবলিক query, key / রেসিডুয়াল q, k [H, B, T, K] কোয়েরি / কী
জনসাধারণের value / অবশিষ্ট v [H, B, T, V] মূল্য
beta [H, B, T] প্রতি টোকেনে ডেল্টা-রুল লেখার শক্তি
জনসাধারণের gate [H, B, T, K] চাঙ্ক-লোকাল কামসামের আগে প্রতি-টোকেন গেট ইনপুট
g_cumsum [H, B, T, K] / None লগ২ স্পেসে অভ্যন্তরীণ পোস্ট-কামসাম গেট; বাদ দেওয়া হলে পুনরায় গণনা করা হয়।
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 পাবলিক কেডিএ চুক্তির অধীনে প্রাথমিক অবস্থা
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 গ্রেডিয়েন্ট

স্থির সীমাবদ্ধতা:

  • প্রদত্ত মোজাইক কনফিগারেশনে BT=chunk_size=64 রয়েছে;
  • প্রস্তুতকৃত T অবশ্যই BT দ্বারা বিভাজ্য হতে হবে; অবশিষ্টাংশ নির্মাণের পূর্বে পরিবর্তনশীল-দৈর্ঘ্যের ইনপুটগুলি সারিবদ্ধ করা হয়;
  • মূল পাথে log2 গেট ব্যবহার করা হয়, অর্থাৎ use_exp2=True ;
  • নন-সিপি এক্সিকিউশন ব্যাকএন্ডের সরবরাহ করা K<=256 কন্ট্রাক্টকে সমর্থন করে, যার মধ্যে নন-128-অ্যালাইনড K/V ও অন্তর্ভুক্ত; সিপি-এর জন্য বর্তমানে K এবং V উভয়কেই 128-এর গুণিতক হতে হয়।

ব্যাকওয়ার্ড অর্কেস্ট্রেটর একটি N=1 অক্ষ সন্নিবেশ করে একটি অভ্যন্তরীণ চতুর্মাত্রিক চূড়ান্ত-অবস্থার কোট্যানজেন্টকে স্বাভাবিক করতে পারে, কিন্তু এটি একটি বাস্তবায়ন সামঞ্জস্যের পথ। পাবলিক টোকাম্যাক্স KDA চুক্তিটি পঞ্চমাত্রিক [B, N, H, K, V] আকারে পুনরাবৃত্ত অবস্থা গ্রহণ করে এবং ফেরত দেয়। নির্দিষ্ট-দৈর্ঘ্যের প্যালাস এক্সিকিউশনের জন্য N=1 প্রয়োজন।


২. কোড-ভিত্তিক কল চেইন এবং কার্নেল কাঠামো

বর্তমান টোকাম্যাক্স কল চেইনটি হলো:

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. গেট গ্রেডিয়েন্টের উপর বিপরীত কামসাম প্রয়োগ করুন।

সংরক্ষিত- 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-নরম্যালাইজেশন ব্যাকওয়ার্ড, এবং ভ্যারলেন আনঅ্যালাইনমেন্ট হলো কোর কার্নেলগুলোর চারপাশের ঐচ্ছিক অপারেশন।

২.১ ফিউশন কার্নেল পুনরায় গণনা করুন

প্রবেশ বিন্দু: 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 সংরক্ষিত থাকে না, তখন স্বল্প-মেমরি পথটি ব্যবহৃত হয়। এটি প্রথমে _recompute_w_u_fwd(...) এর মাধ্যমে w, u, qg, kg পুনরুদ্ধার করে, তারপর স্টেট রিকারেন্স পুনরায় চালানোর জন্য এবং h, v_new পাওয়ার জন্য chunk_gated_delta_rule_fwd_h(...) কল করে। এই পথটি আন্তঃ-চ্যাঙ্ক ক্রমিক নির্ভরতা পুনরায় প্রবর্তন করে এবং কম ফরোয়ার্ড রেসিড্যুয়াল মেমরির জন্য কম্পিউটেশন ব্যয় করে।

২.২ ডিএভি কার্নেল

প্রবেশ বিন্দু: 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}

কোডে, dAqk তৈরি করার সময় scale শুধুমাত্র একবার গুণ করা হয়। পরবর্তী ফিউশন কার্নেল সরাসরি ইতিমধ্যে স্কেল করা dAqk ব্যবহার করে, ফলে বারবার স্কেল করার প্রয়োজন হয় না।

সংরক্ষিত Aqk ইতিমধ্যেই scale অন্তর্ভুক্ত থাকে, তাই value-gradient ব্রাঞ্চটি অন্য কোনো ফ্যাক্টর ছাড়াই Aqk^T @ do ব্যবহার করে। dAqk নামের টেনসরটি fused intra backward-এর আগে স্কেল করা হয়, কারণ ঐ কার্নেলটি অন্তর্নিহিত অস্কেলকৃত কোয়েরি-কী রিলেশনকে সরাসরি ডিফারেনশিয়েট করে। যদিও chunk_kda_bwd_dAv_kernel তার পাবলিক সিগনেচারে q এবং k ধরে রাখে, বর্তমান লঞ্চার এবং কার্নেল শুধুমাত্র v , Aqk , do , scale , এবং tiling প্যারামিটারগুলো ব্যবহার করে।

2.3 dhu/WY/intra/cumsum ফিউশন কার্নেল

প্রবেশ বিন্দু: _fused_dhu_wy_intra_cumsum_pallas_jit(...) .

এই কার্নেলটি বিপরীত খণ্ড ক্রমে কার্যকর হয় এবং চার ধরনের কাজকে একত্রিত করে:

  1. dhu পশ্চাৎমুখী পুনরাবৃত্তি আউটপুট পাথ, ডেল্টা-আপডেট পাথ এবং চাঙ্ক-ক্ষয় পাথ থেকে প্রাপ্ত অবদানসমূহকে একই VMEM স্ক্র্যাচে সঞ্চয় করে ক্রস-চাঙ্ক স্টেট গ্রেডিয়েন্ট dh বজায় রাখে।

  2. WY পশ্চাৎমুখীভাবে v_new = u - w @ h , u = Akk @ (v * beta) , w = Akk @ (...) থেকে q, k, v, beta, g পর্যন্ত ব্যাকপ্রোপাগেট করে এবং Akk এর গ্রেডিয়েন্ট অবদান নির্ণয় করে।

  3. ইন্ট্রা ব্যাকওয়ার্ড dAv কার্নেল দ্বারা উৎপাদিত dAqk গ্রহণ করে, ইন্ট্রা-চ্যাঙ্ক অ্যাটেনশন ম্যাট্রিক্সের গ্রেডিয়েন্টকে q, k, beta, g পর্যন্ত ব্যাকপ্রোপাগেট করা চালিয়ে যায় এবং WY ফলাফলের সাথে সঞ্চিত করে।

  4. রিভার্স কামসাম: ফরওয়ার্ড গেটটি একটি চাঙ্ক-লোকাল কামসাম, তাই ব্যাকওয়ার্ড পাসের জন্য একটি রিভার্স কামসাম প্রয়োজন। বর্তমান ইমপ্লিমেন্টেশনটি এটিকে একই কার্নেলের মধ্যে রাখে এবং আর কোনো পৃথক কামসাম কার্নেল চালু করে না।

চাঙ্ক অক্ষটি পশ্চাৎমুখী পুনরাবৃত্তি অবস্থা বহন করার জন্য arbitrary গ্রিড অর্থ ব্যবহার করে; হেড এবং ব্যাচ অক্ষ দুটি সমান্তরাল মাত্রা।

২.৪ ঐচ্ছিক প্রসঙ্গ সমান্তরাল

যখন কনটেক্সট প্যারালাল সক্রিয় করা হয়, তখন ব্যাকওয়ার্ড পাস dAv এবং ফিউশন কার্নেলগুলির মধ্যে একটি ক্রস-র‍্যাঙ্ক স্টেট-গ্রেডিয়েন্ট মার্জ সন্নিবেশ করে।

প্রবাহটি হলো:

  1. chunk_gated_delta_rule_bwd_dhu_pre_process(...) শুধুমাত্র প্রথম আসল লোকাল সেগমেন্টটি স্ক্যান করে, কারণ সেটিই আপস্ট্রিম র‍্যাঙ্ক থেকে স্টেট গ্রহণ করতে পারে, এবং নিম্নলিখিত ফলাফল দেয়:
    • dS_ext : ইনপুট স্টেট গ্রেডিয়েন্টে এই র‍্যাঙ্কের বাহ্যিক অবদান;
    • dM : এই র‍্যাঙ্কের পশ্চাৎমুখী ট্রানজিশন-ম্যাট্রিক্স চেইন প্রোডাক্ট।
  2. শেষ ডাইমেনশন বরাবর dS_ext এবং dM প্যাক করুন এবং একটিমাত্র all_gather_into_tensor(...) ব্যবহার করে সেগুলোকে বিনিময় করুন।
  3. _merge_dht(...) forward দ্বারা সংরক্ষিত post_num_ranks এবং is_last_rank মেটাডেটা ব্যবহার করে ডাউনস্ট্রিম র‍্যাঙ্কগুলোর অবদানকে দূরতম থেকে নিকটতম ক্রমে একত্রিত করে।
  4. একটি [B,N,H,K,V] dht তৈরি করুন এবং ফিউজড ব্যাকওয়ার্ড কার্নেলে প্রবেশের আগে প্রতিটি ব্যাচ এলিমেন্টের মার্জড স্টেট গ্রেডিয়েন্টকে এর শেষ বাস্তব লোকাল সেগমেন্ট স্লটে রাখুন।

CP-এর পশ্চাৎমুখী দিকটি ডাউনস্ট্রিম র‍্যাঙ্ক থেকে আপস্ট্রিম র‍্যাঙ্কের দিকে যায়, যা অ্যালাইনড segment_ids এবং ContextParallelMetadata তে পুনরুদ্ধার করা র‍্যাঙ্ক মেটাডেটা দ্বারা নির্ধারিত হয়। B>1 প্রতিটি ব্যাচ এলিমেন্টের জন্য স্বাধীনভাবে পরিচালনা করা হয়। CP কন্ট্রাক্ট কোনো বাহ্যিক initial_state গ্রহণ করে না বা কোনো initial-state গ্রেডিয়েন্ট ফেরত দেয় না।

২.৫ ঐচ্ছিক গেট পশ্চাৎমুখী

যখন use_gate_in_kernel=True , তখন chunk_kda_bwd_custom(...) কোর ব্যাকওয়ার্ডের পরে kda_gate_bwd(...) কল করে।

এই পর্যায়ে ফিউশন কার্নেল ইতিমধ্যেই চাঙ্ক-লোকাল রিভার্স কামসাম প্রয়োগ করেছে, তাই এর dg হলো কামসামের আগে সক্রিয় প্রতি-টোকেন গেটের সাপেক্ষে গ্রেডিয়েন্ট। এরপর kda_gate_bwd(...) অ্যাক্টিভেশনকে ডিফারেনশিয়েট করে এবং এই গ্রেডিয়েন্টকে র পাবলিক gate , a_log , এবং ঐচ্ছিক delta_time_bias ম্যাপ করে, এবং যথাক্রমে dg , dA , ও dbias রিটার্ন করে।


৩. টিপিইউ অভিযোজন

এই বাস্তবায়নগত সিদ্ধান্তগুলোর লক্ষ্য হলো ব্যাকওয়ার্ড পাথকে টিপিইউ-এর এমএক্সইউ/ভিপিইউ/এইচবিএম এক্সিকিউশন মডেলের সাথে যথাসম্ভব ঘনিষ্ঠভাবে সামঞ্জস্যপূর্ণ করা।

3.1 log2 Gate এবং exp2

ইন্ট্রা এবং স্টেট-রিকরেন্স কার্নেল দ্বারা ব্যবহৃত কিউমুলেটিভ গেটকে লগ২ স্পেসে প্রকাশ করা হয়, তাই ক্ষয়কে অভিন্নভাবে 2^g আকারে লেখা হয়। পাবলিক গেট ইনপুটকে ফরোয়ার্ড গেট/কিউমসাম স্টেজের মাধ্যমে এই উপস্থাপনায় রূপান্তরিত করা হয়; এটিকে অভ্যন্তরীণ কিউমুলেটিভ মানের সাথে গুলিয়ে ফেলা উচিত নয়।

এর দুটি সুবিধা রয়েছে:

  • এক্সপোনেন্টে যোগ ও বিয়োগের মাধ্যমে ইন্ট্রা-চ্যাঙ্ক এবং ইন্টার-চ্যাঙ্ক ক্ষয়কে একত্রিত করা যেতে পারে;
  • কার্নেল হার্ডওয়্যার exp2 ব্যবহার করে, যার ফলে exp এর ভিত্তির বারবার পরিবর্তন এড়ানো যায়।

৩.২ fp32 সঞ্চয়ন এবং bf16 সংরক্ষণ

ইনপুটগুলো সাধারণত bf16 হয়, কিন্তু matmuls-এ fp32 accumulate ব্যবহৃত হয়।

মূল ফিউশন কার্নেলের ভিতরে, fp32 ইনপুটগুলো jax.lax.Precision.HIGHEST নির্বাচন করে, যেখানে bf16 ইনপুটগুলো fp32 অ্যাকুমুলেশন সহ ডিফল্ট ডট প্রিসিশন ব্যবহার করে। রিভার্স-কামসাম ডট উভয় ইনপুট ডেটাটাইপের জন্যই স্পষ্টভাবে jax.lax.Precision.HIGHEST ব্যবহার করে। এই পছন্দগুলো dg এর জন্য গুরুত্বপূর্ণ, যেখানে ক্রমবর্ধমান গেট কন্ট্রিবিউশনগুলো প্রায় বাতিল হয়ে যেতে পারে।

৩.৩ ভিএমইএম স্ক্র্যাচ ক্রস-স্টেজ ইন্টারমিডিয়েটগুলি ধরে রাখে

ফিউশন কার্নেল ব্যাকওয়ার্ড স্টেট dh একটি [MB, K, V] VMEM স্ক্র্যাচে রাখে এবং এটিকে বিপরীত চাঙ্ক ক্রমে আপডেট করে।

একই সময়ে, dv_new , WY ইন্টারমিডিয়েট গ্রেডিয়েন্ট, intra ইন্টারমিডিয়েট গ্রেডিয়েন্ট এবং রিভার্স কামসামের জন্য প্রয়োজনীয় স্থানীয় ফলাফলগুলো যথাসম্ভব VMEM-এর মধ্যেই সঞ্চালিত হতে থাকে। এর ফলে dhu, WY, intra এবং cumsum লজিক্যাল পর্যায়গুলোর মধ্যে HBM-এ বারবার লেখা এবং পড়া এড়ানো যায়।

৩.৪ মিনি-ব্যাচ ডিএমএ গ্র্যানুলারিটি নিয়ন্ত্রণ করে

একটি প্যালাস প্রোগ্রাম একবারে MB হেড/চাঙ্ক টাইলস প্রসেস করে।

একটি আনুমানিক VMEM ফুটপ্রিন্ট থেকে MB নির্বাচন করা হয়, কিন্তু এর সঠিক বাস্তবায়ন কার্নেল ভেদে ভিন্ন হয়। chunk_kda_bwd_dAv_kernel শেয়ার্ড estimate_mini_batch(...) হেল্পারটি ব্যবহার করে। সেভড- h রিকম্পিউট কার্নেল, প্রধান ফিউশন কার্নেল এবং CP প্রি-প্রসেস বর্তমানে কার্নেল-নির্দিষ্ট ক্যাপ এবং বিভাজ্যতা/অ্যালাইনমেন্ট সমন্বয় সহ স্থানীয় VMEM-বাজেট হিউরিস্টিকস ব্যবহার করে।

  • অত্যধিক ছোট হওয়ার ফলে ডিএমএ গ্র্যানুলারিটি অপর্যাপ্ত হয় এবং এইচবিএম ব্যান্ডউইথের ব্যবহার কম হয়;
  • অতিরিক্ত বড় হলে VMEM ও রেজিস্টারের উপর চাপ এবং স্ট্যাটিক আনরোলিং ওভারহেড বৃদ্ধি পায়।

অতএব প্রতিটি লঞ্চার তার টাইল ফুটপ্রিন্ট অনুমান করে এবং তার গ্রিড, বিভাজ্যতা, এবং টিপিইউ মাইনর-ডাইমেনশন অ্যালাইনমেন্টের প্রয়োজনীয়তা সাপেক্ষে একটি MB নির্বাচন করে। যখন কোনো উপযুক্ত বিভাজক উপলব্ধ থাকে না, তখন সেভড- h রিকম্পিউট লঞ্চারটি তার ফ্ল্যাটেনড চাঙ্ক কাউন্ট প্যাড করতে পারে; প্রধান ফিউশন এবং সিপি লঞ্চারগুলো MB কমাতে থাকে যতক্ষণ না এটি হেড কাউন্টকে ভাগ করে।

৩.৫ ব্লকস্পেক ডেটা মুভমেন্ট এবং এক্সপ্লিসিট সিপি ডিএমএ

তিনটি কোর সেভড- h কার্নেল এইচবিএম টাইলস বর্ণনা করার জন্য প্যালাস গ্রিড এবং BlockSpec অবজেক্ট ব্যবহার করে। এগুলোর লঞ্চারে কোনো সুস্পষ্ট হাতে লেখা অ্যাসিঙ্ক-কপি লুপ থাকে না।

সিপি প্রি-প্রসেসটি ভিন্ন: _chunk_gated_delta_rule_bwd_dhu_pre_process_kernel হলো সক্রিয় সিপি ইমপ্লিমেন্টেশন এবং এটি সুস্পষ্টভাবে pltpu.make_async_copy , ডিএমএ সেমাফোর এবং ডাবল-বাফারড-ভিএমইএম ইনপুট ব্যবহার করে। এটি হেড গ্রুপ এবং চাঙ্কগুলোর উপর একটি একক-প্রোগ্রাম রিভার্স স্ক্যান, যা পরবর্তী ইনপুট ট্রান্সফারকে বর্তমান ম্যাট্রিক্স ওয়ার্কের সাথে ওভারল্যাপ করে এবং প্রতিটি হেড গ্রুপের dS_ext/dM সামারি অ্যাসিঙ্ক্রোনাসভাবে লেখে।

৩.৬ টাইলস বসানো এবং প্যাডিং করার মূলনীতি

TPU টাইলিং নির্দিষ্ট কিছু পশ্চাৎবর্তী মাত্রাকে হার্ডওয়্যার টাইলের সীমানার সাথে সারিবদ্ধ করে। যদিও প্যাড করা উপাদানগুলো অর্থগতভাবে শূন্য, তবুও সেগুলো HBM এবং DMA পর্যায়ে প্রকৃত ব্যান্ডউইথ ব্যবহার করে।

বর্তমান বাস্তবায়নটি একটি বাস্তবসম্মত নীতি অনুসরণ করে: আরও জটিল প্যাকড লেআউট কেবল তখনই চালু করা হয়, যখন লেআউট পরিবর্তনের মাধ্যমে বাদ দেওয়া যায় এমন HBM ট্র্যাফিক অতিরিক্ত রিশেপ/গ্যাদার/কপি খরচকে ছাড়িয়ে যায়।

অতএব:

  • নন-সিপি ব্যাকওয়ার্ড প্রক্রিয়াকরণ সিপি লেন সীমাবদ্ধতা আরোপ না করেই সরবরাহকৃত আনঅ্যালাইনড K/V কেসগুলোকে সমর্থন করে, অপরদিকে সিপি প্রি-প্রসেসিংয়ের জন্য 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 এবং chunk mapping ব্যবহার করে:

  • সেগমেন্ট আইডি ০ যুক্ত খণ্ডগুলো হলো প্যাডিং;
  • প্রতিটি বাস্তব খণ্ডের শেষ অংশটি dht দিয়ে পশ্চাৎমুখী পুনরাবৃত্তির সূচনা করে;
  • প্রতিটি বাস্তব সেগমেন্টের প্রথম অংশ dh0 আউটপুট করে।

শুধুমাত্র উদাহরণের জন্য, B=1, T=12, BT=4 বিবেচনা করুন, যেখানে যথাক্রমে 8 এবং 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] লেআউটের সাথে অসারিবদ্ধ থাকে।


৪. ফিউশন পারফরম্যান্স বিশ্লেষণ

এই অংশের পরিমাপগুলো ঐতিহাসিক অপ্টিমাইজেশনের প্রমাণ হিসেবে সংরক্ষিত আছে। এগুলো বর্তমান টোকাম্যাক্স অ্যাডাপ্টার, কেন্দ্রীভূত অ্যালাইনমেন্ট/রেসিডুয়াল কন্ট্রাক্ট, বর্তমান সিপি মেটাডেটা হ্যান্ডঅফ এবং লঞ্চার-স্তরের কিছু নামকরণের পরিবর্তনের পূর্ববর্তী। গ্রহণযোগ্যতার সীমা বা রিগ্রেশন বেসলাইন হিসেবে ব্যবহার করার আগে এগুলোকে অবশ্যই বর্তমান কমিটে পুনরায় চালাতে হবে।

ঐতিহাসিক পরিমাপ কনফিগারেশন: 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; চাংকিং করার পর লিডিং ডাইমেনশন = 2128 = 16×133 (NT=133 = 128 বেস চাংক + 5 সেগমেন্ট প্যাডিং)। HBM ব্যান্ডউইথ 1.6 TB/s ধরা হয়েছে।

কার্নেল রানটাইম (µs) এইচবিএম (এমবি) BW নিম্ন সীমা (µs) BW ব্যবহার
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 সেগমেন্ট), H=16, B=1, T=8192, K=V=128, BT=64; চাংকিং করার পর লিডিং ডাইমেনশন = 2128 = 16×133 (NT=133 = 128 বেস চাংক + 5 সেগমেন্ট প্যাডিং)। HBM ব্যান্ডউইথ 3.69 TB/s ধরা হয়েছে।

কার্নেল রানটাইম (µs) এইচবিএম (এমবি) BW নিম্ন সীমা (µs) BW ব্যবহার
fused_recompute_w_u_vnew_from_h_pallas ১৮১.৬৭৫ ৪১৮.৯ ১১৩.৫২ ৬২.৫%
chunk_kda_bwd_dAv_kernel ১০০.২৩৭ ২০৯.২ ৫৬.৬৯ ৫৬.৬%
_fused_dhu_wy_intra_cumsum_pallas_jit ৭৮৪.২৯৯ ৯৫১.৮ ২৫৭.৯৪ ৩২.৯%

৪.৩ প্রসঙ্গ সমান্তরালের কর্মক্ষমতার বৈশিষ্ট্য

পরিমাপ: TPU v6e-4 (cp_size=4, 4-কোর SPMD প্যারালেলিজম), bf16, একক সেগমেন্ট per_rank_T=8192 → গ্লোবাল T=32768, H=16, K=V=128, BT=64; চাংকিংয়ের পর প্রতি-র‍্যাঙ্ক লিডিং ডাইমেনশন = 2128 = 16×133।

একক কার্নেল (একই প্যাডিং ভিত্তিতে):

কার্নেল রানটাইম (µs) এইচবিএম (এমবি) BW নিম্ন সীমা (µs) BW ব্যবহার
fused_recompute_w_u_vnew_from_h_pallas ২৭৩.৪ ৪০৬.৩ ২৫৩.৯৫ ৯২.৯%
chunk_kda_bwd_dAv_kernel ১৪৬.৮ ২০২.৯ ১২৬.৮১ ৮৬.৪%
_fused_dhu_wy_intra_cumsum_pallas_jit ৮৭২.৭ ৯১৫.১ ৫৭১.৯৭ ৬৫.৫%
CP chunk_gated_delta_rule_bwd_dhu_pre_process ৫৯১.২ ২৩৮.৮ ১৪৯.২৬ ২৫.২%

পরিমাপ: TPU v7x-4 (cp_size=4, 4-কোর SPMD প্যারালেলিজম), bf16, একক সেগমেন্ট per_rank_T=8192 → গ্লোবাল T=32768, H=16, K=V=128, BT=64; চাংকিংয়ের পর প্রতি-র‍্যাঙ্ক লিডিং ডাইমেনশন = 2128 = 16×133।

কার্নেল রানটাইম (µs) এইচবিএম (এমবি) BW নিম্ন সীমা (µs) BW ব্যবহার
fused_recompute_w_u_vnew_from_h_pallas ১৮১ ৪০৬.৩ ১১০.১১ ৬০.৮%
chunk_kda_bwd_dAv_kernel ১০০ ২০২.৯ ৫৪.৯৯ ৫৫.০%
_fused_dhu_wy_intra_cumsum_pallas_jit ৭৮৫ ৯১৫.১ ২৪৭.৯৯ ৩১.৬%
CP chunk_gated_delta_rule_bwd_dhu_pre_process ৫১১ ২৩৮.৮ ৬৪.৭২ ১২.৭%

পরিমাপকৃত সিপি পাথটি নন-সিপি পাথের মতোই একই তিনটি কোর কার্নেল ব্যবহার করেছিল; এর অতিরিক্ত কাজটি এসেছিল ইন্টার-র‍্যাঙ্ক স্টেট-গ্রেডিয়েন্ট মার্জ থেকে।

ঐতিহাসিক প্রেক্ষাপটে, যোগাযোগের মূল ব্যয়টি ছিল না স্বয়ং সমষ্টিগত বিষয়টি, বরং যোগাযোগের পূর্ববর্তী স্থানীয় প্রাক-প্রক্রিয়াটি:

  • প্রাক-প্রক্রিয়াটির জন্য এই র‍্যাঙ্কের চাঙ্ক সিকোয়েন্স বরাবর ইনপুট স্টেট গ্রেডিয়েন্ট এবং ব্যাকওয়ার্ড ট্রানজিশন ম্যাট্রিক্সকে সঞ্চিত করতে হবে;
  • এই ধাপটির প্রকৃতি স্ক্যান/হ্রাসের মতো এবং এটিকে একটি সাধারণ প্রতি-খণ্ড ম্যাটমালের মতো সম্পূর্ণরূপে সমান্তরাল করা কঠিন;
  • all_gather নিজেই তুলনামূলকভাবে সাশ্রয়ী, কারণ এর মাধ্যমে পূর্ণ-সিকোয়েন্স টোকেন টেনসর নয়, বরং সংকুচিত স্টেট গ্রেডিয়েন্ট এবং ট্রানজিশন ম্যাট্রিক্স বিনিময় করা হয়।

অতএব, সিপি পাথের অপ্টিমাইজেশনের মূল লক্ষ্য হওয়া উচিত প্রি-প্রসেস:

  • সম্পূর্ণ ইনপুট খণ্ডের পুনরাবৃত্তিমূলক পাঠ হ্রাস করুন;
  • চেইন প্রোডাক্ট / স্ক্যানের সমান্তরালকরণ উন্নত করুন;
  • অথবা স্টেট-মার্জ লজিকের একটি অংশকে সামনে এগিয়ে নিয়ে বিদ্যমান কোনো ব্যাকওয়ার্ড স্টেজে একীভূত করুন।

৪.৪ সংখ্যাগত সামঞ্জস্য

বর্তমান compute_intra_backward(...) ইমপ্লিমেন্টেশনটি compute_reverse_cumsum_dg(...) এর আগে dg_acc + dg_intra স্পষ্টভাবে ইনপুট রেফারেন্স dtype-এ এবং তারপর আবার fp32-তে কাস্ট করে। এই dtype সীমানাটি প্রদত্ত নিউমেরিক্যাল আচরণের একটি অংশ এবং ফিউশনটি রিফ্যাক্টর করার সময় এটি সংরক্ষণ করা বা ইচ্ছাকৃতভাবে পুনরায় যাচাই করা উচিত। একটি আনফিউজড বা XLA রেফারেন্সের সাথে তুলনা করার সময় এই ট্রাঙ্কেশন পয়েন্টের জন্য উপযুক্ত টলারেন্স ব্যবহার করা উচিত।

৪.৫ বর্তমান বাস্তবায়নের জন্য পুনঃবৈধকরণ ম্যাট্রিক্স

যেহেতু উপরের পরিমাপগুলো ঐতিহাসিক, তাই বর্তমান-হেড যাচাইকরণে সঠিকতা এবং কর্মক্ষমতা আলাদাভাবে অন্তর্ভুক্ত করা উচিত। ন্যূনতমভাবে, সঠিকতা পুনঃযাচাইকরণে নিম্নলিখিত বিষয়গুলো অনুশীলন করা উচিত:

  • সংরক্ষিত- h পথ ( rematerialize_for_backward=False ) এবং পূর্ণ-পুনঃবস্তুকরণ পথ ( 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 শেপ;
  • সমর্থিত ১২৮-অ্যালাইনড K/V শেপ এবং একাধিক CP সাইজে CP এক্সিকিউশন;
  • বিএফ১৬ ইনপুট এবং সমর্থিত এফপি৩২ পাথ।

পারফরম্যান্স পুনঃযাচাইকরণে সুনির্দিষ্ট কমিট, টিপিইউ জেনারেশন ও টপোলজি, সফটওয়্যার ভার্সন, শেপ, ভারলেন ডিস্ট্রিবিউশন, সিপি সাইজ, ওয়ার্মআপ পলিসি, ইটারেশন কাউন্ট এবং প্রতিটি সংখ্যা ডিভাইস-ভিত্তিক নাকি এন্ড-টু-এন্ড ল্যাটেন্সি, তা রিপোর্ট করা উচিত। কার্নেল-স্তরের ব্যান্ডউইথ অনুমানের সাথে একটি এন্ড-টু-এন্ড ব্যাকওয়ার্ড পরিমাপ থাকা উচিত, যাতে লঞ্চ, কালেক্টিভ এবং অ্যাডাপ্টার ওভারহেডগুলো দৃশ্যমান হয়।


৫. ভবিষ্যৎ অনুকূলকরণের দিকনির্দেশনা

বর্তমান বাস্তবায়নটি ইতিমধ্যেই কার্নেল ফিউশনের মাধ্যমে এইচবিএম রাউন্ড ট্রিপ কমানোর উপর কেন্দ্র করে তৈরি, কিন্তু প্রোফাইলিংয়ের ফলাফল দুটি স্পষ্ট পরবর্তী দিকনির্দেশনা দেয়।

৫.১ v7 এর কর্মক্ষমতা উন্নত করা

ঐতিহাসিক পরিমাপ অনুযায়ী, কোর কার্নেলগুলো v7x-এ দ্রুততর চললেও v6e-এর তুলনায় এর আনুমানিক ব্যান্ডউইথ ব্যবহার কম ছিল। একটি বর্তমান প্রোফাইলের মাধ্যমে প্রথমে এই ফলাফলটি নিশ্চিত করা উচিত; যদি তা অপরিবর্তিত থাকে, তবে পরবর্তী কাজের মূল লক্ষ্য হওয়া উচিত ফিউজড কার্নেলগুলোকে v7-এর এক্সিকিউশন মডেলের সাথে আরও ভালোভাবে মেলানো।

  • v7-এর জন্য টাইল সাইজ, মিনি-ব্যাচ সিলেকশন এবং VMEM প্রেসার পুনরায় পর্যালোচনা করুন;
  • প্রোফাইলে যেখানে স্পষ্টভাবে দৃশ্যমান, সেখানে স্থানীয় লেআউট ওভারহেড এবং প্যাডিং ট্র্যাফিক হ্রাস করুন;
  • প্রধান ফিউশন কার্নেলের কোন অংশগুলো ব্যান্ডউইথ-বাউন্ড না হয়ে কম্পিউট-বাউন্ড বা পাইপলাইন-বাউন্ড, তা শনাক্ত করুন।

লক্ষ্যটি গাণিতিক বিভাজন পরিবর্তন করা নয়, বরং বিদ্যমান ফিউজড পাথটিকে এমনভাবে পুনঃসমন্বয় করা যাতে এটি v7-এর উচ্চতর ব্যান্ডউইথ এবং গণনা ক্ষমতার সাথে আরও ভালোভাবে খাপ খাইয়ে নিতে পারে।

৫.২ প্রসঙ্গ সমান্তরাল অপ্টিমাইজ করুন

সিপি পাথটি ক্রস-র‍্যাঙ্ক এক্সচেঞ্জের আগে একটি স্থানীয় প্রি-প্রসেস যোগ করে। ঐতিহাসিক প্রোফাইল থেকে বোঝা যায় যে, এই স্থানীয় স্ক্যান/হ্রাসের কাজটি স্বয়ং সম্মিলিত যোগাযোগের চেয়েও বড় প্রতিবন্ধকতা ছিল; অপ্টিমাইজেশনের কাজ শুরু করার আগে একটি বর্তমান প্রোফাইলের মাধ্যমে এই ভারসাম্য নিশ্চিত করা উচিত।

অতএব, ভবিষ্যতের সিপি অপ্টিমাইজেশনে নিম্নলিখিত বিষয়গুলোকে অগ্রাধিকার দেওয়া উচিত:

  • প্রি-প্রসেসে পুনরাবৃত্ত পাঠ কমানো;
  • ট্রানজিশন-ম্যাট্রিক্স চেইন প্রোডাক্টের সমান্তরালকরণের উন্নতি সাধন;
  • সিপি স্টেট-গ্রেডিয়েন্ট প্রস্তুতির অংশবিশেষ বিদ্যমান ব্যাকওয়ার্ড কার্নেলগুলোর সাথে একীভূত করা যায় কিনা, তা খতিয়ে দেখা হচ্ছে।

যোগাযোগের ধরণটি সংক্ষিপ্ত থাকা উচিত: পূর্ণ টোকেন-স্তরের টেনসরের পরিবর্তে অবস্থা-স্তরের সারাংশ বিনিময় করা হোক।