RWKV hxa07D アーキテクチャ解説(hxa079 からの進化ポイント)
こんちは、OpenMOSEです。
今回は hxa07D アーキテクチャを、解説します。
以前の hxa079 解説記事(背景・思想の前提として読むとスムーズです):
https://zenn.dev/openmose/articles/bc4804b47bdc94
TL;DR(hxa07Dのキーポイント)
hxa07D は RWKV v7(x070 “Goose”)系の「RNN-Transformer」アーキテクチャをベースに、次の点を強化しました。
-
Decay(忘却選択)を大型化
RWKVブロックが多い構成で、RULERの長文性能とGSM8Kの推論性能を両立させるために、忘却の表現力を強くしました(実験ベースで必要性が判明)。 -
k_first / v_first を再導入(復活)
hxa07A/07B/07C で一度外していましたが、「トレーニング時のctx」と「有効ctx」の伸びが頭打ちになっていたため、再導入。
とくに key残差(k_first) が外挿に強く、効果が大きかったです。 -
MoE 変換を“正式に”サポート
MoEは attention 系の構造が敏感で、Denseほど素直に変換が通りません。
そこで 初期Weightスケーリング等を見直して、変換性能を大きく改善しました。 -
学習フローを体系化(Stage1→選択的NoPE導入)
Stage1でまず全層RWKV化 → Loss/Logから 変換困難層を特定 → 選択的にNoPE層を導入、という流れを前提に設計しています。
背景:RWKV v7系は「SSMとは違う」時間忘却を持つ
RWKV はよく SSM 系(Mamba等)と並べられますが、感覚としては別物で、hxa07D の説明はここが出発点になります。
- RWKV は 指数的な decay(忘却) を持ちます
- その decay を Projection(学習可能な写像)でコントロールして、
“残す/忘れる/混ぜる” を頭の中でやる、という設計です
hxa07D は、この「忘却の設計」を より強く、より扱いやすくする方向に寄せています。
まずコード:hxa07D forward の全体像
以下の forward をベースに解説します(コメントも含めて、そのまま貼ります)。
def forward(self, x, v_first,k_first,attention_mask,position_embeddings,position_ids,x_emb):
B, T, C = x.size()
#removed tokenshift
H = self.num_attention_heads#self.n_head
if self.RKNormMode == True:
r = self.r_norm(self.receptance(x).view(B,T,self.num_attention_heads,-1))
k = self.k_norm(self.key(x).view(B,T,self.num_key_value_heads,-1))
else:
r = self.receptance(x).view(B,T,self.num_attention_heads,-1)
k = self.key(x).view(B,T,self.num_key_value_heads,-1)
log_neglog_w = -F.softplus(-(self.w0 + F.tanh(x @ self.w1) @ self.w2)) - 0.5
w = (-log_neglog_w.exp()).exp()
v = self.value(x)
k = k.view(B, T, self.num_key_value_heads, self.head_size)
v = v.view(B, T, self.num_key_value_heads, self.head_size)
cos, sin = position_embeddings
r, k = apply_rotary_pos_emb(r, k, cos, sin, unsqueeze_dim=2)
if self.layer_id == 0:
v_first = v # store the v of the first layer
k_first = k # store the k of the first layer
else:
v = v + (v_first - v) * torch.sigmoid(self.v0 + (x @ self.v1) @ self.v2).view(B,T,self.num_key_value_heads,-1) # add value residual
k = k + (k_first - k) * torch.sigmoid(self.k0 + (x @ self.k1) @ self.k2).view(B,T,self.num_key_value_heads,-1) # add key residual
# repeat k/v heads if n_kv_heads < n_heads
#modified repeat_kv B,T,H_kv,D) -> B,T,H,D -> B,T,C
k = repeat_kv(k, self.num_key_value_groups)
v = repeat_kv(v, self.num_key_value_groups)
k = k.view(B, T, -1)
v = v.view(B, T, -1)
#so now all B,T,C tensors
g = torch.sigmoid(x @ self.g1) @ self.g2
a = torch.sigmoid(self.a0 + (x @ self.a1) @ self.a2)
kk = F.normalize(k.view(B,T,H,-1), dim=-1, p=2.0).view(B,T,-1)
k = k * (1.0 - w + a)
x = RUN_CUDA_RWKV7g(r, log_neglog_w, k, v, -kk, kk*a,self.head_size,attention_mask)
x = x.view(B,T,-1)
x = x * (self.head_size ** -0.5)
x = x + ((r.view(B,T,H,-1)*k.view(B,T,H,-1)*self.r_k).sum(dim=-1, keepdim=True) * v.view(B,T,H,-1)).view(B,T,-1)
x = self.output(x*g)
return x, v_first, k_first
ブロック分解:hxa07D は「RWKVの核」を太くしつつ、Transformer互換も維持する
1) Tokenshift を削除(引き続き)
#removed tokenshift の通り、hxa079 と同じ思想です。
- 量子化・高速化実装で扱いやすい
- Transformer→RWKV 変換で構造差分が減り、学習が素直になる
ここは hxa07D でも継承しています。
2) r/k の投影 +(オプション)RKNorm
r = receptance(x)
k = key(x)
に加えて、hxa07D では
if RKNormMode:
r = r_norm(...)
k = k_norm(...)
という R/K 正規化モードを持たせています。
意図としてはシンプルで、
-
変換学習の初期に r/k のスケール暴れが起きると、
- decay と混ざったときに勾配が不安定
- MoE のように敏感な分岐で破綻しやすい
ので、“必要なときだけ” 正規化で支えるためのスイッチです。
(常時ONが正義ではなく、モデル/データ/スケール次第で切り替えます)
3) Decay(忘却)を「log space」で安定に作る
log_neglog_w = -softplus(-(w0 + tanh(x@w1)@w2)) - 0.5
w = exp(-exp(log_neglog_w))
ここが hxa07D の心臓です。
-
wは(感覚として) 0〜1の忘却係数 -
それを log(-log(w)) 的な空間で扱うことで
- 数値的に安定
- “忘却が効きすぎる/効かなすぎる” の暴走を抑えやすい
そして hxa07D のテーマはここで、
RWKVブロックが多い構成でも、忘却の表現力が足りなくならないようにする
ために、w0,w1,w2 周りの設計(容量/rank/スケール)を強化しています。
4) RoPE を r/k に適用(Transformer互換を維持)
r, k = apply_rotary_pos_emb(r, k, cos, sin, unsqueeze_dim=2)
「RWKVなのにRoPE?」と思うかもしれませんが、ここは **“変換”**がテーマです。
- 元が Transformer(特に RoPE 系)なら、その座標系のまま RWKV に落とせる
- NoPE 層を後段で差し込む戦略とも相性が良い
結果として、RWKVブロック中心でも「Transformer由来の位置表現」を受け止められます。
5) k_first / v_first(Layer0のK/Vを残差で混ぜる)
if layer0:
v_first = v
k_first = k
else:
v = v*(1-gv) + v_first*gv
k = k*(1-gk) + k_first*gk
hxa07D の大きな変更点のひとつが、**これを“また主役に戻した”**ことです。
なぜ復活したか(実験の肌感)
hxa07A/07B/07C では削っていましたが、
- トレーニングctxに対して有効ctxが伸びない
- とくに長文外挿(“遠いトークン”の扱い)で頭打ち
という問題が出ました。
感覚的には、
-
v_firstは「情報の保管庫」 -
k_firstは「検索キーの座標系」
に近く、特に k_first の残差が「外挿耐性」に効いている印象が強いです。
(“遠い記憶を引く”とき、キーの土台が安定しているのが効く)
6) GQA:KV head を repeat して B,T,C に畳む
k = repeat_kv(k, num_key_value_groups)
v = repeat_kv(v, num_key_value_groups)
k, v -> view(B,T,-1) # B,T,C
- KV head 数 < Q head 数 のときに repeat
- その後すぐ B,T,C に畳む
この設計は、後段の CUDA kernel を “B,T,Cテンソル前提”で通すための都合も大きいです。
(実装が素直になり、最適化ポイントが集中しやすい)
7) g と a:ゲートを “軽量に” 作る(LoRAっぽい2段)
g = sigmoid(x@g1) @ g2
a = sigmoid(a0 + (x@a1)@a2)
ここは hxa07D の “設計の癖” が出ている部分で、
- ゲートを **低ランク(2段)**で作って軽量にする
- MoE のように敏感な構造でも、ゲートの学習が破綻しにくい
狙いがあります。
a はこのあと
k = k * (1 - w + a)
で効いてくるので、忘却(w)と、残す(a)を同居させる役目です。
8) kk 正規化と kernel 入力(RUN_CUDA_RWKV7g)
kk = normalize(k per head)
x = RUN_CUDA_RWKV7g(r, log_neglog_w, k, v, -kk, kk*a, head_size, attention_mask)
ここは「RWKV v7g の本体」に相当します。
-
r:Receptance(ゲート付きQuery的な役) -
log_neglog_w:忘却(指数decay)の制御 -
k, v:Key/Value(ただし B,T,C に畳まれている) -
-kkとkk*a:正規化されたキー由来の追加項(安定化&表現力のため)
attention_mask も入れているので、実運用での可変長やパディングにも対応しやすいです。
9) 仕上げ:スケーリング + “rkv” の追加項 + 出力ゲート
x *= head_size**-0.5
x += ((r*k*r_k).sum * v)
x = output(x * g)
-
head_size**-0.5は Transformer 的なスケーリング - その後の加算項は、直感的には “超ローカルな注意(1-stepっぽい強化)” に近い働きで、
RWKV kernel の出力を補強します - 最後に
gで出力をゲートしてoutputに流す
hxa07D の本題1:Decay(忘却選択)の大型化
hxa07D で一番言いたいのはここです。
RWKV を「少数層」なら忘却の表現力が多少弱くても誤魔化せますが、
RWKVブロックが多い構成(たとえば 45/48 層が RWKV みたいな設計)では、
- 長文タスク(RULER)を伸ばすために “忘れない” を増やすと、推論(GSM8K)が鈍る
- 推論(GSM8K)を尖らせるために “忘れる” を強めると、長文が死ぬ
という綱引きが露骨に出ます。
だから hxa07D では
-
log_neglog_wの表現(学習のしやすさ) -
aとの合成(忘却と保持の同居) - 必要なら RKNorm でスケールを抑える
のセットで、**“忘却そのものを強く賢くする”**方針に寄せました。
hxa07D の本題2:k_first / v_first 再導入(特に k_first が強い)
自分の観測では、k_first 残差は外挿に対してかなり強力です。
- 同じ training ctx でも “有効ctx” が伸びやすい
- 長文で「検索キーの座標系」が崩れにくい
hxa07D は、ここを実用寄りに割り切って戻しました。
「速い」「変換が通る」だけでなく、“長文で壊れない” を優先しています。
hxa07D の本題3:MoE 変換の正式サポート
MoE は構造上、
- 分岐のスケールが敏感
- attention まわりが小さく、少しの崩れが致命傷になりやすい
という性質があり、Dense と同じノリで RWKV 変換をすると失敗しがちです。
hxa07D では、
- 初期 weight スケーリングの見直し
- ゲート(g/a/k_first/v_first の混ぜ方)を低ランクで安定に
- RKNormMode という “安全装置”
あたりをセットで整えて、MoE でも変換の成功率と最終性能が上がるようにしています。
トレーニング戦略:Stage1 → ログで難所特定 → 選択的NoPE導入
hxa07D は “最初から完璧な構成を当てに行く” より、次の流れを前提にしています。
-
Stage1:まず全層RWKVとして変換学習
- まずは「変換が通る土台」を作る
-
Loss / Log を見て、変換が難しい層(難所)を特定
- 層ごとの loss / KL / 勾配の暴れなどを見る
-
その難所にだけ、選択的に NoPE 層を導入
- すべてを attention に戻すのではなく、必要最小限だけ入れる
この手順のメリットは、
- attention を最小限にできる(KV cache を最小化できる)
- “変換できない場所” だけ人間が介入できる
- MoE のように敏感なモデルでも、壊れ方を局所化できる
という点です。
まとめ
hxa07D は、hxa079 の「Tokenshiftを消して、変換しやすく速くする」路線をベースにしつつ、
- 忘却(Decay)を大型化して、長文と推論の両立を狙う
- k_first/v_first を再導入して、有効ctxと外挿耐性を押し上げる
- MoE 変換を“ちゃんと通す”ための設計を入れる
- Stage1→難所特定→選択的NoPE の運用を前提にする
という方向で、実戦投入寄りにアップデートしたアーキテクチャです。
またベンチや学習ログなど、公開できる範囲で続報も書きます。
RWKVの人口、増えるといいなぁ 🪿
Discussion