| 24 | } |
| 25 | |
| 26 | void compute_gather_node(GraphNode& node, const std::vector<std::unique_ptr<GraphNode>>& nodes, const std::unordered_map<size_t, size_t>& node_index_map) { |
| 27 | const auto& tensor_buffer = get_input(node, 0, nodes, node_index_map); |
| 28 | const auto& indices_buffer = get_input(node, 1, nodes, node_index_map); |
| 29 | |
| 30 | size_t first_dim = tensor_buffer.shape[0]; |
| 31 | size_t element_size = 1; |
| 32 | for (size_t i = 1; i < tensor_buffer.shape.size(); i++) { |
| 33 | element_size *= tensor_buffer.shape[i]; |
| 34 | } |
| 35 | |
| 36 | size_t num_indices = indices_buffer.total_size; |
| 37 | size_t bytes_per_element = PrecisionTraits::packed_size_of(tensor_buffer.precision, element_size); |
| 38 | |
| 39 | if (PrecisionTraits::is_integer(tensor_buffer.precision)) { |
| 40 | const char* tensor_data = static_cast<const char*>(tensor_buffer.get_data()); |
| 41 | char* output = static_cast<char*>(node.output_buffer.get_data()); |
| 42 | Precision prec = tensor_buffer.precision; |
| 43 | |
| 44 | const bool is_grouped = tensor_buffer.group_size > 0; |
| 45 | __fp16* gathered_scales = nullptr; |
| 46 | const __fp16* src_scales = nullptr; |
| 47 | size_t num_groups = 0; |
| 48 | |
| 49 | if (is_grouped) { |
| 50 | num_groups = tensor_buffer.num_groups; |
| 51 | src_scales = tensor_buffer.scales_as_fp16(); |
| 52 | size_t scales_bytes = num_indices * num_groups * sizeof(__fp16); |
| 53 | node.output_buffer.owned_scales = std::make_unique<char[]>(scales_bytes); |
| 54 | gathered_scales = reinterpret_cast<__fp16*>(node.output_buffer.owned_scales.get()); |
| 55 | } |
| 56 | |
| 57 | const int8_t* indices = indices_buffer.data_as<int8_t>(); |
| 58 | for (size_t i = 0; i < num_indices; i++) { |
| 59 | size_t idx = static_cast<size_t>(indices[i]); |
| 60 | if (idx >= first_dim) { |
| 61 | throw std::runtime_error("Gather index " + std::to_string(idx) + " out of bounds for dimension " + std::to_string(first_dim)); |
| 62 | } |
| 63 | std::memcpy(output + PrecisionTraits::byte_offset_of(prec, i * element_size), |
| 64 | tensor_data + PrecisionTraits::byte_offset_of(prec, idx * element_size), |
| 65 | bytes_per_element); |
| 66 | if (is_grouped) { |
| 67 | for (size_t g = 0; g < num_groups; g++) { |
| 68 | gathered_scales[i * num_groups + g] = src_scales[idx * num_groups + g]; |
| 69 | } |
| 70 | } |
| 71 | } |
| 72 | |
| 73 | if (is_grouped) { |
| 74 | node.output_buffer.group_size = tensor_buffer.group_size; |
| 75 | node.output_buffer.num_groups = num_groups; |
| 76 | node.output_buffer.scales_data = gathered_scales; |
| 77 | } |
| 78 | } else if (tensor_buffer.precision == Precision::FP16) { |
| 79 | const __fp16* tensor_data = tensor_buffer.data_as<__fp16>(); |
| 80 | __fp16* output = node.output_buffer.data_as<__fp16>(); |
| 81 | |
| 82 | if (indices_buffer.precision == Precision::INT8) { |
| 83 | const int8_t* indices = indices_buffer.data_as<int8_t>(); |
nothing calls this directly
no test coverage detected