Refactor texture.data to be handled using py::bytes
PiperOrigin-RevId: 843226742 Change-Id: I7ab920da3c387467b1fced27f557c6edd1af4bce
This commit is contained in:
committed by
Copybara-Service
parent
87d333c64a
commit
3c19b63ed3
@@ -226,17 +226,17 @@ def _ptr_binding_code(
|
||||
return f"""\
|
||||
{classname}.def_property(
|
||||
"{varname}",
|
||||
[]({rawclassname}& self) -> MjTypeVec<std::byte> {{
|
||||
return MjTypeVec<std::byte>(self.{fullvarname}->data(),
|
||||
self.{fullvarname}->size());
|
||||
}},
|
||||
[]({rawclassname}& self, py::bytes& rhs) {{
|
||||
self.{fullvarname}->clear();
|
||||
self.{fullvarname}->reserve(py::len(rhs));
|
||||
std::string_view rhs_view = py::cast<std::string_view>(rhs);
|
||||
for (auto val : rhs_view) {{
|
||||
self.{fullvarname}->push_back(static_cast<std::byte>(val));
|
||||
}}
|
||||
[]({rawclassname}& self) -> py::bytes {{
|
||||
return py::bytes(reinterpret_cast<const char*>(self.{fullvarname}->data()),
|
||||
self.{fullvarname}->size());
|
||||
}},
|
||||
[]({rawclassname}& self, py::bytes rhs) {{
|
||||
self.{fullvarname}->clear();
|
||||
std::string_view rhs_view = py::cast<std::string_view>(rhs);
|
||||
self.{fullvarname}->reserve(rhs_view.length());
|
||||
for (char val : rhs_view) {{
|
||||
self.{fullvarname}->push_back(static_cast<std::byte>(val));
|
||||
}}
|
||||
}}, py::return_value_policy::move);"""
|
||||
elif vartype == 'mjStringVec':
|
||||
return f"""\
|
||||
|
||||
@@ -134,41 +134,6 @@ void DefineArray(py::module& m, const std::string& typestr) {
|
||||
}, py::keep_alive<0, 1>(), py::return_value_policy::reference_internal);
|
||||
};
|
||||
|
||||
// Specialization for std::byte to convert to int for Python iteration
|
||||
template <>
|
||||
void DefineArray<std::byte>(py::module& m, const std::string& typestr) {
|
||||
using Class = MjTypeVec<std::byte>;
|
||||
py::class_<Class>(m, typestr.c_str())
|
||||
.def(
|
||||
py::init([](std::byte* data, int size) { return Class(data, size); }))
|
||||
.def("__getitem__",
|
||||
[](Class& v, int i) -> int {
|
||||
if (i < 0 || i >= v.size) {
|
||||
throw py::index_error("Index out of range.");
|
||||
}
|
||||
return static_cast<int>(v.ptr[i]);
|
||||
})
|
||||
.def("__setitem__",
|
||||
[](Class& v, int i, int c) {
|
||||
if (i < 0 || i >= v.size) {
|
||||
throw py::index_error("Index out of range.");
|
||||
}
|
||||
if (c < 0 || c > 255) {
|
||||
throw py::value_error("Value out of range [0, 255].");
|
||||
}
|
||||
v.ptr[i] = static_cast<std::byte>(c);
|
||||
})
|
||||
.def("__len__", [](Class& v) { return v.size; })
|
||||
.def(
|
||||
"__iter__",
|
||||
[](Class& v) {
|
||||
return py::make_iterator(
|
||||
reinterpret_cast<unsigned char*>(v.ptr),
|
||||
reinterpret_cast<unsigned char*>(v.ptr + v.size));
|
||||
},
|
||||
py::keep_alive<0, 1>(), py::return_value_policy::reference_internal);
|
||||
};
|
||||
|
||||
py::list FindAllImpl(raw::MjsBody& body, mjtObj objtype, bool recursive) {
|
||||
py::list list;
|
||||
raw::MjsElement* el = mjs_firstChild(&body, objtype, recursive);
|
||||
|
||||
@@ -1040,7 +1040,7 @@ class SpecsTest(absltest.TestCase):
|
||||
spec = mujoco.MjSpec()
|
||||
texture = spec.add_texture(name='texture', height=1, width=2, nchannel=3)
|
||||
texture.data = bytes([1, 2, 3, 4, 5, 6])
|
||||
read_bytes = bytes(texture.data)
|
||||
read_bytes = texture.data
|
||||
self.assertEqual(read_bytes, bytes([1, 2, 3, 4, 5, 6]))
|
||||
|
||||
def test_modify_texture(self):
|
||||
@@ -1048,16 +1048,18 @@ class SpecsTest(absltest.TestCase):
|
||||
spec = mujoco.MjSpec()
|
||||
texture = spec.add_texture(name='texture', height=1, width=3, nchannel=3)
|
||||
texture.data = bytes([255, 0, 0, 0, 255, 0, 0, 0, 255])
|
||||
texture.data[1] = 255
|
||||
data_array = bytearray(texture.data)
|
||||
data_array[1] = 255
|
||||
texture.data = bytes(data_array)
|
||||
self.assertEqual(
|
||||
bytes(texture.data), bytes([255, 255, 0, 0, 255, 0, 0, 0, 255])
|
||||
texture.data, bytes([255, 255, 0, 0, 255, 0, 0, 0, 255])
|
||||
)
|
||||
|
||||
# Assigning values outside the range [0, 255] should raise an error.
|
||||
with self.assertRaises(ValueError):
|
||||
texture.data[3] = 256
|
||||
data_array[0] = 256
|
||||
with self.assertRaises(ValueError):
|
||||
texture.data[3] = -1
|
||||
data_array[0] = -1
|
||||
|
||||
def test_find_unnamed_asset(self):
|
||||
spec = mujoco.MjSpec()
|
||||
|
||||
Reference in New Issue
Block a user