Move Newton Hessian memory allocation from stack to arena.
PiperOrigin-RevId: 684054446 Change-Id: I4be0b216810ebbd9d7fa119c6a84eb08d3c16b83
This commit is contained in:
committed by
Copybara-Service
parent
356d172f56
commit
9a0dc20821
@@ -1360,12 +1360,20 @@ Euler integrator, semi-implicit in velocity.
|
||||
# (e.g. because the internal output buffer is too small)
|
||||
self.assertIn('mujoco', mujoco.mj_printSchema(flg_html, flg_pad))
|
||||
|
||||
def test_pickle_mjdata_before_step(self):
|
||||
data2 = pickle.loads(pickle.dumps(self.data))
|
||||
attr_to_compare = (
|
||||
'time', 'qpos', 'qvel', 'qacc', 'xpos', 'mocap_pos',
|
||||
'warning', 'energy', 'contact', 'efc_J', 'L'
|
||||
)
|
||||
self._assert_attributes_equal(data2, self.data, attr_to_compare)
|
||||
|
||||
def test_pickle_mjdata(self):
|
||||
mujoco.mj_step(self.model, self.data)
|
||||
data2 = pickle.loads(pickle.dumps(self.data))
|
||||
attr_to_compare = (
|
||||
'time', 'qpos', 'qvel', 'qacc', 'xpos', 'mocap_pos',
|
||||
'warning', 'energy', 'contact', 'efc_J'
|
||||
'warning', 'energy', 'contact', 'efc_J', 'L'
|
||||
)
|
||||
self._assert_attributes_equal(data2, self.data, attr_to_compare)
|
||||
for _ in range(10):
|
||||
@@ -1379,7 +1387,7 @@ Euler integrator, semi-implicit in velocity.
|
||||
data2 = pickle.loads(pickle.dumps(self.data))
|
||||
attr_to_compare = (
|
||||
'time', 'qpos', 'qvel', 'qacc', 'xpos', 'mocap_pos',
|
||||
'warning', 'energy', 'contact', 'efc_J'
|
||||
'warning', 'energy', 'contact', 'efc_J', 'L'
|
||||
)
|
||||
self._assert_attributes_equal(data2, self.data, attr_to_compare)
|
||||
for _ in range(10):
|
||||
|
||||
@@ -15,6 +15,7 @@
|
||||
#ifndef MUJOCO_PYTHON_SERIALIZATION_H_
|
||||
#define MUJOCO_PYTHON_SERIALIZATION_H_
|
||||
|
||||
#include <cstddef>
|
||||
#include <iostream>
|
||||
|
||||
#include <mujoco/mjtnum.h>
|
||||
@@ -56,7 +57,9 @@ inline void WriteBytes(std::ostream& output, const void* src,
|
||||
// Start by writing nbytes itself, so it can be validated at the time of
|
||||
// reading.
|
||||
WriteInt(output, nbytes);
|
||||
output.write(reinterpret_cast<const char*>(src), nbytes);
|
||||
if (src) {
|
||||
output.write(reinterpret_cast<const char*>(src), nbytes);
|
||||
}
|
||||
}
|
||||
|
||||
inline void ReadBytes(std::istream& input, void* dest, std::size_t nbytes) {
|
||||
|
||||
+32
-12
@@ -755,7 +755,7 @@ void MjDataWrapper::Serialize(std::ostream& output) const {
|
||||
X(energy);
|
||||
#undef X
|
||||
|
||||
// Write buffer contents
|
||||
// Write buffer and arena contents
|
||||
{
|
||||
MJDATA_POINTERS_PREAMBLE((this->model_->get()))
|
||||
|
||||
@@ -764,18 +764,24 @@ void MjDataWrapper::Serialize(std::ostream& output) const {
|
||||
MJDATA_POINTERS
|
||||
#undef X
|
||||
|
||||
bool is_sparse_newton = this->model_->get()->opt.solver == mjSOL_NEWTON &&
|
||||
mj_isSparse(this->model_->get());
|
||||
|
||||
#undef MJ_M
|
||||
#define MJ_M(x) this->model_->get()->x
|
||||
#undef MJ_D
|
||||
#define MJ_D(x) this->ptr_->x
|
||||
#define X(type, name, nr, nc) \
|
||||
if ((nr) * (nc)) { \
|
||||
WriteBytes(output, ptr_->name, sizeof(type) * (nr) * (nc)); \
|
||||
#define X(type, name, nr, nc) \
|
||||
if ((nr) * (nc)) { \
|
||||
WriteBytes(output, ptr_->name, \
|
||||
ptr_->name ? sizeof(type) * (nr) * (nc) : 0); \
|
||||
}
|
||||
|
||||
MJDATA_ARENA_POINTERS_CONTACT
|
||||
MJDATA_ARENA_POINTERS_SOLVER
|
||||
// MJDATA_ARENA_POINTERS_NEWTON // TODO: tassa - Add after allocation exists
|
||||
if (is_sparse_newton) {
|
||||
MJDATA_ARENA_POINTERS_NEWTON
|
||||
}
|
||||
if (mj_isDual(this->model_->get())) {
|
||||
MJDATA_ARENA_POINTERS_DUAL
|
||||
}
|
||||
@@ -805,6 +811,8 @@ MjDataWrapper MjDataWrapper::Deserialize(std::istream& input) {
|
||||
|
||||
bool is_dual = mj_isDual(&m);
|
||||
|
||||
bool is_sparse_newton = m.opt.solver == mjSOL_NEWTON && mj_isSparse(&m);
|
||||
|
||||
raw::MjData* d = mj_makeData(&m);
|
||||
if (!d) {
|
||||
throw py::value_error("Failed to create mjData.");
|
||||
@@ -833,7 +841,7 @@ MjDataWrapper MjDataWrapper::Deserialize(std::istream& input) {
|
||||
X(energy);
|
||||
#undef X
|
||||
|
||||
// Read buffer contents
|
||||
// Read buffer and arena contents
|
||||
{
|
||||
MJDATA_POINTERS_PREAMBLE((&m))
|
||||
|
||||
@@ -846,16 +854,28 @@ MjDataWrapper MjDataWrapper::Deserialize(std::istream& input) {
|
||||
#define MJ_M(x) m.x
|
||||
#undef MJ_D
|
||||
#define MJ_D(x) d->x
|
||||
#define X(type, name, nr, nc) \
|
||||
if ((nr) * (nc)) { \
|
||||
d->name = static_cast<decltype(d->name)>( \
|
||||
mj_arenaAllocByte(d, sizeof(type) * (nr) * (nc), alignof(type))); \
|
||||
ReadBytes(input, d->name, sizeof(type) * (nr) * (nc)); \
|
||||
// arena pointers might be null, so we need to check the size before allocating.
|
||||
#define X(type, name, nr, nc) \
|
||||
if ((nr) * (nc)) { \
|
||||
std::size_t actual_nbytes = ReadInt(input); \
|
||||
if (actual_nbytes) { \
|
||||
if (actual_nbytes != sizeof(type) * (nr) * (nc)) { \
|
||||
input.setstate(input.rdstate() | std::ios_base::failbit); \
|
||||
} else { \
|
||||
d->name = static_cast<decltype(d->name)>( \
|
||||
mj_arenaAllocByte(d, sizeof(type) * (nr) * (nc), alignof(type))); \
|
||||
input.read(reinterpret_cast<char*>(d->name), actual_nbytes); \
|
||||
} \
|
||||
} else { \
|
||||
d->name = nullptr; \
|
||||
} \
|
||||
}
|
||||
|
||||
MJDATA_ARENA_POINTERS_CONTACT
|
||||
MJDATA_ARENA_POINTERS_SOLVER
|
||||
// MJDATA_ARENA_POINTERS_NEWTON // TODO: tassa - Add after allocation exists
|
||||
if (is_sparse_newton) {
|
||||
MJDATA_ARENA_POINTERS_NEWTON
|
||||
}
|
||||
if (is_dual) {
|
||||
MJDATA_ARENA_POINTERS_DUAL
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user