📑

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 に畳まれている)
  • -kkkk*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 は “最初から完璧な構成を当てに行く” より、次の流れを前提にしています。

  1. Stage1:まず全層RWKVとして変換学習

    • まずは「変換が通る土台」を作る
  2. Loss / Log を見て、変換が難しい層(難所)を特定

    • 層ごとの loss / KL / 勾配の暴れなどを見る
  3. その難所にだけ、選択的に NoPE 層を導入

    • すべてを attention に戻すのではなく、必要最小限だけ入れる

この手順のメリットは、

  • attention を最小限にできる(KV cache を最小化できる)
  • “変換できない場所” だけ人間が介入できる
  • MoE のように敏感なモデルでも、壊れ方を局所化できる

という点です。


まとめ

hxa07D は、hxa079 の「Tokenshiftを消して、変換しやすく速くする」路線をベースにしつつ、

  • 忘却(Decay)を大型化して、長文と推論の両立を狙う
  • k_first/v_first を再導入して、有効ctxと外挿耐性を押し上げる
  • MoE 変換を“ちゃんと通す”ための設計を入れる
  • Stage1→難所特定→選択的NoPE の運用を前提にする

という方向で、実戦投入寄りにアップデートしたアーキテクチャです。

またベンチや学習ログなど、公開できる範囲で続報も書きます。
RWKVの人口、増えるといいなぁ 🪿

Discussion