| 314 | } |
| 315 | |
| 316 | void compute_concat_node(GraphNode& node, const std::vector<std::unique_ptr<GraphNode>>& nodes, const std::unordered_map<size_t, size_t>& node_index_map) { |
| 317 | const auto& input1_buffer = get_input(node, 0, nodes, node_index_map); |
| 318 | const auto& input2_buffer = get_input(node, 1, nodes, node_index_map); |
| 319 | |
| 320 | std::vector<size_t> shape1 = input1_buffer.shape; |
| 321 | std::vector<size_t> shape2 = input2_buffer.shape; |
| 322 | std::vector<size_t> output_shape = node.output_buffer.shape; |
| 323 | |
| 324 | if (input1_buffer.precision != Precision::FP16) { |
| 325 | throw std::runtime_error("Concat operation only supports FP16 precision"); |
| 326 | } |
| 327 | cactus_concat_f16(input1_buffer.data_as<__fp16>(), input2_buffer.data_as<__fp16>(), |
| 328 | node.output_buffer.data_as<__fp16>(), |
| 329 | shape1.data(), shape2.data(), output_shape.data(), |
| 330 | shape1.size(), node.params.axis); |
| 331 | } |
| 332 | |
| 333 | void compute_cat_node( |
| 334 | GraphNode& node, |
nothing calls this directly
no test coverage detected