transformer_lens.model_bridge.supported_architectures.rwkv7 module

RWKV-7 (“Goose”) architecture adapter (RWKV7ForCausalLM).

Family fla-hub/rwkv7-*: attention-free recurrent LM from the flash-linear-attention library, loaded via remote code. Blocks pair a time-mixing (generalized delta rule) and a token-shifted squared-ReLU channel-mixing sublayer under biased pre-LN; no positional embeddings.

Adapter decisions: - Full delegation: recurrence, token shift, LoRAs, and cross-block v_first

threading all run inside the fla forward (v_first is managed by the HF model-level forward, so delegated blocks get it for free). weight_processing_conversions = {}.

  • OpaqueBlockBridge: BlockBridge’s hook aliases hardcode the standard pre-norm attention flow; only hook_in/hook_out are sound here.

  • ffn_norm is a fused add-and-norm (normed, residual) = ffn_norm(x, res, True) under config.fuse_norm — wrapped as a delegating GeneralizedComponent because NormalizationBridge can’t express it.

  • head_dim is read but never assigned (aliases the read-only d_head).

class transformer_lens.model_bridge.supported_architectures.rwkv7.RWKV7ArchitectureAdapter(cfg: Any)

Bases: ArchitectureAdapter

Architecture adapter for RWKV7ForCausalLM (RWKV-7 “Goose”).

Attention-free recurrent decoder: a flat stack of pre-norm blocks, each a generalized-delta-rule time-mixing sublayer plus a token-shifted squared-ReLU channel-mixing sublayer, wrapped by standard biased LayerNorm. The recurrence and the cross-block v_first threading live inside the fla remote-code forward, which the bridge delegates to; see the module docstring for the full set of adapter decisions.

__init__(cfg: Any) → None

Initialize the RWKV-7 architecture adapter.

applicable_phases: list[int] = []
component_mapping: ComponentMapping | None
prepare_loading(model_name: str, model_kwargs: dict) → None

Patch fla’s RWKV-7 remote code for transformers v5 compatibility.

Two defensive patches, mirroring raven:

  1. Tied-weights format. RWKV7ForCausalLM._tied_weights_keys is a list (["lm_head.weight"], the 4.x form). v5’s tie_weights -> get_expanded_tied_weights_keys calls .keys() on the mapping, which raises AttributeError on a list. Rewrite it to the v5 dict form {"lm_head.weight": "model.embeddings.weight"}. RWKV-7 defaults tie_word_embeddings=False (so the list path is usually short- circuited before .keys()), but the rewrite is harmless when untied and prevents the crash on any tied checkpoint.

  2. Weight re-init. Under v5’s meta-device load-then-materialise flow, PreTrainedModel._init_weights is invoked on modules that already hold checkpoint weights, re-randomising them. Guard it to skip modules whose parameters are already on a real (non-meta) device — the same defensive patch openelm.py / raven.py apply.

Parameters:
  • model_name – The HuggingFace model name/path.

  • model_kwargs – The kwargs dict for from_pretrained().

supports_batched_generation: bool = False
supports_kv_cache: bool = False
uses_split_attention: bool
weight_processing_conversions: Dict[str, ParamProcessingConversion | str] | None