File size: 5,797 Bytes
2760d09
05e7ae4
 
 
 
2760d09
 
 
 
 
 
 
 
 
 
 
 
05e7ae4
 
 
2760d09
 
 
 
 
 
05e7ae4
2760d09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
05e7ae4
2760d09
 
 
 
 
 
 
 
 
 
 
 
 
05e7ae4
2760d09
 
 
 
 
05e7ae4
 
 
 
 
62736d3
05e7ae4
2760d09
 
 
 
 
 
62736d3
2760d09
 
 
 
 
05e7ae4
 
 
 
 
 
 
 
2760d09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
05e7ae4
 
 
 
 
 
 
 
2760d09
 
 
 
 
 
 
 
 
 
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
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
{% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
{% if op == "max" or op == "min" %}
{{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
{{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %}
{% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false, reuse=false) %}
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
  loop {
    if ({{ svar }} == 0u) {
      break;
    }
    if ({{ idx }} < {{ svar }}) {
{% for a in arrays %}
      {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
{% endfor %}
    }
    {{ svar }} = {{ svar }} / 2u;
    workgroupBarrier();
  }{% endmacro %}
/* Normalize residual = input + skip, with an optional bias. Reductions use
 * one workgroup per row; closed-form one-element rows use one invocation. */
{% set degenerateRow = (not simplified) and hiddenSize == 1 %}
{% if useSubgroups and not degenerateRow %}
enable subgroups;
{% endif %}
{{ env.wgsl.resourceDeclarations }}

{% if not degenerateRow or writeResidualSum or (writeMean is defined and writeMean) %}
const HIDDEN: u32 = {{ hiddenSize }}u;
{% endif %}
const WG: u32 = {{ workgroupSize }}u;
{% if not degenerateRow %}

var<workgroup> pair_partial: array<vec2<f32>, WG>;

{% if useSubgroups %}
fn reduce_pair(value: vec2<f32>, sg_lane: u32, sg_id: u32, num_sg: u32) -> vec2<f32> {
  let s = vec2<f32>(subgroupAdd(value.x), subgroupAdd(value.y));
  if (num_sg == 1u) {
    return s;
  }
  if (sg_lane == 0u) {
    pair_partial[sg_id] = s;
  }
  workgroupBarrier();
  var total = vec2<f32>(0.0, 0.0);
  for (var i = 0u; i < num_sg; i = i + 1u) {
    total = total + pair_partial[i];
  }
  return total;
}
{% else %}
fn reduce_pair(value: vec2<f32>, tid: u32) -> vec2<f32> {
  pair_partial[tid] = value;
  workgroupBarrier();
{{ wgsl_tree_fold(["pair_partial"], idx="tid", wg="WG", form="head") }}
  return pair_partial[0];
}
{% endif %}
{% endif %}
{% if not degenerateRow or writeResidualSum or (writeMean is defined and writeMean) %}

fn residual_value(row: u32, d: u32) -> f32 {
  let index = row * HIDDEN + d;
  var value = f32(input[index]) + f32(skip[index]);
{% if hasBias %}
  value = value + f32(bias[d]);
{% endif %}
  return value;
}
{% endif %}

@compute @workgroup_size(WG, 1, 1)
fn main(
  @builtin({{ "global_invocation_id" if degenerateRow else "workgroup_id" }}) {{ "gid" if degenerateRow else "wg" }}: vec3<u32>{% if not degenerateRow %},
  @builtin(local_invocation_id) lid: vec3<u32>{% endif %}{% if useSubgroups and not degenerateRow %},
  @builtin(subgroup_invocation_id) sg_lane: u32,
  @builtin(subgroup_id) sg_id: u32,
  @builtin(num_subgroups) num_sg: u32{% endif %}
) {
  // Fold the row grid across workgroups; independent rows also include the
  // invocation offset. The bounds guard drops the final dispatch tail.
{% if degenerateRow %}
  let row = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
{% else %}
  let row = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u;
{% endif %}
  if (row >= params.rows) {
    return;
  }
{% if not degenerateRow %}
  let tid = lid.x;
{% endif %}
{% if degenerateRow %}

  // HIDDEN == 1: the row's mean is its only element, so the centered value and
  // the variance are exactly zero and the output reduces to beta. The closed
  // form avoids computing that zero by subtracting two equal rounded values.
  let row_inv = inverseSqrt(params.epsilon);
{% if packedStatistics is defined and packedStatistics %}
  row_stats[row] = vec2<f32>(residual_value(row, 0u), row_inv);
{% elif writeMean is defined and writeMean %}
  mean[row] = residual_value(row, 0u);
{% endif %}
{% if writeInvStd is defined and writeInvStd and not (packedStatistics is defined and packedStatistics) %}
  inv_std_var[row] = row_inv;
{% endif %}
{% if writeResidualSum %}
  let residual = residual_value(row, 0u);
  input_skip_bias_sum[row] = {{ scalar }}(residual);
{% endif %}
  // 0.0 * row_inv keeps the IEEE result when epsilon == 0 makes row_inv +Inf.
  output[row] = {{ scalar }}(0.0 * row_inv * f32(gamma[0]){% if hasBeta %} + f32(beta[0]){% endif %});
{% else %}

  // Shifted moments: accumulating (x - x[0], (x - x[0])^2) keeps the sums
  // small for rows with a large common offset; every thread reconstructs the
  // row mean and variance from the merged pair.
  let shift = residual_value(row, 0u);
  var acc = vec2<f32>(0.0, 0.0);
  for (var d = tid; d < HIDDEN; d = d + WG) {
    let centered = residual_value(row, d) - shift;
    acc.x = acc.x + centered;
    acc.y = acc.y + centered * centered;
  }

{% if useSubgroups %}
  let totals = reduce_pair(acc, sg_lane, sg_id, num_sg);
{% else %}
  let totals = reduce_pair(acc, tid);
{% endif %}
  let mean_delta = totals.x / f32(HIDDEN);
  let row_mean = shift + mean_delta;
  let variance = max(totals.y / f32(HIDDEN) - mean_delta * mean_delta, 0.0);
  let row_inv = inverseSqrt(variance + params.epsilon);
{% if packedStatistics is defined and packedStatistics %}
  if (tid == 0u) { row_stats[row] = vec2<f32>(row_mean, row_inv); }
{% elif writeMean is defined and writeMean %}
  if (tid == 0u) { mean[row] = row_mean; }
{% endif %}
{% if writeInvStd is defined and writeInvStd and not (packedStatistics is defined and packedStatistics) %}
  if (tid == 0u) { inv_std_var[row] = row_inv; }
{% endif %}
  for (var d = tid; d < HIDDEN; d = d + WG) {
    let index = row * HIDDEN + d;
    let residual = residual_value(row, d);
{% if writeResidualSum %}
    input_skip_bias_sum[index] = {{ scalar }}(residual);
{% endif %}
    output[index] = {{ scalar }}((residual - row_mean) * row_inv * f32(gamma[d]){% if hasBeta %} + f32(beta[d]){% endif %});
  }
{% endif %}
}