Skip to content
Merged
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
87 changes: 43 additions & 44 deletions backends/vulkan/op_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -1435,58 +1435,56 @@ def register_where():
# =============================================================================


@update_features(exir_ops.edge.aten.index.Tensor)
def register_index_tensor():
def _index_tensor_shapes(node: torch.fx.Node):
"""(self_val, index_val) for the supported single-index form, else None."""
self_arg = node.args[0]
indices = node.args[1]

if not isinstance(self_arg, torch.fx.Node):
return None
self_val = self_arg.meta.get("val", None)
if self_val is None:
return None

# Only support exactly one non-None index tensor, applied to dim 0.
if not isinstance(indices, (list, tuple)):
return None
non_none = [idx for idx in indices if idx is not None]
if len(non_none) != 1 or indices[0] is None:
return None
index_arg = non_none[0]
if not isinstance(index_arg, torch.fx.Node):
return None
index_val = index_arg.meta.get("val", None)
if index_val is None:
return None
def _index_tensor_shapes(node: torch.fx.Node):
"""Return self, index, and axis for the supported form, else None."""
self_arg = node.args[0]
indices = node.args[1]

if not isinstance(self_arg, torch.fx.Node):
return None
self_val = self_arg.meta.get("val", None)
if self_val is None or not isinstance(indices, (list, tuple)):
return None

non_none = [(dim, index) for dim, index in enumerate(indices) if index is not None]
if len(non_none) != 1:
return None
index_dim, index_arg = non_none[0]
if index_dim >= len(self_val.size()) or not isinstance(index_arg, torch.fx.Node):
return None
index_val = index_arg.meta.get("val", None)
if index_val is None:
return None

return self_val, index_val, index_dim


def _check_index_tensor_node(node: torch.fx.Node) -> bool:
shapes = _index_tensor_shapes(node)
if shapes is None:
return False
_, index_val, _ = shapes
# The gather is expressed as "one index position per output slice", so
# the index must be 1-D. `self` may be any rank.
return len(index_val.size()) == 1

return self_val, index_val

def check_index_tensor_node(node: torch.fx.Node) -> bool:
shapes = _index_tensor_shapes(node)
if shapes is None:
return False
_, index_val = shapes
# The gather is expressed as "one index position per output slice", so
# the index must be 1-D. `self` may be any rank: the buffer shader
# copies self's trailing dims through unchanged.
return len(index_val.size()) == 1
def _pick_index_tensor_storage(node: torch.fx.Node):
shapes = _index_tensor_shapes(node)
# Only the buffer shader handles a higher-rank `self`.
if shapes is not None and len(shapes[0].size()) > 1:
return utils.CONTIGUOUS_BUFFER, utils.CONTIGUOUS_BUFFER
return utils.ANY_STORAGE, utils.ANY_STORAGE

def pick_index_tensor_storage(node: torch.fx.Node):
shapes = _index_tensor_shapes(node)
# Only the buffer shader handles a higher-rank `self`; the texture
# variant still assumes the 1-D form (it reads self[idx, 0, 0, 0]).
if shapes is not None and len(shapes[0].size()) > 1:
return utils.CONTIGUOUS_BUFFER, utils.CONTIGUOUS_BUFFER
return utils.ANY_STORAGE, utils.ANY_STORAGE

@update_features(exir_ops.edge.aten.index.Tensor)
def register_index_tensor():
return OpFeatures(
inputs_storage=utils.ANY_STORAGE,
inputs_dtypes=utils.FP_INT_T,
supports_resize=True,
are_node_inputs_supported_fn=check_index_tensor_node,
pick_io_storage_fn=pick_index_tensor_storage,
are_node_inputs_supported_fn=_check_index_tensor_node,
pick_io_storage_fn=_pick_index_tensor_storage,
)


Expand All @@ -1500,6 +1498,7 @@ def register_arange():
return OpFeatures(
inputs_storage=utils.ANY_STORAGE,
inputs_dtypes=utils.FP_INT_T,
supports_resize=True,
)


Expand Down
16 changes: 13 additions & 3 deletions backends/vulkan/runtime/graph/ops/glsl/arange_buffer.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -23,18 +23,28 @@ layout(std430) buffer;
${layout_declare_tensor(B, "w", "t_out", DTYPE, "buffer")}

${layout_declare_ubo(B, "BufferMetadata", "outp")}
${layout_declare_ubo(B, "float", "start")}
${layout_declare_ubo(B, "float", "step")}
${layout_declare_ubo(B, "uint", "start")}
${layout_declare_ubo(B, "uint", "step")}

layout(push_constant) uniform restrict Block {
ivec2 params_are_int;
};

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

#include "dispatch.glslh"

float decode_param(const uint value, const int is_int) {
return is_int != 0 ? float(int(value)) : uintBitsToFloat(value);
}

void main() {
const uint out_bufi = linear_idx_from_gid();
if (out_of_bounds(out_bufi, outp)) {
return;
}

t_out[out_bufi] = T(start + out_bufi * step);
const float start_val = decode_param(start, params_are_int.x);
const float step_val = decode_param(step, params_are_int.y);
t_out[out_bufi] = T(start_val + out_bufi * step_val);
}
16 changes: 13 additions & 3 deletions backends/vulkan/runtime/graph/ops/glsl/arange_texture.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -23,14 +23,22 @@ layout(std430) buffer;
${layout_declare_tensor(B, "w", "t_out", DTYPE, "texture3d")}

${layout_declare_ubo(B, "TextureMetadata", "outp")}
${layout_declare_ubo(B, "float", "start")}
${layout_declare_ubo(B, "float", "step")}
${layout_declare_ubo(B, "uint", "start")}
${layout_declare_ubo(B, "uint", "step")}

layout(push_constant) uniform restrict Block {
ivec2 params_are_int;
};

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

${layout_declare_spec_const(C, "int", "out_layout", "CONTIG_LAYOUT_INT")}
const int packed_dim = get_packed_dim(out_layout);

float decode_param(const uint value, const int is_int) {
return is_int != 0 ? float(int(value)) : uintBitsToFloat(value);
}

void main() {
const ivec3 out_pos = ivec3(gl_GlobalInvocationID);

Expand All @@ -44,11 +52,13 @@ void main() {
// arange output is 1D, so the W dimension holds the element index.
// Compute the value for each element in the texel along the packed dim.
VEC4_T outtex = VEC4_T(0);
const float start_val = decode_param(start, params_are_int.x);
const float step_val = decode_param(step, params_are_int.y);
int limit = min(
4, safe_idx(outp.sizes, packed_dim) - out_tidx.data[packed_dim]);
for (int comp = 0; comp < limit; comp++) {
int elem_idx = out_tidx.data[0]; // W index is the linear element index
outtex[comp] = VEC4_T(start + elem_idx * step).x;
outtex[comp] = VEC4_T(start_val + elem_idx * step_val).x;
out_tidx.data[packed_dim]++;
}

Expand Down
88 changes: 68 additions & 20 deletions backends/vulkan/runtime/graph/ops/glsl/index_tensor_buffer.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
#version 450 core

${define_required_extensions("buffer", DTYPE)}
${define_required_extensions(INDEX_STORAGE, "int")}

#define PRECISION ${PRECISION}

Expand All @@ -22,19 +23,68 @@ layout(std430) buffer;

${layout_declare_tensor(B, "w", "t_out", DTYPE, "buffer")}
${layout_declare_tensor(B, "r", "t_self", DTYPE, "buffer")}
${layout_declare_tensor(B, "r", "t_index", "int", "buffer")}
${layout_declare_tensor(B, "r", "t_index", "int", INDEX_STORAGE)}

${layout_declare_ubo(B, "BufferMetadata", "outp")}
${layout_declare_ubo(B, "BufferMetadata", "inp")}
${layout_declare_ubo(B, "BufferMetadata", "index")}
$if INDEX_STORAGE == "buffer":
${layout_declare_ubo(B, "BufferMetadata", "index")}
$else:
${layout_declare_ubo(B, "TextureMetadata", "index")}

layout(push_constant) uniform restrict Block {
ivec2 index_params;
};

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

#include "dispatch.glslh"

// Implements aten.index.Tensor for the case where self is 1D and there is
// exactly one index tensor. Each output element is:
// Implements aten.index.Tensor with exactly one index tensor. Each output
// element is:
// output[...] = self[index[...]]
${layout_declare_spec_const(C, "int", "out_layout", "CONTIG_LAYOUT_INT")}
${layout_declare_spec_const(C, "int", "inp_layout", "CONTIG_LAYOUT_INT")}
${layout_declare_spec_const(C, "int", "index_layout", "CONTIG_LAYOUT_INT")}

int load_index(const TensorIndex out_tidx) {
$if INDEX_STORAGE == "buffer":
uint index_bufi = 0;
for (int d = 0; d < index_params.y; ++d) {
index_bufi +=
stride_at(index, d) * idx_at(out_tidx, index_params.x + d);
}
return t_index[index_bufi];
$else:
TensorIndex4D index_tidx = zero_tensor4d_idx();
index_tidx.data.x = int(idx_at(out_tidx, index_params.x));
if (index_params.y > 1) {
index_tidx.data.y = int(idx_at(out_tidx, index_params.x + 1));
}
if (index_params.y > 2) {
index_tidx.data.z = int(idx_at(out_tidx, index_params.x + 2));
}
if (index_params.y > 3) {
index_tidx.data.w = int(idx_at(out_tidx, index_params.x + 3));
}
const TextureElementIndex index_elem =
tensor4d_idx_to_texture_element_idx_simple(
index, index_tidx, index_layout);
return texelFetch(t_index, index_elem.pos, 0)[index_elem.comp];
}

uint self_idx_at(
const TensorIndex out_tidx,
const int self_axis,
const uint index_value) {
if (self_axis == index_params.x) {
return index_value;
}
const int out_axis = self_axis < index_params.x
? self_axis
: self_axis + index_params.y - 1;
return idx_at(out_tidx, out_axis);
}

void main() {
const uint out_bufi = linear_idx_from_gid();
Expand All @@ -45,22 +95,20 @@ void main() {
// Convert output buffer index to tensor index
TensorIndex out_tidx = linear_idx_to_tensor_idx(outp, out_bufi);

const uint self_rank = ndim(inp);
const uint index_rank = ndim(index);
// WHCN order places self's trailing axes before the index axes.
const uint index_axis_offset = self_rank - 1;

uint index_bufi = 0;
for (uint d = 0; d < index_rank; ++d) {
index_bufi +=
stride_at(index, d) * idx_at(out_tidx, index_axis_offset + d);
}
const int idx = t_index[index_bufi];

uint self_bufi = stride_at(inp, self_rank - 1) * uint(idx);
for (uint d = 0; d + 1 < self_rank; ++d) {
self_bufi += stride_at(inp, d) * idx_at(out_tidx, d);
}
const int idx = load_index(out_tidx);

TensorIndex self_tidx;
initialize(self_tidx);
const int self_rank = int_ndim(inp);
if (self_rank > 0) self_tidx.data[0].x = self_idx_at(out_tidx, 0, uint(idx));
if (self_rank > 1) self_tidx.data[0].y = self_idx_at(out_tidx, 1, uint(idx));
if (self_rank > 2) self_tidx.data[0].z = self_idx_at(out_tidx, 2, uint(idx));
if (self_rank > 3) self_tidx.data[0].w = self_idx_at(out_tidx, 3, uint(idx));
if (self_rank > 4) self_tidx.data[1].x = self_idx_at(out_tidx, 4, uint(idx));
if (self_rank > 5) self_tidx.data[1].y = self_idx_at(out_tidx, 5, uint(idx));
if (self_rank > 6) self_tidx.data[1].z = self_idx_at(out_tidx, 6, uint(idx));
if (self_rank > 7) self_tidx.data[1].w = self_idx_at(out_tidx, 7, uint(idx));
const uint self_bufi = tensor_idx_to_linear_idx(inp, self_tidx);

t_out[out_bufi] = t_self[self_bufi];
}
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,11 @@ index_tensor_buffer:
parameter_names_with_default_values:
DTYPE: float
STORAGE: buffer
INDEX_STORAGE: buffer
generate_variant_forall:
INDEX_STORAGE:
- VALUE: buffer
- VALUE: texture3d
DTYPE:
- VALUE: half
- VALUE: float
Expand Down
33 changes: 28 additions & 5 deletions backends/vulkan/runtime/graph/ops/glsl/unary_op.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -23,23 +23,35 @@ ${define_active_storage_type(STORAGE)}

layout(std430) buffer;

${layout_declare_tensor(0, "w", "t_out", DTYPE, STORAGE)}
${layout_declare_tensor(1, "r", "t_in", DTYPE, STORAGE)}
${layout_declare_tensor(B, "w", "t_out", DTYPE, STORAGE)}
${layout_declare_tensor(B, "r", "t_in", DTYPE, STORAGE)}

$if DYNAMIC_PARAMS:
${layout_declare_ubo(B, "uint", "minimum")}
${layout_declare_ubo(B, "uint", "maximum")}

layout(push_constant) uniform restrict Block {
$if STORAGE == "buffer":
int numel;
$else:
ivec4 out_limits;
float minimum;
float maximum;
$if DYNAMIC_PARAMS:
ivec2 bounds_are_int;
$else:
float minimum;
float maximum;
};

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

#include "dispatch.glslh"
#include "activations.h"

$if DYNAMIC_PARAMS:
float decode_bound(const uint value, const int is_int) {
return is_int != 0 ? float(int(value)) : uintBitsToFloat(value);
}

#ifdef USING_BUFFER

void main() {
Expand All @@ -48,7 +60,13 @@ void main() {
return;
}

float in_val = float(t_in[i]);
$if DYNAMIC_PARAMS:
const T in_val = T(t_in[i]);
const T minimum_val = T(decode_bound(minimum, bounds_are_int.x));
const T maximum_val = T(decode_bound(maximum, bounds_are_int.y));
t_out[i] = T(op(in_val, minimum_val, maximum_val));
$else:
const float in_val = float(t_in[i]);
t_out[i] = T(op(in_val, minimum, maximum));
}

Expand All @@ -62,6 +80,11 @@ void main() {
}

VEC4_T in_texel = texelFetch(t_in, pos, 0);
$if DYNAMIC_PARAMS:
const VEC4_T minimum_val = VEC4_T(decode_bound(minimum, bounds_are_int.x));
const VEC4_T maximum_val = VEC4_T(decode_bound(maximum, bounds_are_int.y));
imageStore(t_out, pos, op(in_texel, minimum_val, maximum_val));
$else:
imageStore(t_out, pos, VEC4_T(op(in_texel, minimum, maximum)));
}

Expand Down
8 changes: 8 additions & 0 deletions backends/vulkan/runtime/graph/ops/glsl/unary_op.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ unary_op:
OPERATOR: clamp(X, A, B)
DTYPE: float
STORAGE: texture3d
DYNAMIC_PARAMS: false
generate_variant_forall:
DTYPE:
- VALUE: half
Expand All @@ -18,6 +19,13 @@ unary_op:
- NAME: clamp_int32
OPERATOR: clamp(X, A, B)
DTYPE: int32
- NAME: clamp_dynamic_int32
OPERATOR: clamp(X, A, B)
DTYPE: int32
DYNAMIC_PARAMS: true
- NAME: clamp_dynamic
OPERATOR: clamp(X, A, B)
DYNAMIC_PARAMS: true
- NAME: cos
OPERATOR: cos(X)
- NAME: exp
Expand Down
Loading
Loading