Parses a mujoco xml scene description file and returns a Skeleton Tree. We use the model attribute at the root as the name of the tree. :param path: :type path: string :return: The skeleton tree constructed from the mjcf file :rtype: SkeletonTree
(cls, path: str)
| 287 | |
| 288 | @classmethod |
| 289 | def from_mjcf(cls, path: str) -> "SkeletonTree": |
| 290 | """ |
| 291 | Parses a mujoco xml scene description file and returns a Skeleton Tree. |
| 292 | We use the model attribute at the root as the name of the tree. |
| 293 | |
| 294 | :param path: |
| 295 | :type path: string |
| 296 | :return: The skeleton tree constructed from the mjcf file |
| 297 | :rtype: SkeletonTree |
| 298 | """ |
| 299 | tree = ET.parse(path) |
| 300 | xml_doc_root = tree.getroot() |
| 301 | xml_world_body = xml_doc_root.find("worldbody") |
| 302 | if xml_world_body is None: |
| 303 | raise ValueError("MJCF parsed incorrectly please verify it.") |
| 304 | # assume this is the root |
| 305 | xml_body_root = xml_world_body.find("body") |
| 306 | if xml_body_root is None: |
| 307 | raise ValueError("MJCF parsed incorrectly please verify it.") |
| 308 | |
| 309 | node_names = [] |
| 310 | parent_indices = [] |
| 311 | local_translation = [] |
| 312 | |
| 313 | # recursively adding all nodes into the skel_tree |
| 314 | def _add_xml_node(xml_node, parent_index, node_index): |
| 315 | node_name = xml_node.attrib.get("name") |
| 316 | # parse the local translation into float list |
| 317 | pos = np.fromstring(xml_node.attrib.get("pos", "0 0 0"), dtype=float, sep=" ") |
| 318 | node_names.append(node_name) |
| 319 | parent_indices.append(parent_index) |
| 320 | local_translation.append(pos) |
| 321 | curr_index = node_index |
| 322 | node_index += 1 |
| 323 | for next_node in xml_node.findall("body"): |
| 324 | node_index = _add_xml_node(next_node, curr_index, node_index) |
| 325 | return node_index |
| 326 | |
| 327 | _add_xml_node(xml_body_root, -1, 0) |
| 328 | |
| 329 | return cls( |
| 330 | node_names, |
| 331 | torch.from_numpy(np.array(parent_indices, dtype=np.int32)), |
| 332 | torch.from_numpy(np.array(local_translation, dtype=np.float32)), |
| 333 | ) |
| 334 | |
| 335 | def parent_of(self, node_name): |
| 336 | """get the name of the parent of the given node |