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

Function compute_gather_node

cactus/graph/graph_ops_tensor.cpp:26–125  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24}
25
26void 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>();

Callers

nothing calls this directly

Calls 4

sizeMethod · 0.80
get_dataMethod · 0.80
scales_as_fp16Method · 0.80
getMethod · 0.45

Tested by

no test coverage detected