Skip to main content

No project description provided

Project description

English Document

RWKV OPS 项目

由于 RWKV 将持续迭代,核心算子会随之更新。
本仓专门维护「算子」本身,不维护 layer 与 model;尽可能提供各框架的 GPU 算子。
目前:
• GPU 算子:PyTorch、JAX(TensorFlow 待 Google 支持 Triton 后上线)
• 原生算子:PyTorch、JAX、TensorFlow、NumPy
未来若 Keras 生态扩展,可能支持 MLX、OpenVINO。
注意:本库依赖 keras


安装方法

pip install rwkv_ops

环境变量

变量名 含义 取值 默认值 优先级
KERAS_BACKEND Keras 后端 jax / torch / tensorflow / numpy
KERNEL_BACKEND 算子后端 jax / torch / tensorflow / numpy torch
KERNEL_TYPE 实现类型 triton / cuda / native

KERNEL_BACKEND 有值,直接采用;若为空,则用 KERAS_BACKEND;两者皆空则默认 torch。
native 为原生算子,无 chunkwise,速度慢且显存高。


rwkv7op 使用方法

from rwkv_ops import generalized_delta_rule  # 或 from rwkv_ops import rwkv7_op,完全等价

def generalized_delta_rule(
    r,
    w,
    k,
    v,
    a,
    b,
    initial_state=None,
    output_final_state: bool = True,
    head_first: bool = False,
):
    """
    分块 Delta Rule 注意力接口。

    Args:
        q:  [B, T, H, K]
        k:  [B, T, H, K]
        v:  [B, T, H, V]
        a:  [B, T, H, K]
        b:  [B, T, H, K]
        gk: [B, T, H, K]  # decay term in log space!
        initial_state: 初始状态 [N, H, K, V],N 为序列数
        output_final_state: 是否返回最终状态
        head_first: 是否 head-first 格式,不支持变长

    Returns:
        o:           输出 [B, T, H, V] 或 [B, H, T, V]
        final_state: 最终状态 [N, H, K, V] 或 None
    """

torch-cuda下head-size也是一个kernel参数,默认是64. 若 head-size ≠ 64,请使用:

from rwkv_ops import get_generalized_delta_rule

generalized_delta_rule, RWKV7_USE_KERNEL = get_generalized_delta_rule(
    your_head_size, KERNEL_TYPE="cuda"
)

RWKV7_USE_KERNEL 为常量,标记是否使用 chunkwise 算子; 因为两者padding 处理逻辑不同,具体如下

if padding_mask is not None:
    if RWKV7_USE_KERNEL:
        w += (1 - padding_mask) * -1e9
    else:
        w = w * padding_mask + 1 - padding_mask

rwkv7op的实现状态

Framework cuda triton native
PyTorch
JAX
TensorFlow
NumPy

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

rwkv_ops-0.1.1.tar.gz (32.7 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

rwkv_ops-0.1.1-py3-none-any.whl (52.4 kB view details)

Uploaded Python 3

File details

Details for the file rwkv_ops-0.1.1.tar.gz.

File metadata

  • Download URL: rwkv_ops-0.1.1.tar.gz
  • Upload date:
  • Size: 32.7 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.1.0 CPython/3.10.14

File hashes

Hashes for rwkv_ops-0.1.1.tar.gz
Algorithm Hash digest
SHA256 10c630d31ee0bf79a961f2ca21f9591a3493a706a3a88d618d78c275d19c9265
MD5 b584690c6d85b57b291d8aacd535fbfe
BLAKE2b-256 dd0acc0a858b0466c58333243a56c670d592845cb6c859d5e7a92d19a33a8855

See more details on using hashes here.

File details

Details for the file rwkv_ops-0.1.1-py3-none-any.whl.

File metadata

  • Download URL: rwkv_ops-0.1.1-py3-none-any.whl
  • Upload date:
  • Size: 52.4 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.1.0 CPython/3.10.14

File hashes

Hashes for rwkv_ops-0.1.1-py3-none-any.whl
Algorithm Hash digest
SHA256 c13510ba6eb78302027736c77f18472d8a5938546275b5e7d4d9c8885e6f23e7
MD5 d00ad12a5acd8ce42097f35ab536bb03
BLAKE2b-256 491cc219d5cd3db4d5190d6929bad95be40dbd8d7437a0172dcc1c6f3d4cd77f

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page