Refactor mjXUtil::WriteAttr using templates.

PiperOrigin-RevId: 453848941
Change-Id: I1f87f0c20f0c89caab89ea4891aee7583a92534e
This commit is contained in:
Alessio Quaglino
2022-06-08 23:33:04 -07:00
committed by Copybara-Service
parent 066367fcc8
commit c7c765b6eb
4 changed files with 58 additions and 183 deletions
+45 -159
View File
@@ -19,8 +19,10 @@
#include <cstdlib>
#include <cstring>
#include <iostream>
#include <limits>
#include <sstream>
#include <string>
#include <type_traits>
#include <vector>
#include "cc/array_safety.h"
@@ -483,33 +485,16 @@ XMLElement* mjXSchema::Check(XMLElement* elem, int level) {
//---------------------------------- class mjXUtil implementation ----------------------------------
// compare two vectors: double
bool mjXUtil::SameVector(const double* vec1, const double* vec2, int n) {
// compare two vectors
template<typename T>
bool mjXUtil::SameVector(const T* vec1, const T* vec2, int n) {
if (!vec1 || !vec2) {
return false;
}
bool same = true;
for (int i=0; i<n; i++) {
if (fabs(vec1[i] - vec2[i]) > 1E-10) {
same = false;
}
}
return same;
}
// compare two vectors: double
bool mjXUtil::SameVector(const float* vec1, const float* vec2, int n) {
if (!vec1 || !vec2) {
return false;
}
bool same = true;
for (int i=0; i<n; i++) {
if (fabs(vec1[i] - vec2[i]) > 1E-7) {
if (fabs(vec1[i] - vec2[i]) > std::numeric_limits<T>::epsilon()) {
same = false;
}
}
@@ -517,6 +502,10 @@ bool mjXUtil::SameVector(const float* vec1, const float* vec2, int n) {
return same;
}
template bool mjXUtil::SameVector(const double* vec1, const double* vec2, int n);
template bool mjXUtil::SameVector(const float* vec1, const float* vec2, int n);
template bool mjXUtil::SameVector(const int* vec1, const int* vec2, int n);
template bool mjXUtil::SameVector(const mjtByte* vec1, const mjtByte* vec2, int n);
// find string in map, return corresponding integer (-1: not found)
@@ -808,163 +797,60 @@ static int Round(double x) {
}
// write attribute- double
void mjXUtil::WriteAttr(XMLElement* elem, string name, int n, double* data,
const double* def) {
char buf[100];
string value;
value.clear();
// write attribute
template<typename T>
void mjXUtil::WriteAttr(XMLElement* elem, string name, int n, T* data, const T* def) {
// make sure all are defined
for (int i=0; i<n; i++) {
if (std::isnan(data[i])) {
return;
}
}
// skip default attributes
if (SameVector(data, def, n)) {
return;
}
// process all numbers
for (int i=0; i<n; i++) {
// add space between numbers
if (i>0) {
value = value + " ";
}
// write integer or float
if (isint(data[i])) {
mju::sprintf_arr(buf, "%d", Round(data[i]));
} else {
mju::sprintf_arr(buf, mujoco::_mjPRIVATE__get_xml_precision(), data[i]);
}
// append number
value = value + buf;
}
// set attribute as string
WriteAttrTxt(elem, name, value);
}
// write attribute- float
void mjXUtil::WriteAttr(XMLElement* elem, string name, int n, float* data,
const float* def) {
char buf[100];
string value;
value.clear();
// skip default attributes
if (SameVector(data, def, n)) {
return;
}
// process all numbers
for (int i=0; i<n; i++) {
// add space between numbers
if (i>0) {
value = value + " ";
}
// write integer or float
if (isint(data[i])) {
mju::sprintf_arr(buf, "%d", Round(data[i]));
} else {
mju::sprintf_arr(buf, "%g", data[i]);
}
// append number
value = value + buf;
}
// set attribute as string
WriteAttrTxt(elem, name, value);
}
// write attribute- byte
void mjXUtil::WriteAttr(XMLElement* elem, string name, int n, mjtByte* data,
const mjtByte* def) {
char buf[100];
string value;
value.clear();
// skip default attributes
if (def) {
bool skip = true;
for (int i=0; i<n; i++)
if (data[i] != def[i]) {
skip = false;
}
if (skip) {
return;
}
}
// process all numbers
for (int i=0; i<n; i++) {
// add space between numbers
if (i>0) {
value = value + " ";
}
// write integer
mju::sprintf_arr(buf, "%d", data[i]);
// append number
value = value + buf;
}
// set attribute as string
WriteAttrTxt(elem, name, value);
}
// write attribute- int
void mjXUtil::WriteAttr(XMLElement* elem, string name, int n, int* data,
const int* def) {
char buf[100];
string value;
value.clear();
// skip default attributes
if (def) {
bool skip = true;
if constexpr (std::is_floating_point_v<T>) {
for (int i=0; i<n; i++) {
if (data[i] != def[i]) {
skip = false;
if (std::isnan(data[i])) {
return;
}
}
if (skip) {
return;
}
}
// skip default attributes
if (SameVector(data, def, n)) {
return;
}
// increase precision for testing
stringstream stream;
stream.precision(mujoco::_mjPRIVATE__get_xml_precision());
// process all numbers
for (int i=0; i<n; i++) {
// add space between numbers
if (i>0) {
value = value + " ";
stream << " ";
}
// write integer
mju::sprintf_arr(buf, "%d", data[i]);
// append number
value = value + buf;
if (isint(data[i])) {
stream << Round(data[i]);
} else {
stream << data[i];
}
}
// set attribute as string
WriteAttrTxt(elem, name, value);
WriteAttrTxt(elem, name, stream.str());
}
template void mjXUtil::WriteAttr(XMLElement* elem, string name, int n,
double* data, const double* def);
template void mjXUtil::WriteAttr(XMLElement* elem, string name, int n,
float* data, const float* def);
template void mjXUtil::WriteAttr(XMLElement* elem, string name, int n,
int* data, const int* def);
template void mjXUtil::WriteAttr(XMLElement* elem, string name, int n,
mjtByte* data, const mjtByte* def);
// write vector<double> attribute, default = zero array
void mjXUtil::WriteVector(XMLElement* elem, string name, vector<double>& vec) {
// proceed only if non-zero found