1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
#include <metal_stdlib>
using namespace metal;
struct Params {
uint count;
float scale;
};
struct DataBuffer {
array<float, 64> values;
};
kernel void main(uint __luma_work_local_index [[thread_index_in_threadgroup]],
uint3 __luma_work_global_id [[thread_position_in_grid]],
uint3 __luma_work_local_id [[thread_position_in_threadgroup]],
uint3 __luma_work_group_id [[threadgroup_position_in_grid]],
device Params* params [[buffer(0)]],
device DataBuffer* data [[buffer(1)]]) {
threadgroup float scratch[64];
uint idx = __luma_work_local_index;
uint gid = __luma_work_global_id.x;
uint lid = __luma_work_local_id.x;
uint group_x = __luma_work_group_id.x;
if ((gid < params.count)) {
scratch[idx] = (data.values[idx] * params.scale);
threadgroup_barrier(mem_flags::mem_threadgroup);
data.values[idx] = (scratch[idx] + static_cast<float>((lid + group_x)));
}
}