| 96 | } |
| 97 | |
| 98 | void compute_binary_op_node(GraphNode& node, const std::vector<std::unique_ptr<GraphNode>>& nodes, const std::unordered_map<size_t, size_t>& node_index_map) { |
| 99 | const auto& lhs = get_input(node, 0, nodes, node_index_map); |
| 100 | const auto& rhs = get_input(node, 1, nodes, node_index_map); |
| 101 | |
| 102 | if (lhs.precision != Precision::FP16) { |
| 103 | throw std::runtime_error("Binary operations only support FP16 precision (got " + std::to_string(static_cast<int>(lhs.precision)) + ")"); |
| 104 | } |
| 105 | |
| 106 | if (node.params.broadcast_info.needs_broadcasting) { |
| 107 | std::vector<size_t> lhs_strides = compute_strides(lhs.shape, node.params.broadcast_info.output_shape); |
| 108 | std::vector<size_t> rhs_strides = compute_strides(rhs.shape, node.params.broadcast_info.output_shape); |
| 109 | |
| 110 | switch (node.op_type) { |
| 111 | case OpType::ADD: |
| 112 | case OpType::ADD_CLIPPED: |
| 113 | cactus_add_broadcast_f16(lhs.data_as<__fp16>(), rhs.data_as<__fp16>(), |
| 114 | node.output_buffer.data_as<__fp16>(), |
| 115 | lhs_strides.data(), rhs_strides.data(), |
| 116 | node.params.broadcast_info.output_shape.data(), |
| 117 | node.params.broadcast_info.output_shape.size()); |
| 118 | break; |
| 119 | case OpType::SUBTRACT: |
| 120 | cactus_subtract_broadcast_f16(lhs.data_as<__fp16>(), rhs.data_as<__fp16>(), |
| 121 | node.output_buffer.data_as<__fp16>(), |
| 122 | lhs_strides.data(), rhs_strides.data(), |
| 123 | node.params.broadcast_info.output_shape.data(), |
| 124 | node.params.broadcast_info.output_shape.size()); |
| 125 | break; |
| 126 | case OpType::MULTIPLY: |
| 127 | cactus_multiply_broadcast_f16(lhs.data_as<__fp16>(), rhs.data_as<__fp16>(), |
| 128 | node.output_buffer.data_as<__fp16>(), |
| 129 | lhs_strides.data(), rhs_strides.data(), |
| 130 | node.params.broadcast_info.output_shape.data(), |
| 131 | node.params.broadcast_info.output_shape.size()); |
| 132 | break; |
| 133 | case OpType::DIVIDE: |
| 134 | cactus_divide_broadcast_f16(lhs.data_as<__fp16>(), rhs.data_as<__fp16>(), |
| 135 | node.output_buffer.data_as<__fp16>(), |
| 136 | lhs_strides.data(), rhs_strides.data(), |
| 137 | node.params.broadcast_info.output_shape.data(), |
| 138 | node.params.broadcast_info.output_shape.size()); |
| 139 | break; |
| 140 | default: break; |
| 141 | } |
| 142 | } else { |
| 143 | dispatch_binary_op_f16(node.op_type, lhs.data_as<__fp16>(), |
| 144 | rhs.data_as<__fp16>(), node.output_buffer.data_as<__fp16>(), |
| 145 | node.output_buffer.total_size); |
| 146 | } |
| 147 | } |
| 148 | |
| 149 | void compute_unary_op_node(GraphNode& node, const std::vector<std::unique_ptr<GraphNode>>& nodes, const std::unordered_map<size_t, size_t>& node_index_map) { |
| 150 | const auto& input = get_input(node, 0, nodes, node_index_map); |
nothing calls this directly
no test coverage detected