// glp_project.metal — per-row projection of activations off a direction.
//
// For each row m of a `[M, N]` row-major f32 activation matrix:
//
// row_m <- row_m - alpha * (row_m . d) * d / ‖d‖²
//
// where `d` is the GLP direction for the layer. The host passes ‖d‖² so
// the kernel does not require a normalized vector.
//
// Launch contract (do not violate — see the S1 fix history): the
// reduction uses a fixed 256-lane threadgroup (`partial[256]`, stride-128
// tree) and threadgroup id maps directly to the row. The host MUST
// dispatch threads_per_threadgroup = 256 and grid = M * 256 (one
// threadgroup per row). Columns are covered strided
// (`col = tid; col < n; col += 256`), so any row width N is correct —
// widths below 256 leave lanes idle with zero contribution; widths above
// 256 are covered by the stride. A host that passes the row width as the
// threadgroup size indexes `partial` out of bounds for N > 256 and
// exceeds Metal's threadgroup limit for N > 1024.
#include <metal_stdlib>
using namespace metal;
struct GlpProjectParams {
uint m;
uint n;
float alpha;
float d_norm_sq;
};
kernel void glp_project_f32(
constant GlpProjectParams& params [[buffer(0)]],
device const float* direction [[buffer(1)]],
device float* hidden [[buffer(2)]],
uint tid_in_tg [[thread_position_in_threadgroup]],
uint tg_size [[threads_per_threadgroup]],
uint tg_id [[threadgroup_position_in_grid]]
) {
const uint row = tg_id; // one threadgroup per row (grid = m * 256)
if (row >= params.m) {
return;
}
device float* row_ptr = hidden + row * params.n;
// accumulate strided elements of THIS row
float local_dot = 0.0f;
for (uint col = tid_in_tg; col < params.n; col += tg_size) {
local_dot += row_ptr[col] * direction[col];
}
// threadgroup tree reduction on the 256-thread window
threadgroup float partial[256];
partial[tid_in_tg] = local_dot;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = 128; stride > 0; stride >>= 1) {
if (tid_in_tg < stride) {
partial[tid_in_tg] += partial[tid_in_tg + stride];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
float dot = partial[0];
float scale = params.alpha * dot / params.d_norm_sq;
for (uint col = tid_in_tg; col < params.n; col += tg_size) {
row_ptr[col] -= scale * direction[col];
}
}