| 34 | namespace metal { |
| 35 | namespace { |
| 36 | std::string GetReshapeCode() { |
| 37 | std::string code = R"( |
| 38 | #include <metal_stdlib> |
| 39 | using namespace metal; |
| 40 | |
| 41 | struct uniforms { |
| 42 | int4 src_size; |
| 43 | int4 dst_size; |
| 44 | }; |
| 45 | |
| 46 | $0 |
| 47 | kernel void ComputeFunction( |
| 48 | $1 |
| 49 | uint3 gid[[thread_position_in_grid]]) { |
| 50 | const int3 igid = int3(gid); |
| 51 | |
| 52 | if (igid.x >= params.dst_size.x || igid.y >= params.dst_size.y || |
| 53 | igid.z * 4 >= params.dst_size.z) return; |
| 54 | |
| 55 | FLT4 value; |
| 56 | |
| 57 | for (int i = 0; i < 4; ++i) { |
| 58 | const int dst_channel = igid.z * 4 + i; |
| 59 | if (dst_channel < params.dst_size.z) { |
| 60 | int p = dst_channel + params.dst_size.z * igid.x + params.dst_size.w * igid.y; |
| 61 | int src_y = p / params.src_size.w; |
| 62 | int t0 = p - src_y * params.src_size.w; // p % params.src_size.w; |
| 63 | int src_x = t0 / params.src_size.z; |
| 64 | int src_z = t0 - src_x * params.src_size.z; // t0 % params.src_size.z; |
| 65 | int src_layer = src_z >> 2; |
| 66 | int src_channel = src_z & 3; |
| 67 | int src_linear_id = (src_layer * params.src_size.y + src_y) * params.src_size.x + src_x; |
| 68 | value[i] = src_buffer[src_linear_id][src_channel]; |
| 69 | } |
| 70 | } |
| 71 | |
| 72 | int linear_index = (igid.z * params.dst_size.y + igid.y) * params.dst_size.x + igid.x; |
| 73 | $2 |
| 74 | dst_buffer[linear_index] = value; |
| 75 | })"; |
| 76 | return code; |
| 77 | } |
| 78 | |
| 79 | std::string GetReshapex4Code() { |
| 80 | std::string code = R"( |