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.