// Elementwise scale: y = x * alpha (used for gradient clipping).
struct Params {
n: u32,
alpha: f32,
_pad0: u32,
_pad1: u32,
}
@group(0) @binding(0) var<uniform> p: Params;
@group(0) @binding(1) var<storage, read> x: array<f32>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@compute @workgroup_size(256)
fn main(
@builtin(workgroup_id) wid: vec3<u32>,
@builtin(num_workgroups) nwg: vec3<u32>,
@builtin(local_invocation_index) li: u32,
) {
let i = (wid.y * nwg.x + wid.x) * 256u + li;
if (i >= p.n) { return; }
y[i] = x[i] * p.alpha;
}