File size: 2,370 Bytes
0af9165
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8f4239b
0af9165
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
{{ env.wgsl.resourceDeclarations }}

// A power-of-two lane cohort reduces one plane while several cohorts share a
// workgroup.
const HIDDEN: u32 = {{ hidden }}u;
const HIDDEN_V4: u32 = {{ hiddenVec }}u;
const CHANNELS: u32 = {{ channels }}u;
const EPSILON: f32 = {{ epsilon }};
const WG: u32 = {{ workgroupSize }}u;
const LANES: u32 = {{ lanesPerPlane }}u;
const PLANES_PER_WG: u32 = {{ planesPerWorkgroup }}u;

var<workgroup> shifted_moments: array<vec2<f32>, WG>;
var<workgroup> plane_shift: array<f32, PLANES_PER_WG>;

@compute @workgroup_size(WG, 1, 1)
fn main(
  @builtin(workgroup_id) workgroup: vec3<u32>,
  @builtin(local_invocation_id) local: vec3<u32>
) {
  let tid = local.x;
  let plane_in_workgroup = tid / LANES;
  let lane = tid % LANES;
  let group = workgroup.x + workgroup.y * {{ DISPATCH_FOLD_WIDTH }}u;
  let row = group * PLANES_PER_WG + plane_in_workgroup;
  let is_active = row < params.rows;
  let base = row * HIDDEN_V4;

  if (is_active && lane == 0u) {
    plane_shift[plane_in_workgroup] = f32(x[base].x);
  }
  workgroupBarrier();
  let shift = plane_shift[plane_in_workgroup];

  var moments = vec2<f32>(0.0);
  if (is_active) {
    for (var i = lane; i < HIDDEN_V4; i = i + LANES) {
      let value = vec4<f32>(x[base + i]);
      let delta = value - vec4<f32>(shift);
      moments.x = moments.x + delta.x + delta.y + delta.z + delta.w;
      moments.y = moments.y + dot(delta, delta);
    }
  }
  shifted_moments[tid] = moments;
  workgroupBarrier();

  var stride = LANES / 2u;
  loop {
    if (stride == 0u) { break; }
    if (lane < stride) {
      shifted_moments[tid] = shifted_moments[tid] + shifted_moments[tid + stride];
    }
    stride = stride / 2u;
    workgroupBarrier();
  }

  if (is_active) {
    let total = shifted_moments[plane_in_workgroup * LANES];
    let mean_delta = total.x / f32(HIDDEN);
    let variance = max(total.y / f32(HIDDEN) - mean_delta * mean_delta, 0.0);
    let mean = shift + mean_delta;
    let inv_std = inverseSqrt(variance + EPSILON);
    let channel = row % CHANNELS;
    let affine_scale = inv_std * f32(scale[channel]);
    let affine_bias = f32(bias[channel]);
    for (var i = lane; i < HIDDEN_V4; i = i + LANES) {
      let value = vec4<f32>(x[base + i]);
      y[base + i] = {{ vectorScalar }}((value - vec4<f32>(mean)) * vec4<f32>(affine_scale) + vec4<f32>(affine_bias));
    }
  }
}