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
+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