MCPcopy Create free account
hub / github.com/cactus-compute/cactus / compute_binary_op_node

Function compute_binary_op_node

cactus/graph/graph_ops_math.cpp:98–147  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

96}
97
98void 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
149void 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);

Callers

nothing calls this directly

Calls 8

compute_stridesFunction · 0.85
cactus_add_broadcast_f16Function · 0.85
dispatch_binary_op_f16Function · 0.85
dataMethod · 0.80
sizeMethod · 0.80

Tested by

no test coverage detected