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

Function compute_topk_node

cactus/graph/graph_ops_sample.cpp:43–89  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

41}
42
43void compute_topk_node(GraphNode& node, const std::vector<std::unique_ptr<GraphNode>>& nodes, const std::unordered_map<size_t, size_t>& node_index_map) {
44 const auto& input_buffer = get_input(node, 0, nodes, node_index_map);
45 if (input_buffer.shape.size() != 2) {
46 throw std::runtime_error("TopK currently only supports 2D tensors [batch, features]");
47 }
48
49 size_t k = node.params.top_k;
50 size_t batch_size = input_buffer.shape[0];
51 size_t feature_size = input_buffer.shape[1];
52 size_t block_size = batch_size * k;
53
54 std::vector<float> input_float(input_buffer.total_size);
55 if (input_buffer.precision == Precision::INT8) {
56 throw std::runtime_error("TopK currently does not support INT8 input");
57 } else if (input_buffer.precision == Precision::FP16) {
58 const __fp16* input_fp16 = input_buffer.data_as<__fp16>();
59 for (size_t i = 0; i < input_buffer.total_size; ++i) {
60 input_float[i] = static_cast<float>(input_fp16[i]);
61 }
62 } else {
63 const float* input_fp32 = input_buffer.data_as<float>();
64 std::memcpy(input_float.data(), input_fp32, input_buffer.total_size * sizeof(float));
65 }
66
67 float* output = node.output_buffer.data_as<float>();
68
69 for (size_t b = 0; b < batch_size; ++b) {
70 const float* row = input_float.data() + b * feature_size;
71
72 std::vector<std::pair<size_t, float>> indexed_values(feature_size);
73 for (size_t i = 0; i < feature_size; ++i) {
74 indexed_values[i] = {i, row[i]};
75 }
76
77 std::partial_sort(indexed_values.begin(),
78 indexed_values.begin() + k,
79 indexed_values.end(),
80 [](const auto& a, const auto& b) { return a.second > b.second; });
81
82 float* idx_out_row = output + b * k;
83 float* val_out_row = output + block_size + b * k;
84 for (size_t i = 0; i < k; ++i) {
85 idx_out_row[i] = static_cast<float>(indexed_values[i].first);
86 val_out_row[i] = indexed_values[i].second;
87 }
88 }
89}
90
91void compute_scatter_topk_node(GraphNode& node, const std::vector<std::unique_ptr<GraphNode>>& nodes, const std::unordered_map<size_t, size_t>& node_index_map) {
92 const auto& indices_buffer = get_input(node, 0, nodes, node_index_map);

Callers

nothing calls this directly

Calls 2

sizeMethod · 0.80
dataMethod · 0.80

Tested by

no test coverage detected