Skip to main content

RMSPROP

Constant RMSPROP 

Source
pub const RMSPROP: &str = r#"
#include <metal_stdlib>
using namespace metal;

kernel void optirs_rmsprop(
    device float* x [[buffer(0)]],
    device const float* y [[buffer(1)]],
    device float* a [[buffer(2)]],
    device float* b [[buffer(3)]],
    device float* result [[buffer(4)]],
    device const float* output [[buffer(5)]],
    uint idx [[thread_position_in_grid]])
{
    uint n = as_type<uint>(output[6]);
    if (idx >= n) { return; }

    float lr = output[0];
    float alpha = output[1];
    float eps = output[2];
    float wd = output[3];
    float momentum = output[4];
    float centered = output[5];

    float p = x[idx];
    float g = y[idx];
    if (wd > 0.0f) { g = g + wd * p; }

    float sq = alpha * a[idx] + (1.0f - alpha) * g * g;
    a[idx] = sq;

    float avg = sq;
    if (centered > 0.5f) {
        float ga = alpha * b[idx] + (1.0f - alpha) * g;
        b[idx] = ga;
        avg = sq - ga * ga;
    }

    float denom = sqrt(max(avg, 0.0f)) + eps;

    if (momentum > 0.0f) {
        float buf = momentum * result[idx] + g / denom;
        result[idx] = buf;
        x[idx] = p - lr * buf;
    } else {
        x[idx] = p - lr * g / denom;
    }
}
"#;
Expand description

RMSprop with centering and momentum. See super::wgsl::RMSPROP.