File size: 2,984 Bytes
7ab05dd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 | """
PartNetE Datasets preprocessing
Author: Yujia Zhang (yujia.zhang.cs@gmail.com)
Please cite our work if the code is helpful to you.
"""
import os
import shutil
import argparse
import numpy as np
import trimesh
import open3d as o3d
def process_folder(target_dir):
ply_path = os.path.join(target_dir, "pc.ply")
label_path = os.path.join(target_dir, "label.npy")
if os.path.exists(ply_path):
mesh = trimesh.load(ply_path, process=False)
coords = np.array(mesh.vertices, dtype=np.float32)
np.save(os.path.join(target_dir, "coord.npy"), coords)
if (
hasattr(mesh, "vertex_normals")
and mesh.vertex_normals is not None
and len(mesh.vertex_normals) == len(coords)
):
normals = np.array(mesh.vertex_normals, dtype=np.float32)
else:
pcd = o3d.geometry.PointCloud()
pcd.points = o3d.utility.Vector3dVector(coords)
search_param = o3d.geometry.KDTreeSearchParamHybrid(radius=0.1, max_nn=30)
pcd.estimate_normals(search_param=search_param)
normals = np.asarray(pcd.normals).astype(np.float32)
np.save(os.path.join(target_dir, "normal.npy"), normals)
if (
hasattr(mesh.visual, "vertex_colors")
and mesh.visual.vertex_colors is not None
):
colors = np.array(mesh.visual.vertex_colors[:, :3], dtype=np.uint8)
np.save(os.path.join(target_dir, "color.npy"), colors)
elif hasattr(mesh, "colors") and mesh.colors is not None:
colors = np.array(mesh.colors[:, :3], dtype=np.uint8)
np.save(os.path.join(target_dir, "color.npy"), colors)
else:
pass
label_data = np.load(label_path, allow_pickle=True).item()
segment = label_data["semantic_seg"]
instance = label_data["instance_seg"]
assert coords.shape[0] == segment.shape[0]
segment_path = os.path.join(target_dir, "segment.npy")
instance_path = os.path.join(target_dir, "instance.npy")
np.save(segment_path, segment)
np.save(instance_path, instance)
else:
print(f"Warning: pc.ply not found in {target_dir}")
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="Preprocess PartNetE dataset by splitting PLY files and copying labels."
)
parser.add_argument(
"--dataset_root",
required=True,
help="Path to the ScanNet dataset containing scene folders",
)
args = parser.parse_args()
subfolders = ["few_shot", "test"]
total_processed_count = 0
for subfolder in subfolders:
current_root = os.path.join(args.dataset_root, subfolder)
processed_in_subfolder = 0
for dirpath, dirnames, filenames in os.walk(current_root):
if "pc.ply" in filenames and "label.npy" in filenames:
process_folder(dirpath)
processed_in_subfolder += 1
|