Add mjPARTIAL_SORT macro to engine_sort.h

PiperOrigin-RevId: 782999687
Change-Id: I986aec0e1faedae90f1891b5f3233d5262d167f3
This commit is contained in:
Yuval Tassa
2025-07-14 12:25:42 -07:00
committed by Copybara-Service
parent 2a700c98aa
commit ef013a0633
2 changed files with 90 additions and 8 deletions
+36
View File
@@ -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_
+54 -8
View File
@@ -16,8 +16,8 @@
#include <algorithm>
#include <array>
#include <utility>
#include <vector>
#include <utility>
#include <gtest/gtest.h>
#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<int, kPermutationSize>& arr) {
bool NextPermutation(array<int, kPermutationSize>& 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<int, kPermutationSize> initial, buf;
array<int, kPermutationSize> initial, buf;
for (int i = 0; i < kPermutationSize; ++i) {
initial[i] = i + 1;
}
std::array<int, kPermutationSize> arr = initial;
array<int, kPermutationSize> arr = initial;
do {
++total;
std::array<int, kPermutationSize> sorted_arr = arr;
array<int, kPermutationSize> 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<int, 2500> arr, buf;
array<int, 2500> 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<IntStruct> y = {{1}, {3}, {2}};
std::vector<IntStruct> buf(3);
vector<IntStruct> y = {{1}, {3}, {2}};
vector<IntStruct> 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<int, kPermutationSize> initial, buf;
for (int i = 0; i < kPermutationSize; ++i) {
initial[i] = i + 1;
}
array<int, kPermutationSize> arr = initial;
do {
++total;
for (int k = 1; k <= kPermutationSize; ++k) {
array<int, kPermutationSize> 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<int, 2500> arr;
array<int, 100> 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<IntStruct> y = {{1}, {3}, {2}};
vector<IntStruct> 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