import os
import sys
import argparse
import time
import numpy as np
import torch
import graspnet_npu_adaptor
from torch.utils.data import DataLoader
from graspnetAPI import GraspGroup, GraspNetEval
from graspnet_utils import collision_detection, get_om_path
from graspnet_om import GraspNetOM
GRASPNET_BASELINE = os.environ.get(
'GRASPNET_BASELINE',
os.path.join(os.path.dirname(os.path.abspath(__file__)), 'graspnet-baseline'),
)
for sub in ('models', 'utils', 'dataset'):
p = os.path.join(GRASPNET_BASELINE, sub)
if p not in sys.path:
sys.path.insert(0, p)
from graspnet import pred_decode
from graspnet_dataset import GraspNetDataset, collate_fn
parser = argparse.ArgumentParser()
parser.add_argument('--dataset_root', required=True, help='Dataset root')
parser.add_argument('--om_path', required=True)
parser.add_argument('--dump_dir', required=True, help='Dump dir to save outputs')
parser.add_argument('--camera', required=True, help='Camera split [realsense/kinect]')
parser.add_argument('--num_point', type=int, default=20000, help='Point Number [default: 20000]')
parser.add_argument('--batch_size', type=int, default=1, help='Batch Size during inference [default: 1]')
parser.add_argument('--collision_thresh', type=float, default=0.01, help='Collision Threshold in collision detection [default: 0.01]')
parser.add_argument('--voxel_size', type=float, default=0.01, help='Voxel Size to process point clouds before collision detection [default: 0.01]')
parser.add_argument('--num_workers', type=int, default=30, help='Number of workers used in evaluation [default: 30]')
args = parser.parse_args()
if not os.path.exists(args.dump_dir):
os.makedirs(args.dump_dir)
def my_worker_init_fn(worker_id):
np.random.seed(np.random.get_state()[1][0] + worker_id)
pass
TEST_DATASET = GraspNetDataset(args.dataset_root, valid_obj_idxs=None, grasp_labels=None, split='test_seen',
camera=args.camera, num_points=args.num_point, remove_outlier=True, augment=False, load_label=False)
print(f"num samples: {len(TEST_DATASET)}")
SCENE_LIST = TEST_DATASET.scene_list()
TEST_DATALOADER = DataLoader(TEST_DATASET, batch_size=args.batch_size, shuffle=False,
num_workers=4, worker_init_fn=my_worker_init_fn, collate_fn=collate_fn)
device = "cpu"
dtype = torch.float32
sa1_mlp, sa2_mlp, sa3_mlp, sa4_mlp, fp1, fp2_vp, grasp_generator = get_om_path(args.om_path)
graspnet_model = GraspNetOM(sa1_mlp, sa2_mlp, sa3_mlp, sa4_mlp, fp1, fp2_vp, grasp_generator)
def inference():
batch_interval = 100
start_time = time.time()
for batch_idx, batch_data in enumerate(TEST_DATALOADER):
for key in batch_data:
if 'list' in key:
for i in range(len(batch_data[key])):
for j in range(len(batch_data[key][i])):
batch_data[key][i][j] = batch_data[key][i][j].to(device)
else:
batch_data[key] = batch_data[key].to(device)
batch_data['point_clouds'] = batch_data['point_clouds'].to(dtype)
with torch.no_grad():
end_points = graspnet_model.infer(batch_data)
grasp_preds = pred_decode(end_points)
for i in range(args.batch_size):
data_idx = batch_idx * args.batch_size + i
preds = grasp_preds[i].detach().cpu().numpy()
grasp_group = GraspGroup(preds)
if args.collision_thresh > 0:
cloud, _ = TEST_DATASET.get_data(data_idx, return_raw_cloud=True)
grasp_group = collision_detection(grasp_group, cloud, args.voxel_size, args.collision_thresh)
save_dir = os.path.join(args.dump_dir, SCENE_LIST[data_idx], args.camera)
save_path = os.path.join(save_dir, str(data_idx % 256).zfill(4) + '.npy')
if not os.path.exists(save_dir):
os.makedirs(save_dir)
grasp_group.save_npy(save_path)
if batch_idx % batch_interval == 0:
end_time = time.time()
print('Eval batch: %d, time: %fs' % (batch_idx, (end_time - start_time) / batch_interval))
start_time = time.time()
def evaluate():
grasp_eval = GraspNetEval(root=args.dataset_root, camera=args.camera, split='test_seen')
result, ap = grasp_eval.eval_seen(args.dump_dir, proc=args.num_workers)
save_dir = os.path.join(args.dump_dir, 'ap_{}.npy'.format(args.camera))
np.save(save_dir, result)
print(f"test_seen AP: {ap}")
def main():
inference()
evaluate()
if __name__ == '__main__':
main()