Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 28 additions & 0 deletions backends/vulkan/runtime/graph/ops/glsl/common.glslh
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,34 @@
#define mod_4(x) ((x) & 3)
#define mod_8(x) ((x) & 7)

// Branch access avoids Adreno compiler crashes from runtime-indexed UBO vectors.
int safe_idx(const ivec4 v, const int idx) {
if (idx == 0) return v.x;
if (idx == 1) return v.y;
if (idx == 2) return v.z;
return v.w;
}

uint safe_idx(const uvec4 v, const int idx) {
if (idx == 0) return v.x;
if (idx == 1) return v.y;
if (idx == 2) return v.z;
return v.w;
}

int safe_idx(const ivec3 v, const int idx) {
if (idx == 0) return v.x;
if (idx == 1) return v.y;
return v.z;
}

void safe_set(inout ivec4 v, const int idx, const int val) {
if (idx == 0) { v.x = val; }
else if (idx == 1) { v.y = val; }
else if (idx == 2) { v.z = val; }
else { v.w = val; }
}


int sign_extend_8bit(const int val) {
if ((val & 0x80) != 0) {
Expand Down
63 changes: 16 additions & 47 deletions backends/vulkan/runtime/graph/ops/glsl/indexing.glslh
Original file line number Diff line number Diff line change
Expand Up @@ -80,40 +80,9 @@ bool is_channels_last(const int hashed_layout) {
#define get_outer_packed_dim_block_size(layout) (((layout) >> 28) & 0xF)
#define get_batch_concat_dim(layout) (((layout) >> 12) & 0xF)

// Safe ivec4 component access via if/else chain. Avoids dynamic vector
// indexing (v[index]) on UBO-backed ivec4 members, which crashes on Adreno 740
// when the index is a specialization constant with value 1 or 2.
int safe_idx(const ivec4 v, const int idx) {
if (idx == 0) return v.x;
if (idx == 1) return v.y;
if (idx == 2) return v.z;
return v.w;
}

// Safe uvec4 component access via if/else chain. Same rationale as safe_idx
// for ivec4 — avoids dynamic vector indexing on UBO-backed uvec4 members.
uint safe_idx(const uvec4 v, const int idx) {
if (idx == 0) return v.x;
if (idx == 1) return v.y;
if (idx == 2) return v.z;
return v.w;
}

// Safe ivec3 component access via if/else chain. Same rationale as safe_idx
// for ivec4.
int safe_idx(const ivec3 v, const int idx) {
if (idx == 0) return v.x;
if (idx == 1) return v.y;
return v.z;
}

// Safe ivec4 component write via if/else chain. Companion to safe_idx for
// cases where we need to set a component by a spec-const-derived index.
void safe_set(inout ivec4 v, const int idx, const int val) {
if (idx == 0) { v.x = val; }
else if (idx == 1) { v.y = val; }
else if (idx == 2) { v.z = val; }
else { v.w = val; }
uint safe_idx(const uvec4 lower, const uvec4 upper, const int idx) {
if (idx < 4) return safe_idx(lower, idx);
return safe_idx(upper, idx - 4);
}

//
Expand All @@ -140,19 +109,19 @@ uint numel(const BufferMetadata meta) {
}

uint dim_order_at(const BufferMetadata meta, const int dim) {
return meta.dim_order[div_4(dim)][mod_4(dim)];
return safe_idx(meta.dim_order[0], meta.dim_order[1], dim);
}

uint dim_order_at(const BufferMetadata meta, const uint dim) {
return meta.dim_order[div_4(dim)][mod_4(dim)];
return safe_idx(meta.dim_order[0], meta.dim_order[1], int(dim));
}

uint stride_at(const BufferMetadata meta, const int dim) {
return meta.strides[div_4(dim)][mod_4(dim)];
return safe_idx(meta.strides[0], meta.strides[1], dim);
}

uint stride_at(const BufferMetadata meta, const uint dim) {
return meta.strides[div_4(dim)][mod_4(dim)];
return safe_idx(meta.strides[0], meta.strides[1], int(dim));
}

uint width(const BufferMetadata meta) {
Expand All @@ -164,11 +133,11 @@ uint height(const BufferMetadata meta) {
}

uint size_at(const BufferMetadata meta, const int dim) {
return meta.sizes[div_4(dim)][mod_4(dim)];
return safe_idx(meta.sizes[0], meta.sizes[1], dim);
}

uint size_at(const BufferMetadata meta, const uint dim) {
return meta.sizes[div_4(dim)][mod_4(dim)];
return safe_idx(meta.sizes[0], meta.sizes[1], int(dim));
}

bool are_equal(const BufferMetadata meta1, const BufferMetadata meta2) {
Expand Down Expand Up @@ -322,8 +291,8 @@ TensorIndex linear_idx_to_tensor_idx(
int dim = int_ndim(meta);
int i = 0;
for (int d = max(dim - 1, 0); d >= 0; d--) {
uint dim_idx = meta.dim_order[div_4(d)][mod_4(d)];
uint dim_stride = meta.strides[div_4(dim_idx)][mod_4(dim_idx)];
uint dim_idx = dim_order_at(meta, d);
uint dim_stride = stride_at(meta, dim_idx);

tidx.data[div_4(dim_idx)][mod_4(dim_idx)] = linear_idx / dim_stride;
linear_idx = linear_idx % dim_stride;
Expand Down Expand Up @@ -424,7 +393,7 @@ TensorIndex4D texel_idx_to_tensor4d_idx(
uint remaining = block_idx;
[[unroll]] for (int d = 3; d >= 0; d--) {
int dim_idx = extract_4b(hashed_layout, d);
uint dim_stride = meta.strides[0][dim_idx];
uint dim_stride = safe_idx(meta.strides[0], dim_idx);
tidx.data[dim_idx] = int(remaining / dim_stride);
remaining = remaining % dim_stride;
}
Expand Down Expand Up @@ -562,7 +531,7 @@ uint tensor4d_idx_to_linear_idx(
const TensorIndex4D tidx) {
uint lin_idx = 0;
for (int d = 0; d < 4; ++d) {
lin_idx += meta.strides[0][d] * tidx.data[d];
lin_idx += safe_idx(meta.strides[0], d) * tidx.data[d];
}
return lin_idx;
}
Expand Down Expand Up @@ -637,7 +606,7 @@ int tensor4d_idx_to_block_idx(
// Compute block-space linear index
int block_idx = 0;
[[unroll]] for (int d = 0; d < 4; ++d) {
block_idx += int(meta.strides[0][d]) * tidx.data[d];
block_idx += int(safe_idx(meta.strides[0], d)) * tidx.data[d];
}
return block_idx;
}
Expand Down Expand Up @@ -732,10 +701,10 @@ int tensor_idx_to_block_idx(
// Compute block-space linear index over all 8 dims
int block_idx = 0;
[[unroll]] for (int d = 0; d < 4; ++d) {
block_idx += int(meta.strides[0][d]) * int(tidx.data[0][d]);
block_idx += int(safe_idx(meta.strides[0], d)) * int(tidx.data[0][d]);
}
[[unroll]] for (int d = 0; d < 4; ++d) {
block_idx += int(meta.strides[1][d]) * int(tidx.data[1][d]);
block_idx += int(safe_idx(meta.strides[1], d)) * int(tidx.data[1][d]);
}
return block_idx;
}
Expand Down
35 changes: 21 additions & 14 deletions backends/vulkan/runtime/graph/ops/glsl/indexing_utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@
#ifndef INDEXING_UTILS_H
#define INDEXING_UTILS_H

#include "common.glslh"

/*
* The functions defined in this header file use the following shorthand to
* represent tensor related data structures.
Expand Down Expand Up @@ -94,7 +96,7 @@ ivec4 tidx_to_nchwi(const ivec4 tidx, const ivec4 sizes, const int packed_dim) {
int base_i = tidx.x * strides.x + tidx.y * strides.y + tidx.z * strides.z +
tidx.w * strides.w;

return base_i + ivec4(0, 1, 2, 3) * strides[packed_dim];
return base_i + ivec4(0, 1, 2, 3) * safe_idx(strides, packed_dim);
}

/*
Expand Down Expand Up @@ -135,9 +137,10 @@ int tidx_to_nchwi(const ivec4 tidx, const ivec4 sizes) {
ivec4 bufi_to_tidx(int bufi, const ivec4 strides, const ivec4 dim_order) {
ivec4 idx;
for (int i = 3; i >= 0; i--) {
int dim = dim_order[i];
idx[dim] = bufi / strides[dim];
bufi %= strides[dim];
int dim = safe_idx(dim_order, i);
int dim_stride = safe_idx(strides, dim);
idx[dim] = bufi / dim_stride;
bufi %= dim_stride;
}
return idx;
}
Expand All @@ -148,8 +151,9 @@ ivec4 bufi_to_tidx(int bufi, const ivec4 strides, const ivec4 dim_order) {
ivec4 contiguous_bufi_to_tidx(int bufi, const ivec4 strides) {
ivec4 idx;
for (int i = 3; i >= 0; i--) {
idx[i] = bufi / strides[i];
bufi %= strides[i];
int dim_stride = safe_idx(strides, i);
idx[i] = bufi / dim_stride;
bufi %= dim_stride;
}
return idx;
}
Expand All @@ -165,15 +169,16 @@ ivec4 lpos_to_tidx(
const int batch_inner_dim,
const int packed_dim) {
// Align packed dim to next multiple of 4 to account for texel padding
sizes[packed_dim] = alignup4(sizes[packed_dim]);
safe_set(sizes, packed_dim, alignup4(safe_idx(sizes, packed_dim)));
// Moving 1 texel along the packed dim traverses 4 tensor elements
lpos[packed_dim] *= 4;

ivec4 tidx = ivec4(lpos, 0);

if (sizes.w > 1) {
tidx.w = tidx[batch_inner_dim] / sizes[batch_inner_dim];
tidx[batch_inner_dim] %= sizes[batch_inner_dim];
int batch_inner_size = safe_idx(sizes, batch_inner_dim);
tidx.w = tidx[batch_inner_dim] / batch_inner_size;
tidx[batch_inner_dim] %= batch_inner_size;
}
return tidx;
}
Expand All @@ -184,13 +189,13 @@ ivec3 tidx_to_lpos(
const int batch_inner_dim,
const int packed_dim) {
// Align packed dim to next multiple of 4 to account for texel padding
sizes[packed_dim] = alignup4(sizes[packed_dim]);
safe_set(sizes, packed_dim, alignup4(safe_idx(sizes, packed_dim)));

ivec3 lpos = tidx.xyz;

// Adjust batch inner dim by batch index if needed
if (sizes.w > 1) {
lpos[batch_inner_dim] += tidx.w * sizes[batch_inner_dim];
lpos[batch_inner_dim] += tidx.w * safe_idx(sizes, batch_inner_dim);
}
// Fast division by 4, since moving 1 texel along the packed dim traverses 4
// tensor elements.
Expand All @@ -204,7 +209,7 @@ ivec3 tidx_to_pos(
const ivec4 axis_map,
const int packed_dim) {
// Align packed dim to next multiple of 4 to account for texel padding
sizes[packed_dim] = alignup4(sizes[packed_dim]);
safe_set(sizes, packed_dim, alignup4(safe_idx(sizes, packed_dim)));

ivec3 pos;
for (int dim = 0; dim < 3; ++dim) {
Expand All @@ -213,11 +218,13 @@ ivec3 tidx_to_pos(

// Adjust batch inner dim by batch index if needed
if (sizes.w > 1) {
pos[axis_map[axis_map.w]] += tidx.w * sizes[axis_map.w];
int batch_inner_dim = axis_map.w;
pos[safe_idx(axis_map, batch_inner_dim)] +=
tidx.w * safe_idx(sizes, batch_inner_dim);
}
// Fast division by 4, since moving 1 texel along the packed dim traverses 4
// tensor elements.
pos[axis_map[packed_dim]] >>= 2;
pos[safe_idx(axis_map, packed_dim)] >>= 2;
return pos;
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,7 @@ ivec4 read_texel(ivec4 tidx) {
ivec4 out_tex = ivec4(0);

[[unroll]] for (int i = 0; i < 4; ++i) {
if (tidx[packed_dim] + i < sizes[packed_dim]) {
if (tidx[packed_dim] + i < safe_idx(sizes, packed_dim)) {
const int in_texel = nchw_in[buf_indices[i] >> 2];
int extracted_val = (in_texel >> (8 * (buf_indices[i] & 3))) & mask;
extracted_val = extend_sign(extracted_val);
Expand Down
2 changes: 1 addition & 1 deletion backends/vulkan/runtime/graph/ops/glsl/pad_buffer.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ void main() {
// value that fails the out_of_bounds check below.
TensorIndex in_tidx = out_tidx;
[[unroll]] for (int d = 0; d < 4; d++) {
in_tidx.data[0][d] -= uint(pad_per_dim[d]);
in_tidx.data[0][d] -= uint(safe_idx(pad_per_dim, d));
}

if (out_of_bounds(in_tidx, inp)) {
Expand Down
2 changes: 1 addition & 1 deletion backends/vulkan/runtime/graph/ops/glsl/reduce.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ reduce:
shader_variants:
- NAME: sum
- NAME: mean
POSTPROCESS: (accum / tin_sizes[reduce_dim])
POSTPROCESS: (accum / safe_idx(tin_sizes, reduce_dim))
- NAME: amax
INIT_ACCUM: first_val
UPDATE_ACCUM: max(accum, new_val)
Expand Down
2 changes: 1 addition & 1 deletion backends/vulkan/runtime/graph/ops/glsl/reduce2d.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ reduce2d:
shader_variants:
- NAME: sum2d
- NAME: mean2d
POSTPROCESS: (accum / (tin_sizes[reduce_dim1] * tin_sizes[reduce_dim2]))
POSTPROCESS: (accum / (safe_idx(tin_sizes, reduce_dim1) * safe_idx(tin_sizes, reduce_dim2)))
- NAME: amax2d
INIT_ACCUM: first_val
UPDATE_ACCUM: max(accum, new_val)
Expand Down
2 changes: 1 addition & 1 deletion backends/vulkan/runtime/graph/ops/glsl/var_buffer.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@ void main() {
shared_count[tid] = 0;
barrier();

const int R = in_sizes[reduce_dim];
const int R = safe_idx(in_sizes, reduce_dim);
const uint N = gl_WorkGroupSize[reduce_dim];

// Each workgroup processes a contiguous chunk of the input tensor
Expand Down
Loading