| 163 | // Mostly modeled on tensorflow/core/kernels/topk_op.cc for CPU. |
| 164 | template <typename T> |
| 165 | void TopK(int32 row_size, int32 num_rows, const T* data, int32 k, |
| 166 | int32* output_indexes, T* output_values) { |
| 167 | TopContainer<T> topc(k, row_size); |
| 168 | for (int row = 0; row < num_rows; ++row) { |
| 169 | const T* values_row = data + row * row_size; |
| 170 | topc.start_collecting(values_row); |
| 171 | for (int32 c = 0; c < row_size; ++c) { |
| 172 | topc.push(c); |
| 173 | } |
| 174 | |
| 175 | // Prepare output buffers. |
| 176 | int32* indexes_row = output_indexes + row * k; |
| 177 | T* output_row = output_values + row * k; |
| 178 | // We always assume that the output is sorted. |
| 179 | const auto& top_k = topc.sorted_result(); |
| 180 | std::copy(top_k.begin(), top_k.end(), indexes_row); |
| 181 | std::transform(top_k.begin(), top_k.end(), output_row, |
| 182 | [values_row](const int32 loc) { return values_row[loc]; }); |
| 183 | } |
| 184 | } |
| 185 | |
| 186 | } // namespace |
| 187 | |