nbv_reconstruction/runners/strategy_generator.py

154 lines
8.0 KiB
Python
Raw Normal View History

2024-08-21 17:11:56 +08:00
import os
2024-08-22 20:27:21 +08:00
import json
import numpy as np
2024-08-21 17:11:56 +08:00
from PytorchBoot.runners.runner import Runner
from PytorchBoot.config import ConfigManager
from PytorchBoot.utils import Log
2024-08-21 17:11:56 +08:00
import PytorchBoot.stereotype as stereotype
2024-09-02 23:47:52 +08:00
from PytorchBoot.status import status_manager
2024-08-21 17:11:56 +08:00
2024-08-22 20:27:21 +08:00
from utils.data_load import DataLoadUtil
from utils.reconstruction import ReconstructionUtil
from utils.pts import PtsUtil
2024-08-22 20:27:21 +08:00
2024-08-22 22:28:20 +08:00
@stereotype.runner("strategy_generator")
2024-08-21 17:11:56 +08:00
class StrategyGenerator(Runner):
def __init__(self, config):
super().__init__(config)
2024-09-19 00:14:26 +08:00
self.load_experiment("generate_strategy")
2024-09-02 23:47:52 +08:00
self.status_info = {
"status_manager": status_manager,
2024-09-19 00:14:26 +08:00
"app_name": "generate_strategy",
2024-09-02 23:47:52 +08:00
"runner_name": "strategy_generator"
}
2024-09-12 15:11:09 +08:00
self.overwrite = ConfigManager.get("runner", "generate", "overwrite")
2024-09-23 14:30:51 +08:00
self.seq_num = ConfigManager.get("runner","generate","seq_num")
2024-10-23 02:58:58 -05:00
self.overlap_area_threshold = ConfigManager.get("runner","generate","overlap_area_threshold")
self.compute_with_normal = ConfigManager.get("runner","generate","compute_with_normal")
self.scan_points_threshold = ConfigManager.get("runner","generate","scan_points_threshold")
2024-09-10 20:12:46 +08:00
2024-09-08 19:43:01 +08:00
2024-08-21 17:11:56 +08:00
def run(self):
dataset_name_list = ConfigManager.get("runner", "generate", "dataset_list")
2024-10-23 02:58:58 -05:00
voxel_threshold = ConfigManager.get("runner","generate","voxel_threshold")
2024-09-02 23:47:52 +08:00
for dataset_idx in range(len(dataset_name_list)):
dataset_name = dataset_name_list[dataset_idx]
2024-09-19 00:14:26 +08:00
status_manager.set_progress("generate_strategy", "strategy_generator", "dataset", dataset_idx, len(dataset_name_list))
2024-08-22 20:27:21 +08:00
root_dir = ConfigManager.get("datasets", dataset_name, "root_dir")
2024-09-30 10:04:53 +08:00
from_idx = ConfigManager.get("datasets",dataset_name,"from")
to_idx = ConfigManager.get("datasets",dataset_name,"to")
2024-09-10 20:12:46 +08:00
scene_name_list = os.listdir(root_dir)
2024-09-30 10:04:53 +08:00
if to_idx == -1:
to_idx = len(scene_name_list)
cnt = 0
2024-10-02 16:24:13 +08:00
total = len(scene_name_list[from_idx:to_idx])
2024-09-30 10:04:53 +08:00
Log.info(f"Processing Dataset: {dataset_name}, From: {from_idx}, To: {to_idx}")
for scene_name in scene_name_list[from_idx:to_idx]:
Log.info(f"({dataset_name})Processing [{cnt}/{total}]: {scene_name}")
2024-09-19 00:14:26 +08:00
status_manager.set_progress("generate_strategy", "strategy_generator", "scene", cnt, total)
2024-09-23 14:30:51 +08:00
output_label_path = DataLoadUtil.get_label_path(root_dir, scene_name,0)
2024-09-12 15:11:09 +08:00
if os.path.exists(output_label_path) and not self.overwrite:
Log.info(f"Scene <{scene_name}> Already Exists, Skip")
cnt += 1
continue
2024-10-02 16:24:13 +08:00
2024-10-23 02:58:58 -05:00
self.generate_sequence(root_dir, scene_name,voxel_threshold)
cnt += 1
2024-09-19 00:14:26 +08:00
status_manager.set_progress("generate_strategy", "strategy_generator", "scene", total, total)
status_manager.set_progress("generate_strategy", "strategy_generator", "dataset", len(dataset_name_list), len(dataset_name_list))
2024-08-21 17:11:56 +08:00
def create_experiment(self, backup_name=None):
super().create_experiment(backup_name)
output_dir = os.path.join(str(self.experiment_path), "output")
os.makedirs(output_dir)
def load_experiment(self, backup_name=None):
super().load_experiment(backup_name)
2024-10-23 02:58:58 -05:00
def generate_sequence(self, root, scene_name, voxel_threshold):
2024-09-19 00:14:26 +08:00
status_manager.set_status("generate_strategy", "strategy_generator", "scene", scene_name)
frame_num = DataLoadUtil.get_scene_seq_length(root, scene_name)
2024-10-17 14:28:19 +00:00
2024-09-08 19:43:01 +08:00
model_points_normals = DataLoadUtil.load_points_normals(root, scene_name)
model_pts = model_points_normals[:,:3]
2024-10-23 02:58:58 -05:00
down_sampled_model_pts, idx = PtsUtil.voxel_downsample_point_cloud(model_pts, voxel_threshold, require_idx=True)
down_sampled_model_nrm = model_points_normals[idx, 3:]
2024-08-21 17:11:56 +08:00
pts_list = []
2024-10-23 02:58:58 -05:00
nrm_list = []
scan_points_indices_list = []
2024-10-02 16:24:13 +08:00
non_zero_cnt = 0
2024-10-05 15:10:31 -05:00
for frame_idx in range(frame_num):
2024-10-02 16:24:13 +08:00
status_manager.set_progress("generate_strategy", "strategy_generator", "loading frame", frame_idx, frame_num)
2024-10-05 15:10:31 -05:00
pts_path = os.path.join(root,scene_name, "pts", f"{frame_idx}.npy")
2024-10-23 02:58:58 -05:00
nrm_path = os.path.join(root,scene_name, "nrm", f"{frame_idx}.npy")
2024-10-05 15:10:31 -05:00
idx_path = os.path.join(root,scene_name, "scan_points_indices", f"{frame_idx}.npy")
2024-10-28 16:48:34 +00:00
2024-10-23 02:58:58 -05:00
pts = np.load(pts_path)
2024-10-28 18:25:53 +00:00
if self.compute_with_normal:
if pts.shape[0] == 0:
nrm = np.zeros((0,3))
else:
nrm = np.load(nrm_path)
nrm_list.append(nrm)
2024-10-23 02:58:58 -05:00
pts_list.append(pts)
2024-10-28 18:25:53 +00:00
indices = np.load(idx_path)
2024-10-03 01:59:13 +08:00
scan_points_indices_list.append(indices)
2024-10-23 02:58:58 -05:00
if pts.shape[0] > 0:
2024-10-05 15:10:31 -05:00
non_zero_cnt += 1
2024-09-19 00:14:26 +08:00
status_manager.set_progress("generate_strategy", "strategy_generator", "loading frame", frame_num, frame_num)
2024-10-05 15:10:31 -05:00
2024-10-02 16:24:13 +08:00
seq_num = min(self.seq_num, non_zero_cnt)
init_view_list = []
2024-10-05 15:10:31 -05:00
idx = 0
2024-10-11 16:34:16 +08:00
while len(init_view_list) < seq_num and idx < len(pts_list):
2024-10-23 02:58:58 -05:00
if pts_list[idx].shape[0] > 50:
2024-10-05 15:10:31 -05:00
init_view_list.append(idx)
idx += 1
2024-09-10 20:12:46 +08:00
2024-09-23 14:30:51 +08:00
seq_idx = 0
2024-10-05 15:10:31 -05:00
import time
2024-09-23 14:30:51 +08:00
for init_view in init_view_list:
status_manager.set_progress("generate_strategy", "strategy_generator", "computing sequence", seq_idx, len(init_view_list))
2024-10-05 15:10:31 -05:00
start = time.time()
2024-10-23 02:58:58 -05:00
if not self.compute_with_normal:
limited_useful_view, _, _ = ReconstructionUtil.compute_next_best_view_sequence(down_sampled_model_pts, pts_list, scan_points_indices_list = scan_points_indices_list,init_view=init_view,
threshold=voxel_threshold, scan_points_threshold=self.scan_points_threshold, overlap_area_threshold=self.overlap_area_threshold, status_info=self.status_info)
else:
limited_useful_view, _, _ = ReconstructionUtil.compute_next_best_view_sequence_with_normal(down_sampled_model_pts, down_sampled_model_nrm, pts_list, nrm_list, scan_points_indices_list = scan_points_indices_list,init_view=init_view,
threshold=voxel_threshold, scan_points_threshold=self.scan_points_threshold, overlap_area_threshold=self.overlap_area_threshold, status_info=self.status_info)
2024-10-05 15:10:31 -05:00
end = time.time()
print(f"Time: {end-start}")
2024-09-23 14:30:51 +08:00
data_pairs = self.generate_data_pairs(limited_useful_view)
seq_save_data = {
"data_pairs": data_pairs,
"best_sequence": limited_useful_view,
"max_coverage_rate": limited_useful_view[-1][1]
}
status_manager.set_status("generate_strategy", "strategy_generator", "max_coverage_rate", limited_useful_view[-1][1])
Log.success(f"Scene <{scene_name}> Finished, Max Coverage Rate: {limited_useful_view[-1][1]}, Best Sequence length: {len(limited_useful_view)}")
2024-09-10 20:12:46 +08:00
2024-09-23 14:30:51 +08:00
output_label_path = DataLoadUtil.get_label_path(root, scene_name, seq_idx)
with open(output_label_path, 'w') as f:
json.dump(seq_save_data, f)
seq_idx += 1
status_manager.set_progress("generate_strategy", "strategy_generator", "computing sequence", len(init_view_list), len(init_view_list))
2024-09-10 20:12:46 +08:00
2024-08-22 20:27:21 +08:00
def generate_data_pairs(self, useful_view):
data_pairs = []
2024-08-30 17:57:47 +08:00
for next_view_idx in range(1, len(useful_view)):
2024-08-22 20:27:21 +08:00
scanned_views = useful_view[:next_view_idx]
next_view = useful_view[next_view_idx]
data_pairs.append((scanned_views, next_view))
return data_pairs
2024-08-21 17:11:56 +08:00