| 31 | } |
| 32 | |
| 33 | static TinyNNStatus execute(Instruction* inst, VM* vm) { |
| 34 | Tensor* output = inst->workload.dimshuffle.output; |
| 35 | Dimshuffle* dimshuffle = &inst->workload.dimshuffle; |
| 36 | int32_t nr_dim = dimshuffle->pattern_dim; |
| 37 | Tensor input = *dimshuffle->input; |
| 38 | Layout origin_layout = input.layout; |
| 39 | TINYNN_ASSERT(nr_dim == origin_layout.nr_dim); |
| 40 | for (int32_t i = 0; i < nr_dim; i++) { |
| 41 | int32_t axis = dimshuffle->pattern[i]; |
| 42 | input.layout.dims[i] = origin_layout.dims[axis]; |
| 43 | input.layout.stride[i] = origin_layout.stride[axis]; |
| 44 | } |
| 45 | output->dtype = input.dtype; |
| 46 | output->layout = input.layout; |
| 47 | //! init output stride |
| 48 | output->layout.stride[output->layout.nr_dim - 1] = 1; |
| 49 | for (int index = output->layout.nr_dim - 2; index >= 0; index--) { |
| 50 | output->layout.stride[index] = |
| 51 | output->layout.dims[index + 1] * output->layout.stride[index + 1]; |
| 52 | } |
| 53 | alloc_tensor(output, vm); |
| 54 | //! do dimshuffle naive |
| 55 | size_t nr_elem = 1; |
| 56 | for (int i = 0; i < output->layout.nr_dim; ++i) { |
| 57 | nr_elem *= output->layout.dims[i]; |
| 58 | } |
| 59 | NoconIter src_iter = init_iter(input.layout); |
| 60 | NoconIter dst_iter = init_iter(output->layout); |
| 61 | if (dtype_length((input).dtype.type_enum, NULL) == 1) { |
| 62 | char* dst_data = output->ptr; |
| 63 | char* src_data = input.ptr; |
| 64 | for (size_t i = 0; i < nr_elem; ++i) { |
| 65 | dst_data[dst_iter.offset] = src_data[src_iter.offset]; |
| 66 | inc_iter(input.layout, &src_iter); |
| 67 | inc_iter(output->layout, &dst_iter); |
| 68 | } |
| 69 | } else if (dtype_length((input).dtype.type_enum, NULL) == 2) { |
| 70 | int16_t* dst_data = output->ptr; |
| 71 | int16_t* src_data = input.ptr; |
| 72 | for (size_t i = 0; i < nr_elem; ++i) { |
| 73 | dst_data[dst_iter.offset] = src_data[src_iter.offset]; |
| 74 | inc_iter(input.layout, &src_iter); |
| 75 | inc_iter(output->layout, &dst_iter); |
| 76 | } |
| 77 | } else if (dtype_length(input.dtype.type_enum, NULL) == 4) { |
| 78 | int32_t* dst_data = output->ptr; |
| 79 | int32_t* src_data = input.ptr; |
| 80 | for (size_t i = 0; i < nr_elem; ++i) { |
| 81 | dst_data[dst_iter.offset] = src_data[src_iter.offset]; |
| 82 | inc_iter(input.layout, &src_iter); |
| 83 | inc_iter(output->layout, &dst_iter); |
| 84 | } |
| 85 | } else { |
| 86 | LOG_ERROR("unsupport dtype in dimshuffle.\n"); |
| 87 | return TinyNN_ERROR_UNSUPPORTED_DTYPE_TYPE; |
| 88 | } |
| 89 | #if TINYNN_DUMP_TENSOR |
| 90 | log_tensor(dimshuffle->output, "dimshuffle", dimshuffle->input); |
nothing calls this directly
no test coverage detected