diff --git a/DeepSDFStruct/optimization.py b/DeepSDFStruct/optimization.py index 95d587cc..4fa05be6 100644 --- a/DeepSDFStruct/optimization.py +++ b/DeepSDFStruct/optimization.py @@ -92,14 +92,18 @@ def get_mesh_from_torchfem(Solid: torchfem.Solid) -> pyvista.UnstructuredGrid: if not isinstance(Solid, torchfem.Solid): raise NotImplementedError("Currently only solid mesh is supported.") # VTK cell types - if Solid.etype is Tetra1: - cell_types = Solid.n_elem * [pyvista.CellType.TETRA] - elif Solid.etype is Tetra2: - cell_types = Solid.n_elem * [pyvista.CellType.QUADRATIC_TETRA] - elif Solid.etype is Hexa1: - cell_types = Solid.n_elem * [pyvista.CellType.HEXAHEDRON] - elif Solid.etype is Hexa2: - cell_types = Solid.n_elem * [pyvista.CellType.QUADRATIC_HEXAHEDRON] + etype = Solid.etype + + if etype is Tetra1 or isinstance(etype, Tetra1): + cell_types = [pyvista.CellType.TETRA] * Solid.n_elem + elif etype is Tetra2 or isinstance(etype, Tetra2): + cell_types = [pyvista.CellType.QUADRATIC_TETRA] * Solid.n_elem + elif etype is Hexa1 or isinstance(etype, Hexa1): + cell_types = [pyvista.CellType.HEXAHEDRON] * Solid.n_elem + elif etype is Hexa2 or isinstance(etype, Hexa2): + cell_types = [pyvista.CellType.QUADRATIC_HEXAHEDRON] * Solid.n_elem + else: + raise TypeError(f"Unsupported element type: {etype} ({type(etype)})") # VTK element list el = len(Solid.elements[0]) * torch.ones(Solid.n_elem, dtype=Solid.elements.dtype)