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


环境变量

变量名 含义 取值 默认值 优先级
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.0.tar.gz (33.1 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.0-py3-none-any.whl (53.0 kB view details)

Uploaded Python 3

File details

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

File metadata

  • Download URL: rwkv_ops-0.1.0.tar.gz
  • Upload date:
  • Size: 33.1 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.0.tar.gz
Algorithm Hash digest
SHA256 d90caf64d668ff79da39be3df7356f5876402e5b9e3c8bbdfd76bd100e018548
MD5 fb5d382d026d5a619b32deb9b379169b
BLAKE2b-256 33cea5b3db777ebdcb23e7d601020275819fdcdb62c948bdbd0eaf3816efd611

See more details on using hashes here.

File details

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

File metadata

  • Download URL: rwkv_ops-0.1.0-py3-none-any.whl
  • Upload date:
  • Size: 53.0 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.0-py3-none-any.whl
Algorithm Hash digest
SHA256 0e52dd8fbe4890b045fb3e9a84e80e49f3a190168e018a27080a1a0025a2aeec
MD5 e6c2a134b0045d1f769b135f1e029681
BLAKE2b-256 82d6868c58c0dea29e30927985de5553f96d81033a6666dd37fe0612211c630c

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