Team Ai
Datasetpublic

TaxonomyProject/SimulationImage

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes1.5kdownloads
data_collection_helpers.py1526 linesDownload Raw Back to code
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

Showing the first 1,200 of 1526 lines. Download the file for the rest.

TaxonomyProject/SimulationImage · Team Ai