import numpy as np
__input__ = {"kernel": {"scatter_elements": "scatter_elements_input"}}
def scatter_elements_input(
data, indices, updates, *, axis=0, reduction="none", **kwargs
):
"""
Input function for scatter_elements.
All the parameters (names and order) follow @scatter_elements_def.cpp without outputs.
All the input Tensors are numpy.ndarray.
Args:
**kwargs: {input,output}_{dtypes,ori_shapes,formats,ori_formats},
input_ranges, full_soc_version, short_soc_version, testcase_name
Returns:
List of input tensors (length must match Input count in _def.cpp)
"""
shape, dtype = indices.shape, indices.dtype
if axis < 0:
axis = axis + len(shape)
value_max = data.shape[axis]
if reduction in ["none", "mul", "add"]:
batch_shape = shape[:axis] + shape[axis + 1 :]
tmp_shape = batch_shape + (shape[axis],)
batch_num = 1
for num in batch_shape:
batch_num *= num
shape_order = list(range(len(shape) - 1))
shape_order.insert(axis, len(shape) - 1)
x = np.random.choice(range(0, value_max), size=shape[axis], replace=False)
x = np.expand_dims(x, axis=0)
indices = np.repeat(x, batch_num, axis=0)
np.apply_along_axis(np.random.shuffle, axis=1, arr=indices)
indices = indices.reshape(tmp_shape).transpose(tuple(shape_order)).astype(dtype)
else:
indices = np.random.uniform(0, value_max, shape).astype(dtype)
return [data, indices, updates]