From ef013a0633523305b1644f173b3d86ac5e3fcf69 Mon Sep 17 00:00:00 2001 From: Yuval Tassa Date: Mon, 14 Jul 2025 12:25:42 -0700 Subject: [PATCH] Add mjPARTIAL_SORT macro to engine_sort.h PiperOrigin-RevId: 782999687 Change-Id: I986aec0e1faedae90f1891b5f3233d5262d167f3 --- src/engine/engine_sort.h | 36 +++++++++++++++++++ test/engine/engine_sort_test.cc | 62 ++++++++++++++++++++++++++++----- 2 files changed, 90 insertions(+), 8 deletions(-) diff --git a/src/engine/engine_sort.h b/src/engine/engine_sort.h index 4e2970ae..a5ac28f3 100644 --- a/src/engine/engine_sort.h +++ b/src/engine/engine_sort.h @@ -72,4 +72,40 @@ if (src != arr) memcpy(arr, src, n * sizeof(type)); \ } + +// sub-macro that sifts down a node in a max heap to its correct position +#define _mjSIFT_DOWN(type, buf, start, end, cmp, context) \ +{ \ + int root = start; \ + while (2 * root + 1 < end) { \ + int child = 2 * root + 1; \ + int swap = root; \ + if (cmp(buf + swap, buf + child, context) < 0) swap = child; \ + if (child + 1 < end && cmp(buf + swap, buf + child + 1, context) < 0) swap = child + 1; \ + if (swap == root) break; \ + type tmp = buf[root]; buf[root] = buf[swap]; buf[swap] = tmp; \ + root = swap; \ + } \ +} + +// defines an inline function that selects the bottom k elements using partial heap sort +// buf needs to be of size k +#define mjPARTIAL_SORT(name, type, cmp) \ + static inline void name(type* arr, type* buf, int n, int k, void* context) { \ + if (k <= 0 || n < k) return; \ + /* fill initial heap */ \ + for (int i = 0; i < k; i++) buf[i] = arr[i]; \ + for (int j = (k - 2) / 2; j >= 0; j--) _mjSIFT_DOWN(type, buf, j, k, cmp, context); \ + /* scan remaining elements */ \ + for (int i = k; i < n; i++) { \ + if (cmp(arr + i, buf, context) < 0) { \ + buf[0] = arr[i]; \ + _mjSIFT_DOWN(type, buf, 0, k, cmp, context); \ + } \ + } \ + /* copy back and sort the result */ \ + for (int j = 0; j < k; j++) arr[j] = buf[j]; \ + _mjINSERTION_SORT(type, arr, 0, k, cmp, context); \ + } + #endif // MUJOCO_SRC_ENGINE_ENGINE_SORT_H_ diff --git a/test/engine/engine_sort_test.cc b/test/engine/engine_sort_test.cc index a68d77c8..23410a21 100644 --- a/test/engine/engine_sort_test.cc +++ b/test/engine/engine_sort_test.cc @@ -16,8 +16,8 @@ #include #include -#include #include +#include #include #include "test/fixture.h" @@ -25,6 +25,8 @@ namespace mujoco { namespace { +using std::array; +using std::vector; using EngineSortTest = MujocoTest; constexpr int kPermutationSize = 5; // 5! = 120 checks @@ -43,7 +45,7 @@ struct IntStruct { // computes next permutation in lexicographical order in place // returns false if there is no next permutation -bool NextPermutation(std::array& arr) { +bool NextPermutation(array& arr) { int j; for (j = arr.size() - 2; j >= 0; --j) { if (arr[j] < arr[j + 1]) { @@ -76,6 +78,7 @@ int IntCompare(int* i, int* j, void* context) { } } mjSORT(IntSort, int, IntCompare) +mjPARTIAL_SORT(IntSelect, int, IntCompare) int IntStructCompare(const IntStruct* x, const IntStruct* y, void* context) { IntStruct* a = (IntStruct*)x; @@ -89,17 +92,18 @@ int IntStructCompare(const IntStruct* x, const IntStruct* y, void* context) { } } mjSORT(IntStructSort, IntStruct, IntStructCompare) +mjPARTIAL_SORT(IntStructSelect, IntStruct, IntStructCompare) TEST_F(EngineSortTest, Sort) { int total = 0; - std::array initial, buf; + array initial, buf; for (int i = 0; i < kPermutationSize; ++i) { initial[i] = i + 1; } - std::array arr = initial; + array arr = initial; do { ++total; - std::array sorted_arr = arr; + array sorted_arr = arr; IntSort(sorted_arr.data(), buf.data(), sorted_arr.size(), nullptr); EXPECT_EQ(sorted_arr, initial); } while (NextPermutation(arr)); @@ -107,7 +111,7 @@ TEST_F(EngineSortTest, Sort) { } TEST_F(EngineSortTest, LargeSort) { - std::array arr, buf; + array arr, buf; int n = 1; for (int i = arr.size() - 1; i >= 0; --i) { arr[i] = n++; @@ -119,13 +123,55 @@ TEST_F(EngineSortTest, LargeSort) { } TEST_F(EngineSortTest, SortStruct) { - std::vector y = {{1}, {3}, {2}}; - std::vector buf(3); + vector y = {{1}, {3}, {2}}; + vector buf(3); IntStructSort(y.data(), buf.data(), y.size(), nullptr); EXPECT_EQ(y[0].value, 1); EXPECT_EQ(y[1].value, 2); EXPECT_EQ(y[2].value, 3); } +TEST_F(EngineSortTest, Select) { + int total = 0; + array initial, buf; + for (int i = 0; i < kPermutationSize; ++i) { + initial[i] = i + 1; + } + array arr = initial; + do { + ++total; + for (int k = 1; k <= kPermutationSize; ++k) { + array select_arr = arr; + IntSelect(select_arr.data(), buf.data(), select_arr.size(), k, nullptr); + for (int i = 0; i < k; ++i) { + EXPECT_EQ(select_arr[i], initial[i]); + } + } + } while (NextPermutation(arr)); + ASSERT_EQ(total, factorial()); +} + +TEST_F(EngineSortTest, LargeSelect) { + array arr; + array buf; + int n = 1; + for (int i = arr.size() - 1; i >= 0; --i) { + arr[i] = n++; + } + IntSelect(arr.data(), buf.data(), arr.size(), buf.size(), nullptr); + for (int i = 0; i < buf.size(); ++i) { + EXPECT_EQ(arr[i], i + 1); + } +} + +TEST_F(EngineSortTest, SelectStruct) { + vector y = {{1}, {3}, {2}}; + vector buf(3); + const int k = 2; + IntStructSelect(y.data(), buf.data(), y.size(), k, nullptr); + EXPECT_EQ(y[0].value, 1); + EXPECT_EQ(y[1].value, 2); +} + } // namespace } // namespace mujoco