File size: 3,359 Bytes
69e74ac
5ec57d5
 
 
 
69e74ac
 
 
 
 
 
 
 
 
 
5ec57d5
69e74ac
 
 
 
 
f1a8138
 
 
69e74ac
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f1a8138
69e74ac
f1a8138
69e74ac
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f1a8138
69e74ac
 
 
 
 
 
 
 
f1a8138
 
 
69e74ac
f1a8138
69e74ac
 
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
{% 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 %}
{% if useSubgroups %}
enable subgroups;
{% endif %}
{{ env.wgsl.resourceDeclarations }}

const HIDDEN: u32 = {{ hidden }}u;
const HIDDEN_V: u32 = {{ hiddenVec }}u;
const WG: u32 = {{ wg }}u;

var<workgroup> sg_partials: array<f32, WG>;

fn reduce_scalar(value: f32{% if useSubgroups %}, sg_lane: u32, sg_id: u32, num_sg: u32{% else %}, tid: u32{% endif %}) -> f32 {
{% if useSubgroups %}
  let s = subgroupAdd(value);
  if (num_sg == 1u) {
    return s;
  }
  if (sg_lane == 0u) {
    sg_partials[sg_id] = s;
  }
  workgroupBarrier();
  var total = 0.0;
  for (var i = 0u; i < num_sg; i = i + 1u) {
    total = total + sg_partials[i];
  }
  return total;
{% else %}
  // No-subgroup tier: workgroup barrier tree-reduction (WG is a power of two).
  sg_partials[tid] = value;
  workgroupBarrier();
{{ wgsl_tree_fold(["sg_partials"], idx="tid", wg="WG", form="head", breakInline=true) }}
  return sg_partials[0];
{% endif %}
}

// 4 contiguous residual elements (input[idx] + skip[skip_idx] [+ bias]) at vec4
// index `vi`. skip_idx == idx for the normal (non-broadcast) path; for a skip
// that broadcasts across the leading/batch dim uses a folded index.
fn residual_value(idx: u32, skip_idx: u32{% if hasBias %}, vi: u32{% endif %}) -> vec4<f32> {
  var value = vec4<f32>(input[idx]) + vec4<f32>(skip[skip_idx]);
{% if hasBias %}
  value = value + vec4<f32>(bias[vi]);
{% endif %}
  return value;
}

@compute @workgroup_size(WG, 1, 1)
fn main(
  @builtin(workgroup_id) wg_id: vec3<u32>,
  @builtin(local_invocation_id) lid: vec3<u32>{% if useSubgroups %},
  @builtin(subgroup_invocation_id) sg_lane: u32,
  @builtin(subgroup_id) sg_id: u32,
  @builtin(num_subgroups) num_sg: u32{% endif %}
) {
  let row = wg_id.x + wg_id.y * params.rowStride;
  if (row >= params.rows) {
    return;
  }
  let tid = lid.x;
  let base = row * HIDDEN_V;
  let skip_base = base;

  var acc = 0.0;
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
    let v = residual_value(base + i, skip_base + i{% if hasBias %}, i{% endif %});
    acc = acc + dot(v, v);
  }

  let total = reduce_scalar(acc{% if useSubgroups %}, sg_lane, sg_id, num_sg{% else %}, tid{% endif %});
  let row_inv = inverseSqrt(total / f32(HIDDEN) + params.epsilon);

  for (var i = tid; i < HIDDEN_V; i = i + WG) {
    let idx = base + i;
    let residual = residual_value(idx, skip_base + i{% if hasBias %}, i{% endif %});
{% if writeResidualSum %}
    input_skip_bias_sum[idx] = {{ vecType }}(residual);
{% endif %}
    output[idx] = {{ vecType }}(residual * row_inv * vec4<f32>(gamma[i]));
  }
}