Skip to content
Merged
110 changes: 110 additions & 0 deletions backends/vulkan/runtime/graph/ops/glsl/fused_ce.glsl
Original file line number Diff line number Diff line change
@@ -0,0 +1,110 @@
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
*
* This source code is licensed under the BSD-style license found in the
* LICENSE file in the root directory of this source tree.
*/

#version 450 core

${define_required_extensions(STORAGE, DTYPE)}

#define PRECISION ${PRECISION}
#define T ${buffer_scalar_type(DTYPE)}

${define_active_storage_type(STORAGE)}

layout(std430) buffer;

${layout_declare_tensor(B, "w", "t_dlogits", DTYPE, "buffer")}
${layout_declare_tensor(B, "w", "t_loss_partial", DTYPE, "buffer")}
${layout_declare_tensor(B, "r", "t_logits", DTYPE, "buffer")}
${layout_declare_tensor(B, "r", "t_labels", "int", "buffer")}

${layout_declare_ubo(B, "int", "vocab")}
${layout_declare_ubo(B, "int", "n_rows")}
${layout_declare_ubo(B, "float", "n_valid")}

layout(local_size_x_id = 0, local_size_y_id = 1, local_size_z_id = 2) in;

#define NWORKERS 64

shared float red_m[NWORKERS];
shared float red_l[NWORKERS];

// Fused cross-entropy: one workgroup per row cooperatively reduces the vocab
// dimension with a single-pass online softmax (running max + rescaled running
// sum), then writes per-row loss and dlogits. Rows with label < 0 are masked.
void main() {
const uint row = gl_GlobalInvocationID.y;
if (int(row) >= n_rows) {
return;
}

const uint tid = gl_LocalInvocationID.x;
const uint V = uint(vocab);
const uint base = row * V;
const int lbl = t_labels[row];

if (lbl < 0) {
for (uint j = tid; j < V; j += NWORKERS) {
t_dlogits[base + j] = T(0);
}
if (tid == 0u) {
t_loss_partial[row] = T(0);
}
return;
}

// Single read pass: maintain a running max m and the sum l of exp(x - m),
// rescaling l whenever a larger value updates m. Finite -3.4e38 init.
float m = -3.4e38;
float l = 0.0;
for (uint j = tid; j < V; j += NWORKERS) {
float x = float(t_logits[base + j]);
if (x > m) {
l = l * exp(m - x) + 1.0;
m = x;
} else {
l = l + exp(x - m);
}
}
red_m[tid] = m;
red_l[tid] = l;
memoryBarrierShared();
barrier();

// Tree-combine the (m, l) pairs: m becomes the max, l is rescaled to it.
for (uint s = NWORKERS / 2u; s > 0u; s >>= 1u) {
if (tid < s) {
float ma = red_m[tid];
float la = red_l[tid];
float mb = red_m[tid + s];
float lb = red_l[tid + s];
float mm = max(ma, mb);
red_m[tid] = mm;
red_l[tid] = la * exp(ma - mm) + lb * exp(mb - mm);
}
memoryBarrierShared();
barrier();
}

const float row_max = red_m[0];
const float denom = red_l[0];
const float inv = 1.0 / denom;
const float scale = 1.0 / n_valid;

if (tid == 0u) {
const float lse = row_max + log(denom);
t_loss_partial[row] = T((lse - float(t_logits[base + uint(lbl)])) * scale);
}

for (uint j = tid; j < V; j += NWORKERS) {
float g = exp(float(t_logits[base + j]) - row_max) * inv * scale;
if (j == uint(lbl)) {
g = g - scale;
}
t_dlogits[base + j] = T(g);
}
}
15 changes: 15 additions & 0 deletions backends/vulkan/runtime/graph/ops/glsl/fused_ce.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

fused_ce:
parameter_names_with_default_values:
DTYPE: float
STORAGE: buffer
generate_variant_forall:
DTYPE:
- VALUE: float
shader_variants:
- NAME: fused_ce_buffer
55 changes: 55 additions & 0 deletions backends/vulkan/runtime/graph/ops/glsl/fused_ce_sum.glsl
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
*
* This source code is licensed under the BSD-style license found in the
* LICENSE file in the root directory of this source tree.
*/

#version 450 core

${define_required_extensions(STORAGE, DTYPE)}

#define PRECISION ${PRECISION}
#define T ${buffer_scalar_type(DTYPE)}

${define_active_storage_type(STORAGE)}

layout(std430) buffer;

${layout_declare_tensor(B, "w", "t_loss", DTYPE, "buffer")}
${layout_declare_tensor(B, "r", "t_loss_partial", DTYPE, "buffer")}

${layout_declare_ubo(B, "int", "n_rows")}

layout(local_size_x_id = 0, local_size_y_id = 1, local_size_z_id = 2) in;

#define NWORKERS 64

shared float red[NWORKERS];

// Self-contained [N] -> [1] tree-sum of the per-row losses in one workgroup, so
// fused_ce carries no cross-op reduce dependency.
void main() {
const uint tid = gl_LocalInvocationID.x;

float s = 0.0;
for (uint j = tid; j < uint(n_rows); j += NWORKERS) {
s += float(t_loss_partial[j]);
}
red[tid] = s;
memoryBarrierShared();
barrier();

for (uint k = NWORKERS / 2u; k > 0u; k >>= 1u) {
if (tid < k) {
red[tid] += red[tid + k];
}
memoryBarrierShared();
barrier();
}

if (tid == 0u) {
t_loss[0] = T(red[0]);
}
}
15 changes: 15 additions & 0 deletions backends/vulkan/runtime/graph/ops/glsl/fused_ce_sum.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

fused_ce_sum:
parameter_names_with_default_values:
DTYPE: float
STORAGE: buffer
generate_variant_forall:
DTYPE:
- VALUE: float
shader_variants:
- NAME: fused_ce_sum_buffer
78 changes: 78 additions & 0 deletions backends/vulkan/runtime/graph/ops/glsl/linear_dW.glsl
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
*
* This source code is licensed under the BSD-style license found in the
* LICENSE file in the root directory of this source tree.
*/

#version 450 core

#define PRECISION ${PRECISION}

#define T ${texel_load_component_type(DTYPE, STORAGE)}

#define TILE_N 4
#define TILE_K 4

${define_required_extensions(STORAGE, DTYPE)}

layout(std430) buffer;

${layout_declare_tensor(B, "w", "t_dw", DTYPE, STORAGE, is_scalar_array=True)}
${layout_declare_tensor(B, "r", "t_dout", DTYPE, STORAGE, is_scalar_array=True)}
${layout_declare_tensor(B, "r", "t_x", DTYPE, STORAGE, is_scalar_array=True)}

${layout_declare_ubo(B, "ivec4", "dout_sizes")}
${layout_declare_ubo(B, "ivec4", "x_sizes")}

layout(local_size_x_id = 0, local_size_y_id = 1, local_size_z_id = 2) in;

void main() {
// dW[N, K] = sum_m d_out[m, N] * x[m, K]; contraction over the flattened M.
const int N = dout_sizes.x;
const int M = dout_sizes.y * dout_sizes.z * dout_sizes.w;
const int K = x_sizes.x;

const int nnt = (N + TILE_N - 1) / TILE_N;
const int nkt = (K + TILE_K - 1) / TILE_K;
const int tiles = nnt * nkt;

const int tile_idx = int(gl_GlobalInvocationID.x);
if (tile_idx >= tiles) {
return;
}

const int n0 = (tile_idx / nkt) * TILE_N;
const int k0 = (tile_idx % nkt) * TILE_K;

float acc[TILE_N * TILE_K];
for (int i = 0; i < TILE_N * TILE_K; ++i) {
acc[i] = 0.0;
}

for (int m = 0; m < M; ++m) {
float dout_reg[TILE_N];
for (int nl = 0; nl < TILE_N; ++nl) {
const int n_eff = min(n0 + nl, N - 1);
dout_reg[nl] = float(t_dout[m * N + n_eff]);
}
for (int kl = 0; kl < TILE_K; ++kl) {
const int k_eff = min(k0 + kl, K - 1);
const float xv = float(t_x[m * K + k_eff]);
for (int nl = 0; nl < TILE_N; ++nl) {
acc[nl * TILE_K + kl] += dout_reg[nl] * xv;
}
}
}

for (int nl = 0; nl < TILE_N; ++nl) {
const int n = n0 + nl;
for (int kl = 0; kl < TILE_K; ++kl) {
const int k = k0 + kl;
if (n < N && k < K) {
t_dw[n * K + k] = T(acc[nl * TILE_K + kl]);
}
}
}
}
17 changes: 17 additions & 0 deletions backends/vulkan/runtime/graph/ops/glsl/linear_dW.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

linear_dW:
parameter_names_with_default_values:
DTYPE: float
STORAGE: buffer
generate_variant_forall:
STORAGE:
- VALUE: buffer
DTYPE:
- VALUE: float
shader_variants:
- NAME: linear_dW
Loading
Loading