TaxonomyProject/SimulationImage
01.5k
1import copy
2import json
3import os
4from typing import Any
5
6import cv2
7import matplotlib.pyplot as plt
8import numpy as np
9from PIL import Image
10
11from lychsim.api import LychSim
12from lychsim.utils.camera_projection_utils import project_3d_to_2d, get_bbox3d
13
14from dataclasses import dataclass
15from typing import List, Optional, Dict, Tuple
16from scipy.spatial import cKDTree
17from collections import defaultdict
18
19
20
21class EasyDict(dict):
22 """Convenience class that behaves like a dict but allows access with the attribute syntax."""
23
24 def __getattr__(self, name: str) -> Any:
25 try:
26 return self[name]
27 except KeyError:
28 raise AttributeError(name)
29
30 def __setattr__(self, name: str, value: Any) -> None:
31 self[name] = value
32
33 def __delattr__(self, name: str) -> None:
34 del self[name]
35
36
37def init_sampling_params(state):
38 # list of table and floor objects
39 # will be provided by Xingrui and Siyi
40 state.floor_objects = [
41 "/Game/ManagerOffice/Meshes/Props/SM_AmchairTreadle.SM_AmchairTreadle",
42 "/Game/ManagerOffice/Meshes/Props/SM_ArmchairManager.SM_ArmchairManager",
43 "/Game/ManagerOffice/Meshes/Props/SM_ColumnTable.SM_ColumnTable",
44 "/Game/ManagerOffice/Meshes/Props/SM_Decorative17.SM_Decorative17",
45 "/Game/ManagerOffice/Meshes/Props/SM_Komod.SM_Komod",
46 "/Game/ManagerOffice/Meshes/Props/SM_KomodB.SM_KomodB",
47 "/Game/ManagerOffice/Meshes/Props/SM_Plant2.SM_Plant2",
48 "/Game/ManagerOffice/Meshes/Props/SM_Plant1.SM_Plant1",
49 "/Game/ManagerOffice/Meshes/Props/SM_TeaTable.SM_TeaTable",
50 ]
51 state.table_objects = [
52 "/Game/ManagerOffice/Meshes/Props/SM_Ashtray.SM_Ashtray",
53 "/Game/ManagerOffice/Meshes/Props/SM_Award3.SM_Award3",
54 "/Game/ManagerOffice/Meshes/Props/SM_Award9.SM_Award9",
55 "/Game/ManagerOffice/Meshes/Props/SM_Book2.SM_Book2",
56 "/Game/ManagerOffice/Meshes/Props/SM_CalendarDesk.SM_CalendarDesk",
57 "/Game/ManagerOffice/Meshes/Props/SM_Decorative10.SM_Decorative10",
58 "/Game/ManagerOffice/Meshes/Props/SM_Decorative37.SM_Decorative37",
59 "/Game/ManagerOffice/Meshes/Props/SM_Fruits.SM_Fruits",
60 "/Game/ManagerOffice/Meshes/Props/SM_PC.SM_PC",
61 "/Game/ManagerOffice/Meshes/Props/SM_MarkerMug.SM_MarkerMug",
62 ]
63
64 mesh_extents = state.sim.get_mesh_extent(state.floor_objects + state.table_objects)[
65 "outputs"
66 ]
67 state.mesh_extents = {
68 x["mesh_path"]: x["extent"] for x in mesh_extents if x["status"] == "ok"
69 }
70
71 for x in state.floor_objects:
72 if x not in state.mesh_extents:
73 print(f"Warning: Floor object {x} not found in the scene.")
74 for x in state.table_objects:
75 if x not in state.mesh_extents:
76 print(f"Warning: Table object {x} not found in the scene.")
77
78 state.floor_objects = [x for x in state.floor_objects if x in state.mesh_extents]
79 state.table_objects = [x for x in state.table_objects if x in state.mesh_extents]
80
81 state.table_height_margin_low, state.table_height_margin_high = (
82 -30.0,
83 50.0,
84 ) # table object hit box [top-30, top+50]
85 state.table_object_threshold = (
86 0.75 # IoA threshold: intersection over object volume
87 )
88
89 # number of trials to sample floor objects
90 state.max_floor_sampling_trials = 20
91 # IoU threshold for floor object collision detection
92 state.floor_object_collision_iou_thr = 0.1
93 # threshold for worst addition on floor: if worse than this, skip adding floor objects
94 state.worst_floor_addition = -10
95
96 # number of trials to sample table objects
97 state.max_table_sampling_trials = 20
98 # IoU threshold for table object collision detection
99 state.table_object_collision_iou_thr = 0.1
100 # threshold for worst addition on table: if worse than this, skip adding table objects
101 state.worst_table_addition = -10
102
103
104def add_selection_as_floor(state, num_objects):
105 objects = state.sim.list_selected()
106 if objects["status"] != "ok":
107 raise RuntimeError(f"Failed to get selected objects. Response: {objects}")
108
109 new_floors = []
110 for obj in objects["outputs"]:
111 obj_id = obj["object_id"]
112 new_floors.append((obj_id, num_objects))
113
114 before_count = len(state.floors)
115 state.floors.update(new_floors)
116
117 print(
118 f"Added {len(new_floors)} object(s) to the floor list (prev={before_count} "
119 f"-> now={len(state.floors)}):\n{state.floors}"
120 )
121
122
123def add_selection_as_table(state, num_objects):
124 objects = state.sim.list_selected()
125 if objects["status"] != "ok":
126 raise RuntimeError(f"Failed to get selected objects. Response: {objects}")
127
128 new_tables = []
129 for obj in objects["outputs"]:
130 obj_id = obj["object_id"]
131 new_tables.append((obj_id, num_objects))
132
133 before_count = len(state.tables)
134 state.tables.update(new_tables)
135
136 print(
137 f"Added {len(new_tables)} object(s) to the table list (prev={before_count} "
138 f"-> now={len(state.tables)}):\n{state.tables}"
139 )
140
141
142def add_camera_location(state):
143 cam_id = state.cam_id
144 loc = state.sim.get_cam_loc(0)
145
146 before_count = len(state.cam_locations)
147 state.cam_locations.append(loc)
148
149 print(f"New location added (prev={before_count} -> {len(state.cam_locations)}):")
150 for loc in state.cam_locations:
151 print(f"\t{loc}")
152
153
154def get_objects_on_aabb(state, table_aabb, objs_aabb):
155 table_aabb, objs_aabb = copy.deepcopy(table_aabb), copy.deepcopy(objs_aabb)
156 object_list = []
157 target_center, target_extent = table_aabb["center"], table_aabb["extent"]
158
159 # we compute the space above the table
160 state.table_height_margin_low, state.table_height_margin_high = -30.0, 50.0
161 target_center[2] = (
162 target_center[2]
163 + target_extent[2]
164 + (state.table_height_margin_low + state.table_height_margin_high) / 2.0
165 )
166 target_extent[2] = (
167 state.table_height_margin_high - state.table_height_margin_low
168 ) / 2.0
169
170 tgt_min = np.array(target_center) - np.array(target_extent)
171 tgt_max = np.array(target_center) + np.array(target_extent)
172
173 for aabb in objs_aabb:
174 if aabb["status"] != "ok" or aabb["object_id"] == table_aabb["object_id"]:
175 continue
176 aabb["extent"] = [max(x, 1e-6) for x in aabb["extent"]]
177 obj_min = np.array(aabb["center"]) - np.array(aabb["extent"])
178 obj_max = np.array(aabb["center"]) + np.array(aabb["extent"])
179
180 inter_min = np.maximum(obj_min, tgt_min)
181 inter_max = np.minimum(obj_max, tgt_max)
182 inter_extent = np.maximum(0.0, inter_max - inter_min)
183 inter_vol = np.prod(inter_extent)
184
185 obj_vol = np.prod(2 * np.array(aabb["extent"]))
186
187 if inter_vol / obj_vol >= state.table_object_threshold:
188 object_list.append(aabb["object_id"])
189
190 return object_list
191
192
193def clear_table_objects(state, table_id, objs_aabb):
194 objs_aabb = copy.deepcopy(objs_aabb)
195 table_aabb = [x for x in objs_aabb if x["object_id"] == table_id][0]
196 objects_on_table = get_objects_on_aabb(state, table_aabb, objs_aabb)
197 for obj_id in objects_on_table:
198 state.sim.del_obj(obj_id)
199
200
201def collide(center1, extent1, center2, extent2, thr):
202 center1, extent1 = np.array(center1), np.array(extent1)
203 center2, extent2 = np.array(center2), np.array(extent2)
204
205 min1, max1 = center1 - extent1, center1 + extent1
206 min2, max2 = center2 - extent2, center2 + extent2
207
208 inter_min = np.maximum(min1, min2)
209 inter_max = np.minimum(max1, max2)
210 inter_extent = np.maximum(0.0, inter_max - inter_min)
211 inter_vol = np.prod(inter_extent)
212
213 vol1, vol2 = np.prod(2 * extent1), np.prod(2 * extent2)
214 union_vol = vol1 + vol2 - inter_vol
215
216 iou = inter_vol / union_vol if union_vol > 0 else 0.0
217 return iou >= thr
218
219
220def compute_addition_from_collision(state, objs_aabb, sampling):
221 addition = len(sampling)
222
223 # first check mutual collisions
224 for obj1 in sampling:
225 for obj2 in sampling:
226 if obj1 >= obj2:
227 continue
228 if collide(
229 sampling[obj1]["center"],
230 sampling[obj1]["extent"],
231 sampling[obj2]["center"],
232 sampling[obj2]["extent"],
233 state.floor_object_collision_iou_thr,
234 ):
235 return -1e5, []
236
237 tables = [x[0] for x in state.tables]
238 all_collided_objects = []
239 for obj in sampling:
240 collided_objects = [
241 x
242 for x in objs_aabb
243 if collide(
244 x["center"],
245 x["extent"],
246 sampling[obj]["center"],
247 sampling[obj]["extent"],
248 state.floor_object_collision_iou_thr,
249 )
250 ]
251 for x in collided_objects:
252 if x["object_id"] in tables:
253 return -1e5, []
254 addition -= len(collided_objects)
255 all_collided_objects.extend([x["object_id"] for x in collided_objects])
256 return addition, all_collided_objects
257
258
259def sample_floor_objects(state, floor_id, num_objects, objs_aabb):
260 floor_aabb = state.sim.get_obj_aabb(floor_id)["outputs"][0]
261 target_center, target_extent = np.array(floor_aabb["center"]), np.array(
262 floor_aabb["extent"]
263 )
264 target_extent[0] *= 0.9
265 target_extent[1] *= 0.9
266
267 best_sampling, best_addition, best_collisions = None, -1e6, None
268 for _ in range(state.max_floor_sampling_trials):
269 sampling = {}
270 sampled_object_ids = [
271 state.floor_objects[i]
272 for i in np.random.choice(
273 len(state.floor_objects), num_objects, replace=False
274 )
275 ]
276 for soi in sampled_object_ids:
277 horizontal_location = target_center[:2] + np.random.uniform(
278 -target_extent[:2] * 0.5, target_extent[:2] * 0.5
279 )
280 vertical_location = target_center[2] + target_extent[2]
281 sampling[soi] = dict(
282 center=list(horizontal_location) + [vertical_location],
283 extent=state.mesh_extents[soi],
284 )
285 addition, collisions = compute_addition_from_collision(
286 state, objs_aabb, sampling
287 )
288 if addition > best_addition:
289 best_addition = addition
290 best_sampling = sampling
291 best_collisions = collisions
292 if best_addition < state.worst_floor_addition:
293 # print(f"Best addition: {best_addition}, collisions: {best_collisions}")
294 return None
295
296 for obj_id in best_collisions:
297 state.sim.del_obj(obj_id)
298 # print(f"del {obj_id}")
299 for obj_id in best_sampling:
300 loc = best_sampling[obj_id]["center"]
301 rot = [0.0, float(np.random.uniform(0, 360)), 0.0]
302 state.sim.add_obj(f"{obj_id.split('.')[-1]}_{random_uuid()}", obj_id, loc, rot)
303 # print(f"add {obj_id}, {loc}, {rot}")
304
305
306def sample_table_objects(state, table_id, num_objects, objs_aabb):
307 table_aabb = state.sim.get_obj_aabb(table_id)["outputs"][0]
308 target_center, target_extent = np.array(table_aabb["center"]), np.array(
309 table_aabb["extent"]
310 )
311 target_extent[0] *= 0.9
312 target_extent[1] *= 0.9
313
314 best_sampling, best_addition, best_collisions = None, -1e6, None
315 for _ in range(state.max_table_sampling_trials):
316 sampling = {}
317 sampled_object_ids = [
318 state.table_objects[i]
319 for i in np.random.choice(
320 len(state.table_objects), num_objects, replace=False
321 )
322 ]
323 for soi in sampled_object_ids:
324 horizontal_location = target_center[:2] + np.random.uniform(
325 -target_extent[:2] * 0.5, target_extent[:2] * 0.5
326 )
327 vertical_location = target_center[2] + target_extent[2]
328 sampling[soi] = dict(
329 center=list(horizontal_location) + [vertical_location],
330 extent=state.mesh_extents[soi],
331 )
332 addition, collisions = compute_addition_from_collision(
333 state, objs_aabb, sampling
334 )
335 if addition > best_addition:
336 best_addition = addition
337 best_sampling = sampling
338 best_collisions = collisions
339 if best_addition < state.worst_table_addition:
340 print(f"Best addition: {best_addition}, collisions: {best_collisions}")
341 return None
342
343 for obj_id in best_collisions:
344 state.sim.del_obj(obj_id)
345 print(f"del {obj_id}")
346 for obj_id in best_sampling:
347 loc = best_sampling[obj_id]["center"]
348 rot = [0.0, float(np.random.uniform(0, 360)), 0.0]
349 state.sim.add_obj(f"{obj_id.split('.')[-1]}_{random_uuid()}", obj_id, loc, rot)
350 print(f"add {obj_id}, {loc}, {rot}")
351
352
353def sample_random_placement(state):
354 objs_aabb = state.sim.get_obj_aabb()["outputs"]
355
356 for floor_id, num_objects in state.floors:
357 sample_floor_objects(state, floor_id, num_objects, objs_aabb)
358
359 for table_id, num_objects in state.tables:
360 clear_table_objects(state, table_id, objs_aabb)
361 sample_table_objects(state, table_id, num_objects, objs_aabb)
362
363
364def get_random_camera_rotations(state):
365 def sample_rotation():
366 pitch = float(np.random.uniform(state.min_pitch, state.max_pitch))
367 yaw = float(np.random.uniform(0, 360))
368 roll = 0.0
369 return [pitch, yaw, roll]
370
371 return [sample_rotation() for _ in range(state.random_viewpoints_per_location)]
372
373def get_random_camera_rotations_fixed_yaw(state):
374 yaw_list = np.arange(0, 360, 60)
375 def sample_rotation(i):
376 pitch = 0.0
377 yaw = yaw_list[i]
378 roll = 0.0
379 return [pitch, yaw, roll]
380 return [sample_rotation(i) for i in range(len(yaw_list))]
381
382def add_random_camera_height_offset(loc, state):
383 offset = float(
384 np.random.uniform(
385 -state.random_camera_height_offset, state.random_camera_height_offset
386 )
387 )
388 new_loc = loc.copy()
389 new_loc[2] += offset
390 return new_loc
391
392
393def set_camera_location_and_rotation(scene_state, cam_loc_final, cam_rot):
394 cam_id = scene_state.cam_id
395 sim = scene_state.sim
396
397 sim.set_cam_loc(cam_id, cam_loc_final)
398 sim.set_cam_rot(cam_id, cam_rot)
399
400
401def save_state(scene_state):
402 save_state = {}
403 for k in scene_state:
404 if isinstance(scene_state[k], LychSim):
405 save_state[k] = str(type(scene_state[k]))
406 elif isinstance(scene_state[k], set):
407 save_state[k] = list(scene_state[k])
408 else:
409 save_state[k] = scene_state[k]
410
411 save_path = os.path.join(scene_state.save_path, scene_state.scene_name)
412 os.makedirs(save_path, exist_ok=True)
413
414 with open(os.path.join(save_path, "state.json"), "w") as f:
415 json.dump(save_state, f, indent=4)
416
417
418def capture_and_save(scene_state, view_name, camera_warmup_steps=10):
419 scene_output_path = os.path.join(
420 scene_state.save_path, scene_state.scene_name, view_name
421 )
422 os.makedirs(scene_output_path, exist_ok=True)
423
424 scene_state.sim.warmup_cam(scene_state.cam_id, camera_warmup_steps)
425 image = scene_state.sim.get_cam_lit(scene_state.cam_id)
426 image.save(os.path.join(scene_output_path, "lit.png"))
427
428 seg = scene_state.sim.get_cam_seg(scene_state.cam_id)
429 seg.save(os.path.join(scene_output_path, "seg.png"))
430
431 depth = scene_state.sim.get_cam_depth(scene_state.cam_id)
432 np.save(os.path.join(scene_output_path, "depth.npy"), depth)
433
434 normal = scene_state.sim.get_cam_normal(scene_state.cam_id)
435 normal.save(os.path.join(scene_output_path, "normal.png"))
436
437 annots_obj = scene_state.sim.get_obj_annots()
438 with open(os.path.join(scene_output_path, "object_annots.json"), "w") as f:
439 json.dump(annots_obj, f)
440
441 annots_cam = scene_state.sim.get_cam_annots(scene_state.cam_id)
442 fov = annots_cam["outputs"]["fov"]
443 w = annots_cam["outputs"]["width"]
444 h = annots_cam["outputs"]["height"]
445 fovx = np.deg2rad(fov)
446 fx = 0.5 * w / np.tan(0.5 * fovx)
447 fovy = 2.0 * np.arctan((h / float(w)) * np.tan(0.5 * fovx))
448 fy = 0.5 * h / np.tan(0.5 * fovy)
449 annots_cam["outputs"]["fxfycxcy"] = [fx, fy, w / 2.0, h / 2.0]
450 with open(os.path.join(scene_output_path, "camera_annots.json"), "w") as f:
451 json.dump(annots_cam, f)
452
453 scene_state.sim.clear_annot_comps()
454
455def capture_and_save_filter(scene_state, view_name, camera_warmup_steps=10):
456 scene_output_path = os.path.join(
457 scene_state.save_path, scene_state.scene_name, view_name
458 )
459 os.makedirs(scene_output_path, exist_ok=True)
460
461
462 seg = scene_state.sim.get_cam_seg(scene_state.cam_id)
463 seg.save(os.path.join(scene_output_path, "seg.png"))
464
465 depth = scene_state.sim.get_cam_depth(scene_state.cam_id)
466 np.save(os.path.join(scene_output_path, "depth.npy"), depth)
467
468 annots_obj = scene_state.sim.get_obj_annots()
469 with open(os.path.join(scene_output_path, "object_annots.json"), "w") as f:
470 json.dump(annots_obj, f)
471
472 annots_cam = scene_state.sim.get_cam_annots(scene_state.cam_id)
473 fov = annots_cam["outputs"]["fov"]
474 w = annots_cam["outputs"]["width"]
475 h = annots_cam["outputs"]["height"]
476 fovx = np.deg2rad(fov)
477 fx = 0.5 * w / np.tan(0.5 * fovx)
478 fovy = 2.0 * np.arctan((h / float(w)) * np.tan(0.5 * fovx))
479 fy = 0.5 * h / np.tan(0.5 * fovy)
480 annots_cam["outputs"]["fxfycxcy"] = [fx, fy, w / 2.0, h / 2.0]
481 with open(os.path.join(scene_output_path, "camera_annots.json"), "w") as f:
482 json.dump(annots_cam, f)
483
484 scene_state.sim.clear_annot_comps()
485
486
487def capture_and_save_image(scene_state, view_name, camera_warmup_steps=10):
488 scene_output_path = os.path.join(
489 scene_state.save_path, scene_state.scene_name, view_name
490 )
491 os.makedirs(scene_output_path, exist_ok=True)
492
493 scene_state.sim.warmup_cam(scene_state.cam_id, camera_warmup_steps)
494 image = scene_state.sim.get_cam_lit(scene_state.cam_id)
495 image.save(os.path.join(scene_output_path, "lit.png"))
496
497
498def visualize_bbox(img, corners_2d, edges, color=(255, 255, 0, 255), thickness=2):
499 for i, j in edges:
500 pt1 = (int(corners_2d[i, 0]), int(corners_2d[i, 1]))
501 pt2 = (int(corners_2d[j, 0]), int(corners_2d[j, 1]))
502 cv2.line(img, pt1, pt2, color, thickness)
503 plt.imshow(img)
504 return img
505
506
507def draw_bbox_3d(img, center, extent, c2w, fov):
508 if isinstance(img, Image.Image):
509 img = np.array(img)
510 vis_img = np.array(img).copy()
511 corners, edges = get_bbox3d(center=center, extent=extent)
512 pts2d, in_front = project_3d_to_2d(corners, c2w, fov, 1920, 1080)
513 vis_img = visualize_bbox(vis_img, pts2d, edges, color=(0, 255, 0, 255))
514 return Image.fromarray(vis_img)
515
516
517def random_uuid(length=4):
518 return "".join(
519 np.random.choice(list("abcdefghijklmnopqrstuvwxyz0123456789"), size=length)
520 )
521
522
523class CameraPositionEvaluator:
524 """
525 相机位置质量评估器 - 判断深度图和分割掩码质量是否合格
526
527 评分权重:深度40% + 分割60%
528
529 分割要求(非常严格):
530 - ⚠️ 物体总数<6个,分割评分直接返回0,必定不合格
531 - ⚠️ 任何物体占比>50%,分割评分直接返回0,必定不合格
532 - 物体总数≥20个为满分,12-20个部分得分
533 - 小物体(占比<5%)需要≥6个
534 - 平均物体占比2-8%为理想
535 - 最大物体占比理想范围10%-30%
536
537 深度异常值检测策略(极其严格):
538 - 自动过滤常见的无效深度值(65504, 65535, 0等)
539 - ⚠️ 如果有效深度<90%(即无效值>10%),深度评分直接返回0,必定不合格
540 - 检测深度单一性:如果深度值过于集中(如一面墙),会被降分
541 - 智能判断:如果深度有足够变化(标准差/熵高,说明墙前有物品),则放宽集中度要求
542 - 使用中位数而非均值计算比值(更鲁棒,不受极端值影响)
543 - 最大深度/中位数比:检测单个极端异常值
544 - 离群值占比:检测多个异常大的深度值(室外空旷区域)
545 """
546
547 def __init__(self, threshold: float = 0.6, background_color: Tuple[int, int, int] = (0, 0, 0)):
548 """
549 参数:
550 threshold: 合格阈值,0-1之间,默认0.6
551 background_color: 背景颜色RGB值,默认为黑色(0, 0, 0)
552 """
553 self.threshold = threshold
554 self.depth_weight = 0.4 # 深度权重
555 self.seg_weight = 0.6 # 分割权重
556 self.background_color = background_color
557
558 def evaluate(self, depth_map: np.ndarray, seg_mask: np.ndarray) -> Dict:
559 """
560 评估相机位置是否合格
561
562 参数:
563 depth_map: 深度图 (H, W),单位米
564 seg_mask: 分割掩码 (H, W, 4),值为RGBA颜色,格式为(r, g, b, 255)
565
566 返回:
567 包含评估结果的字典:
568 {
569 'is_qualified': bool, # 是否合格
570 'score': float, # 总评分 0-1
571 'depth_score': float, # 深度评分
572 'seg_score': float, # 分割评分
573 'details': dict # 详细指标
574 }
575 """
576 # 验证分割掩码的形状
577 if len(seg_mask.shape) != 3 or seg_mask.shape[2] != 4:
578 raise ValueError(f"分割掩码形状应为 (H, W, 4),但得到 {seg_mask.shape}")
579
580 # 深度评估
581 depth_metrics = self._evaluate_depth(depth_map)
582 depth_score = self._score_depth(depth_metrics)
583
584 # 分割评估
585 seg_metrics = self._evaluate_segmentation(seg_mask)
586 seg_score = self._score_segmentation(seg_metrics)
587
588 # 综合评分 (深度40%,分割60%)
589 total_score = (depth_score * self.depth_weight +
590 seg_score * self.seg_weight)
591
592 # 判断是否合格
593 is_qualified = total_score >= self.threshold
594
595 return {
596 'is_qualified': is_qualified,
597 'score': round(total_score, 3),
598 'depth_score': round(depth_score, 3),
599 'seg_score': round(seg_score, 3),
600 'details': {
601 'depth': depth_metrics,
602 'segmentation': seg_metrics
603 }
604 }
605
606 def _evaluate_depth(self, depth_map: np.ndarray) -> Dict[str, float]:
607 """评估深度图特征"""
608 # 常见的无效深度标记值
609 INVALID_DEPTH_VALUES = [65504.0, 65535.0, 0.0]
610
611 # 过滤无效深度值
612 valid_mask = depth_map > 0
613 for invalid_val in INVALID_DEPTH_VALUES:
614 valid_mask = valid_mask & (np.abs(depth_map - invalid_val) > 1.0)
615
616 valid_depth = depth_map[valid_mask]
617
618 if len(valid_depth) == 0:
619 return {
620 'coverage': 0.0,
621 'range_mean_ratio': 0.0,
622 'std_mean_ratio': 0.0,
623 'entropy': 0.0,
624 'max_depth': 0.0,
625 'far_pixel_ratio': 0.0,
626 'max_median_ratio': 0.0,
627 'outlier_ratio': 0.0,
628 'valid_depth_ratio': 0.0,
629 'depth_concentration': 0.0
630 }
631
632 # 1. 有效深度覆盖率 - 真正有效的深度像素占比
633 valid_depth_ratio = len(valid_depth) / depth_map.size
634
635 # 2. 深度覆盖率(向后兼容)
636 coverage = valid_depth_ratio
637
638 # 3. 深度范围与均值的比值
639 depth_range = float(np.max(valid_depth) - np.min(valid_depth))
640 mean_depth = float(np.mean(valid_depth))
641 range_mean_ratio = depth_range / mean_depth if mean_depth > 0 else 0.0
642
643 # 4. 深度标准差与均值的比值
644 std_depth = float(np.std(valid_depth))
645 std_mean_ratio = std_depth / mean_depth if mean_depth > 0 else 0.0
646
647 # 5. 深度分布熵
648 hist, _ = np.histogram(valid_depth, bins=20)
649 hist = hist / hist.sum()
650 hist = hist[hist > 0]
651 entropy = -np.sum(hist * np.log(hist))
652
653 # 6. 最大深度值
654 max_depth = float(np.max(valid_depth))
655
656 # 7. 使用中位数检测远距离像素
657 median_depth = float(np.median(valid_depth))
658
659 # 远距离像素占比
660 far_threshold = median_depth * 5.0
661 far_pixels = valid_depth > far_threshold
662 far_pixel_ratio = float(np.sum(far_pixels) / len(valid_depth))
663
664 # 8. 最大深度/中位数比值
665 max_median_ratio = max_depth / median_depth if median_depth > 0 else 0.0
666
667 # 9. 离群值占比
668 percentile_75 = float(np.percentile(valid_depth, 75))
669 outlier_threshold = percentile_75 * 10.0
670 outliers = valid_depth > outlier_threshold
671 outlier_ratio = float(np.sum(outliers) / len(valid_depth))
672
673 # 10. 深度集中度 - 检测深度值是否过于单一(比如大部分是一面墙)
674 # 计算在中位数±15%范围内的像素占比
675 median_threshold_low = median_depth * 0.85
676 median_threshold_high = median_depth * 1.15
677 concentrated_pixels = (valid_depth >= median_threshold_low) & (valid_depth <= median_threshold_high)
678 depth_concentration = float(np.sum(concentrated_pixels) / len(valid_depth))
679
680 return {
681 'coverage': float(coverage),
682 'range_mean_ratio': float(range_mean_ratio),
683 'std_mean_ratio': float(std_mean_ratio),
684 'entropy': float(entropy),
685 'max_depth': float(max_depth),
686 'far_pixel_ratio': float(far_pixel_ratio),
687 'max_median_ratio': float(max_median_ratio),
688 'outlier_ratio': float(outlier_ratio),
689 'valid_depth_ratio': float(valid_depth_ratio),
690 'depth_concentration': float(depth_concentration)
691 }
692
693 def _evaluate_segmentation(self, seg_mask: np.ndarray) -> Dict[str, float]:
694 """
695 评估分割掩码特征
696
697 参数:
698 seg_mask: 分割掩码 (H, W, 4),RGBA格式
699 """
700 # 提取RGB通道(忽略alpha通道)
701 rgb_mask = seg_mask[:, :, :3]
702
703 # 重塑为(H*W, 3)以便处理
704 h, w = rgb_mask.shape[:2]
705 total_pixels = h * w
706 rgb_flat = rgb_mask.reshape(-1, 3)
707
708 # 使用字典统计每个颜色的像素数
709 color_counts = defaultdict(int)
710 for pixel in rgb_flat:
711 color_tuple = tuple(pixel)
712 color_counts[color_tuple] += 1
713
714 # 过滤背景颜色
715 if self.background_color in color_counts:
716 del color_counts[self.background_color]
717
718 # 获取唯一颜色(物体)数量
719 num_objects = len(color_counts)
720
721 if num_objects == 0:
722 return {
723 'num_objects': 0,
724 'num_small_objects': 0,
725 'max_coverage': 0.0,
726 'min_coverage': 0.0,
727 'avg_coverage': 0.0,
728 'has_large_object': False,
729 'color_distribution': {}
730 }
731
732 # 计算每个物体的覆盖率
733 coverages = []
734 small_object_threshold = 0.05 # 占比<5%的算小物体
735 large_object_threshold = 0.5 # 占比>50%的算大物体
736 num_small_objects = 0
737 has_large_object = False
738 color_distribution = {}
739
740 for color, count in color_counts.items():
741 coverage = count / total_pixels
742 coverages.append(coverage)
743
744 # 统计小物体数量
745 if coverage < small_object_threshold:
746 num_small_objects += 1
747
748 # 检测大物体
749 if coverage > large_object_threshold:
750 has_large_object = True
751
752 # 记录颜色分布(可选,用于调试)
753 color_str = f"RGB{color}"
754 color_distribution[color_str] = round(coverage, 4)
755
756 # 对覆盖率排序,便于查看分布
757 color_distribution = dict(sorted(color_distribution.items(),
758 key=lambda x: x[1], reverse=True))
759
760 return {
761 'num_objects': float(num_objects),
762 'num_small_objects': float(num_small_objects),
763 'max_coverage': float(max(coverages)) if coverages else 0.0,
764 'min_coverage': float(min(coverages)) if coverages else 0.0,
765 'avg_coverage': float(np.mean(coverages)) if coverages else 0.0,
766 'has_large_object': has_large_object, # 添加大物体标记
767 'color_distribution': color_distribution # 添加颜色分布信息
768 }
769
770 def _score_segmentation(self, metrics: Dict[str, float]) -> float:
771 """计算分割评分 (0-1) - 严格要求物体数量、小物体数量,并惩罚大面积物体"""
772 num_objects = metrics['num_objects']
773 num_small_objects = metrics['num_small_objects']
774 max_coverage = metrics['max_coverage']
775
776 # 硬性要求1:物体<6个直接不合格
777 if num_objects < 6:
778 return 0.0
779
780 # 硬性要求2:任何物体占比超过50%直接不合格
781 if max_coverage > 0.5:
782 return 0.0
783
784 score = 0.0
785
786 # 物体总数量评分 (12-20个部分得分,≥20个满分) - 权重30%
787 if num_objects >= 20:
788 score += 0.3
789 elif num_objects >= 12:
790 # 12-20个之间线性增长
791 score += ((num_objects - 12) / 8) * 0.3
792 else:
793 # 6-12个之间降低得分
794 score += ((num_objects - 6) / 6) * 0.15
795
796 # 小物体数量评分 (≥6个满分,<6个按比例) - 权重30%
797 if num_small_objects >= 6:
798 score += 0.3
799 else:
800 score += (num_small_objects / 6) * 0.3
801
802 # 最大物体占比评分 (理想范围10%-30%) - 权重20%
803 # 由于已经在50%处设置了硬性门槛,这里优化30%-50%之间的评分
804 if max_coverage <= 0.1:
805 # 太小也不理想(可能是分割过于碎片化)
806 score += max_coverage / 0.1 * 0.1
807 elif max_coverage <= 0.3:
808 # 10%-30%是理想范围
809 score += 0.2
810 else:
811 # 30%-50%之间线性下降
812 score += (0.5 - max_coverage) / 0.2 * 0.2
813
814 # 最小物体占比 (至少0.3%) - 权重10%
815 min_coverage = metrics['min_coverage']
816 if min_coverage >= 0.003:
817 score += 0.1
818 else:
819 score += min_coverage / 0.003 * 0.1
820
821 # 平均物体占比 (2-8%为理想,物体多所以占比要小) - 权重10%
822 avg_coverage = metrics['avg_coverage']
823 if 0.02 <= avg_coverage <= 0.08:
824 score += 0.1
825 elif avg_coverage < 0.02:
826 score += avg_coverage / 0.02 * 0.1
827 else:
828 score += max(0, (1 - (avg_coverage - 0.08) / 0.12)) * 0.1
829
830 return min(score, 1.0)
831
832 def _score_depth(self, metrics: Dict[str, float]) -> float:
833 """计算深度评分 (0-1) - 严格惩罚无效值和单一深度场景"""
834
835 # 严格检查有效深度比例 - 无效值>10%直接不合格
836 valid_depth_ratio = metrics['valid_depth_ratio']
837 if valid_depth_ratio < 0.9:
838 # 有效深度<90%(即无效值>10%),直接返回0分
839 return 0.0
840
841 score = 0.0
842
843 # 有效深度覆盖率评分 (>98%为好) - 权重15%
844 if valid_depth_ratio >= 0.98:
845 score += 0.15
846 else:
847 # 90-98%之间线性评分
848 score += ((valid_depth_ratio - 0.9) / 0.08) * 0.15
849
850 # 深度范围/均值比评分 (0.5-2.0为理想) - 权重10%
851 range_mean_ratio = metrics['range_mean_ratio']
852 if 0.5 <= range_mean_ratio <= 2.0:
853 score += 0.1
854 elif range_mean_ratio < 0.5:
855 score += range_mean_ratio / 0.5 * 0.1
856 else:
857 score += max(0, (1 - (range_mean_ratio - 2.0) / 3.0)) * 0.1
858
859 # 深度标准差/均值比评分 (0.2-0.6为理想) - 权重10%
860 std_mean_ratio = metrics['std_mean_ratio']
861 if 0.2 <= std_mean_ratio <= 0.6:
862 score += 0.1
863 elif std_mean_ratio < 0.2:
864 score += std_mean_ratio / 0.2 * 0.1
865 else:
866 score += max(0, (1 - (std_mean_ratio - 0.6) / 0.6)) * 0.1
867
868 # 深度分布熵评分 (越高越好) - 权重10%
869 entropy = metrics['entropy']
870 max_entropy = 3.0
871 score += min(entropy / max_entropy, 1.0) * 0.1
872
873 # 深度集中度惩罚 - 权重15%(检测单一深度场景如一面墙)
874 depth_concentration = metrics['depth_concentration']
875
876 # 如果标准差/熵都比较高,说明有物品,放宽集中度要求
877 has_variation = (std_mean_ratio >= 0.25) or (entropy >= 2.0)
878
879 if has_variation:
880 # 有足够的深度变化(墙前有物品),集中度要求宽松
881 if depth_concentration <= 0.6:
882 score += 0.15
883 elif depth_concentration <= 0.8:
884 score += (0.8 - depth_concentration) / 0.2 * 0.15
885 else:
886 score += 0.05 # 即使有变化,但集中度过高也要扣一些分
887 else:
888 # 深度变化不足,严格要求集中度
889 if depth_concentration <= 0.5:
890 score += 0.15
891 elif depth_concentration <= 0.7:
892 score += (0.7 - depth_concentration) / 0.2 * 0.15
893 else:
894 # 集中度>70%且无变化,严重扣分
895 score += 0.0
896
897 # 最大深度/中位数比值惩罚 - 权重20%
898 max_median_ratio = metrics['max_median_ratio']
899 if max_median_ratio <= 5.0:
900 score += 0.2
901 elif max_median_ratio <= 10.0:
902 score += (10.0 - max_median_ratio) / 5.0 * 0.2
903 else:
904 penalty = max(0, 1 - (max_median_ratio - 10.0) / 50.0)
905 score += penalty * 0.2
906
907 # 离群值占比惩罚 - 权重20%
908 outlier_ratio = metrics['outlier_ratio']
909 if outlier_ratio <= 0.01:
910 score += 0.2
911 elif outlier_ratio <= 0.05:
912 score += (0.05 - outlier_ratio) / 0.04 * 0.2
913 else:
914 penalty = max(0, 1 - (outlier_ratio - 0.05) / 0.15)
915 score += penalty * 0.2
916
917 return min(score, 1.0)
918
919
920
921
922@dataclass
923class CameraConfig:
924 """相机配置类,统一管理相机参数"""
925 width: float = 40.0
926 height: float = 40.0
927 depth: float = 40.0
928
929 @property
930 def size(self) -> List[float]:
931 return [self.width, self.height, self.depth]
932
933 @property
934 def half_extents(self) -> List[float]:
935 return [self.width/2, self.height/2, self.depth/2]
936
937
938# 全局默认相机配置
939DEFAULT_CAMERA = CameraConfig()
940
941
942def compute_aabb_from_vertices(vertices):
943 """
944 从顶点计算AABB(轴对齐包围盒)的中心和半长
945
946 Args:
947 vertices: (N, 3) array, 物体的顶点
948
949 Returns:
950 dict: {
951 'center': (3,) array,
952 'extent': (3,) array (半长),
953 'radius': float (包围球半径,用于快速排除)
954 }
955 """
956 min_point = vertices.min(axis=0)
957 max_point = vertices.max(axis=0)
958
959 center = (min_point + max_point) / 2
960 extent = (max_point - min_point) / 2
961
962 # 计算包围球半径(用于快速排除)
963 radius = np.linalg.norm(extent)
964
965 return {
966 'center': center,
967 'extent': extent,
968 'radius': radius,
969 'min': min_point,
970 'max': max_point
971 }
972
973
974def estimate_aabb_distance(aabb1_info, aabb2_info):
975 """
976 估算两个AABB之间的距离
977 使用包围球距离减去半径作为下界估计
978
979 Args:
980 aabb1_info, aabb2_info: AABB信息字典
981
982 Returns:
983 float: 估算的最小距离(可能为负表示重叠)
984 """
985 center_dist = np.linalg.norm(aabb2_info['center'] - aabb1_info['center'])
986 return center_dist - (aabb1_info['radius'] + aabb2_info['radius'])
987
988
989def create_camera_aabb_vertices(position, camera_config=None):
990 """
991 创建相机的AABB顶点
992
993 Args:
994 position: [x, y, z] 相机中心位置
995 camera_config: CameraConfig实例,None则使用默认配置
996
997 Returns:
998 (8, 3) array: 8个顶点坐标
999 """
1000 if camera_config is None:
1001 camera_config = DEFAULT_CAMERA
1002
1003 x, y, z = position
1004 w, h, d = camera_config.half_extents
1005
1006 # 创建8个顶点(AABB)
1007 vertices = np.array([
1008 [x - w, y - h, z - d], # 0: 底面左下
1009 [x + w, y - h, z - d], # 1: 底面右下
1010 [x + w, y + h, z - d], # 2: 底面右上
1011 [x - w, y + h, z - d], # 3: 底面左上
1012 [x - w, y - h, z + d], # 4: 顶面左下
1013 [x + w, y - h, z + d], # 5: 顶面右下
1014 [x + w, y + h, z + d], # 6: 顶面右上
1015 [x - w, y + h, z + d], # 7: 顶面左上
1016 ])
1017
1018 return vertices
1019
1020
1021def check_camera_collision(camera_position,
1022 object_vertices_list,
1023 camera_config=None,
1024 check_nearest=10,
1025 collision_threshold=0.0,
1026 use_improved_search=True):
1027 """
1028 检查相机位置是否与场景中的物体发生碰撞
1029
1030 Args:
1031 camera_position: [x, y, z] 相机位置
1032 object_vertices_list: list of (N, 3) arrays,场景中所有物体的顶点
1033 camera_config: CameraConfig实例,None则使用默认配置
1034 check_nearest: 检查最近的几个物体
1035 collision_threshold: IoU碰撞阈值,默认0(任何重叠都算碰撞)
1036 use_improved_search: 是否使用改进的搜索方法
1037
1038 Returns:
1039 dict: {
1040 'collision': bool,
1041 'colliding_indices': list,
1042 'collision_ious': list, # 每个碰撞的IoU值
1043 'nearest_indices': list,
1044 'nearest_distances': list,
1045 'checked_count': int
1046 }
1047 """
1048 if camera_config is None:
1049 camera_config = DEFAULT_CAMERA
1050
1051 # 创建相机AABB
1052 camera_center = np.array(camera_position)
1053 camera_extent = np.array(camera_config.half_extents)
1054
1055 # 预计算所有物体的AABB信息
1056 object_aabb_infos = [compute_aabb_from_vertices(verts)
1057 for verts in object_vertices_list]
1058
1059 if use_improved_search:
1060 # 改进的方法:使用包围球距离估算
1061 camera_aabb_info = {
1062 'center': camera_center,
1063 'extent': camera_extent,
1064 'radius': np.linalg.norm(camera_extent)
1065 }
1066
1067 distances = []
1068 for i, aabb_info in enumerate(object_aabb_infos):
1069 # 使用包围球距离作为估算
1070 dist = estimate_aabb_distance(camera_aabb_info, aabb_info)
1071 distances.append((dist, i))
1072
1073 # 按距离排序
1074 distances.sort(key=lambda x: x[0])
1075
1076 # 选择最近的物体进行精确检查
1077 indices_to_check = [idx for _, idx in distances[:check_nearest]]
1078 nearest_distances = [dist for dist, _ in distances[:check_nearest]]
1079 else:
1080 # 原始方法:使用中心点距离
1081 object_centers = np.array([info['center'] for info in object_aabb_infos])
1082 kdtree = cKDTree(object_centers)
1083 center_distances, indices = kdtree.query(camera_position,
1084 k=min(check_nearest, len(object_centers)))
1085
1086 if not isinstance(center_distances, np.ndarray):
1087 center_distances = np.array([center_distances])
1088 indices = np.array([indices])
1089
1090 indices_to_check = indices
1091 nearest_distances = center_distances.tolist()
1092
1093 # 检查碰撞
1094 colliding_indices = []
1095 collision_ious = []
1096 checked_count = 0
1097
1098 for idx in indices_to_check:
1099 if not 0 <= idx < len(object_aabb_infos):
1100 continue
1101
1102 checked_count += 1
1103
1104 # 使用新的collide函数检查碰撞
1105 obj_info = object_aabb_infos[idx]
1106
1107 # 计算IoU用于记录
1108 iou = compute_iou(camera_center, camera_extent,
1109 obj_info['center'], obj_info['extent'])
1110
1111 if collide(camera_center, camera_extent,
1112 obj_info['center'], obj_info['extent'],
1113 collision_threshold):
1114 colliding_indices.append(int(idx))
1115 collision_ious.append(float(iou))
1116
1117 return {
1118 'collision': len(colliding_indices) > 0,
1119 'colliding_indices': colliding_indices,
1120 'collision_ious': collision_ious,
1121 'nearest_indices': [int(idx) for idx in indices_to_check],
1122 'nearest_distances': nearest_distances,
1123 'checked_count': checked_count
1124 }
1125
1126
1127def compute_iou(center1, extent1, center2, extent2):
1128 """
1129 计算两个AABB的IoU值
1130
1131 Args:
1132 center1, extent1: 第一个AABB的中心和半长
1133 center2, extent2: 第二个AABB的中心和半长
1134
1135 Returns:
1136 float: IoU值(0到1之间)
1137 """
1138 center1, extent1 = np.array(center1), np.array(extent1)
1139 center2, extent2 = np.array(center2), np.array(extent2)
1140
1141 min1, max1 = center1 - extent1, center1 + extent1
1142 min2, max2 = center2 - extent2, center2 + extent2
1143
1144 inter_min = np.maximum(min1, min2)
1145 inter_max = np.minimum(max1, max2)
1146 inter_extent = np.maximum(0.0, inter_max - inter_min)
1147 inter_vol = np.prod(inter_extent)
1148
1149 vol1, vol2 = np.prod(2 * extent1), np.prod(2 * extent2)
1150 union_vol = vol1 + vol2 - inter_vol
1151
1152 iou = inter_vol / union_vol if union_vol > 0 else 0.0
1153 return iou
1154
1155
1156def compute_scene_bounds(object_vertices_list,
1157 margin=30,
1158 trim_percent=10,
1159 camera_config=None):
1160 """
1161 计算包含所有物体的边界框,去掉极值
1162
1163 Args:
1164 object_vertices_list: list of (N, 3) arrays
1165 margin: 边界内缩距离(cm)
1166 trim_percent: 去掉的极值百分比(0-50)
1167 camera_config: CameraConfig实例,用于确保边界足够大
1168
1169 Returns:
1170 dict: 边界信息
1171 """
1172 if camera_config is None:
1173 camera_config = DEFAULT_CAMERA
1174
1175 # 收集所有顶点
1176 all_vertices = np.vstack(object_vertices_list)
1177 total_vertices = len(all_vertices)
1178
1179 # 计算要修剪的百分位数
1180 lower_percentile = trim_percent
1181 upper_percentile = 100 - trim_percent
1182
1183 # 对每个轴分别计算修剪后的范围
1184 x_min = np.percentile(all_vertices[:, 0], lower_percentile)
1185 x_max = np.percentile(all_vertices[:, 0], upper_percentile)
1186 y_min = np.percentile(all_vertices[:, 1], lower_percentile)
1187 y_max = np.percentile(all_vertices[:, 1], upper_percentile)
1188 z_min = np.percentile(all_vertices[:, 2], lower_percentile)
1189 z_max = np.percentile(all_vertices[:, 2], upper_percentile)
1190
1191 # 确保边界至少能容纳相机
1192 min_width = camera_config.width + 2 * margin
1193 min_height = camera_config.height + 2 * margin
1194 min_depth = camera_config.depth + 2 * margin
1195
1196 # 应用边界内缩
1197 x_min += margin
1198 x_max -= margin
1199 y_min += margin
1200 y_max -= margin
