RWKV

RWKV-FLA 使用教程

RWKV-FLA 提供基于 PyTorch 与 Triton 的 RWKV 算子、模型组件和 Hugging Face 风格接口,适合跨硬件实验与已有 FLA 模型推理。

本页适用于 RWKV-FLA 算子实验和已有 fla-hub 模型推理。

FLA 的 RWKV-7 实现尚未与官方参考实现完全对齐,训练效果和性能也存在差距。需要复现或训练正式 RWKV-7 模型时,请使用 RWKV-7 train_temp 参考实现

特性与优势

  • 跨平台支持:支持多种硬件后端,包括 NVIDIA、Intel、AMD、摩尔线程、沐曦等
  • Triton 算子:提供 RWKV 的并行与循环计算实现
  • 灵活的API:提供友好的接口,易于与现有代码集成
  • Transformers 接口:可以加载兼容的 fla-hub 模型进行推理

安装指南

先根据 PyTorch 安装页面 安装适合当前硬件的 PyTorch,再选择对应的 RWKV-FLA 可选依赖:

# 创建新环境(支持 Python 3.10 ~ 3.13)
conda create -n rwkv-fla python=3.12
conda activate rwkv-fla

# 先按 PyTorch 官网命令安装 CUDA 版 PyTorch,再安装 RWKV-FLA
python -m pip install --upgrade "rwkv-fla[cuda]"
python -m pip install --upgrade "rwkv-fla[xpu]"
python -m pip install --upgrade "rwkv-fla[rocm]"

遇到算子或兼容性问题时,先运行 python -m pip install --upgrade rwkv-fla triton

如果升级后问题仍然存在,请在反馈时附上 PyTorch、Triton、显卡和驱动版本。

模型推理示例

下面的示例使用 Hugging Face Transformers 风格接口加载模型。将代码保存为 xx.py,然后运行 python xx.py

from transformers import AutoModelForCausalLM, AutoTokenizer
import torch
# 加载模型和分词器
model = AutoModelForCausalLM.from_pretrained('fla-hub/rwkv7-1.5B-g1',torch_dtype=torch.bfloat16, trust_remote_code=True)
tokenizer = AutoTokenizer.from_pretrained('fla-hub/rwkv7-1.5B-g1', trust_remote_code=True)
model = model.cuda()
# 准备对话历史和提示词
prompt = "什么是深度学习?\n\n"
messages = [
    {"role": "user", "content": "你是谁?"},
    {"role": "assistant", "content": "我是RWKV,一个基于RWKV架构的AI模型。"},
    {"role": "user", "content": prompt}
]
# 应用对话模板
text = tokenizer.apply_chat_template(
    messages,
    tokenize=False,
    add_generation_prompt=True,
    enable_thinking=True  # 默认为 True,设为 False 则禁用思考
)
# 模型推理,可以通过调整 temperature 等解码参数来调整生成质量
model_inputs = tokenizer([text], return_tensors="pt").to(model.device)
generated_ids = model.generate(
    **model_inputs,
    max_new_tokens=500,
    do_sample=True,
    temperature=1.0,
    top_p=0.3,
    repetition_penalty=1.3
)
# 处理输出
generated_ids = [
    output_ids[len(input_ids):] for input_ids, output_ids in zip(model_inputs.input_ids, generated_ids)
]
response = tokenizer.batch_decode(generated_ids, skip_special_tokens=False)[0]
print('\n')
print(response)
print('\n')

成功运行后,终端会输出如下内容:

fla-example-ch

此处以 RWKV7-1.5B-g1 为例,可选模型还有很多,详细请到 https://huggingface.co/fla-hub 查看。

使用 RWKV 组件

RWKV-FLA 提供了各种组件,可以单独使用以构建自定义模型架构。

使用 RWKV7 注意力层

from fla.layers.rwkv7 import RWKV7Attention

attention_layer = RWKV7Attention(
    mode=config.attn_mode,
    hidden_size=config.hidden_size,
    head_dim=config.head_dim,
    num_heads=config.num_heads,
    decay_low_rank_dim=config.decay_low_rank_dim,
    gate_low_rank_dim=config.gate_low_rank_dim,
    a_low_rank_dim=config.a_low_rank_dim,
    v_low_rank_dim=config.v_low_rank_dim,
    norm_eps=config.norm_eps,
    fuse_norm=config.fuse_norm,
    layer_idx=layer_idx,
    value_dim=config.value_dim[layer_idx],
    num_hidden_layers=config.num_hidden_layers
)

使用 RWKV7 前馈网络

from fla.models.rwkv7.modeling_rwkv7 import RWKV7FeedForward

ffn_layer = RWKV7FeedForward(
    hidden_size=config.hidden_size,
    hidden_ratio=config.hidden_ratio,
    intermediate_size=config.intermediate_size,
    hidden_act=config.hidden_act,
    layer_idx=layer_idx,
    num_hidden_layers=config.num_hidden_layers
)

注意事项

  • 不要将 RWKV-FLA 的 RWKV-7 训练结果当作官方参考实现的等价结果
  • 如遇到平台特定问题(如 Triton 版本、精度异常等),建议优先向相应硬件厂商反映
  • 不建议直接调用底层计算内核,除非有特殊需求并已充分了解代码实现

版本要求

  • Python 3.10 及以上版本
  • PyTorch 2.5 及以上版本
  • Triton 3.0 及以上版本
  • Transformers 4.45 及以上版本
这份文档对您有帮助吗?