MCPcopy Create free account
hub / github.com/Derious/cuMPC / TopKselection

Function TopKselection

test/test_CORE/test_topk.cpp:160–233  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

158
159
160void TopKselection(std::vector<int>& result, std::vector<int>& arr, int K, int &iter){
161 int n = arr.size();
162 result.clear();
163 while (K > 0) {
164 printf("K: %d\n", K);
165 printf("input size: %d\n", arr.size());
166 // for(int i = 0; i < arr.size(); i++){
167 // printf("%d ", arr[i]);
168 // }
169 // printf("\n");
170 // 确定 pivots 数量
171 int p = 2;
172
173 // 选择 pivots
174 std::vector<int> pivots = selectPivots(arr, p - 1);
175 arr.pop_back();
176 // printf("pivots: ");
177 // for(int i = 0; i < pivots.size(); i++){
178 // printf("%d ", pivots[i]);
179 // }
180 // printf("\n");
181
182 // 根据 pivots 对数组进行分区
183 std::vector<std::vector<int>> blocks = partition(arr, pivots);
184 printf("blocks: ");
185 for(int i = 0; i < blocks.size(); i++){
186 printf("[");
187 printf("%d", blocks[i].size());
188 printf("] ");
189 }
190 printf("\n");
191
192 // 统计区块大小,找到包含 TopK 的区块
193 int total = 0;
194 arr.clear();
195 for (int i = blocks.size() - 1; i >= 0; i--) {
196 if (total + blocks[i].size() >= K) {
197 arr.insert(arr.end(), blocks[i].begin(), blocks[i].end());
198 break;
199 }
200 result.insert(result.end(), blocks[i].begin(), blocks[i].end());
201 result.insert(result.end(), pivots.begin(), pivots.end());
202 total += blocks[i].size() + pivots.size();
203 }
204 // printf("result: ");
205 // for(int i = 0; i < result.size(); i++){
206 // printf("%d ", result[i]);
207 // }
208 // printf("\n");
209 // 更新 K 和 arr
210 K -= total;
211 n = arr.size();
212
213 //remove pivots from arr
214
215
216 // 如果当前区块大小小于等于 K,直接返回
217 if (n <= K) {

Callers 1

mainFunction · 0.85

Calls 8

selectPivotsFunction · 0.85
partitionFunction · 0.70
sizeMethod · 0.45
clearMethod · 0.45
insertMethod · 0.45
endMethod · 0.45
beginMethod · 0.45
resizeMethod · 0.45

Tested by

no test coverage detected