# Copyright 2026 Huawei Technologies Co., Ltd
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# adopted from graspnet-baseline/test.py

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)


# Init datasets and dataloaders 
def my_worker_init_fn(worker_id):
    np.random.seed(np.random.get_state()[1][0] + worker_id)
    pass


# load test_seen dataset
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
# Init the model
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)

        # Forward pass
        with torch.no_grad():
            end_points = graspnet_model.infer(batch_data)
            grasp_preds = pred_decode(end_points)
        
        # Dump results for evaluation
        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)

            # collision detection
            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 grasps
            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()