如何修改代码为图像不同实例设置差异化深度四叉树
实现不同实例差异化四叉树划分的代码修改方案
核心逻辑是给每个实例预设最大允许划分深度,节点分裂前先判断当前深度是否已经触达节点内所有实例允许的最大深度上限,达到上限就不再分裂保留粗粒度块,没到上限才继续分裂得到细粒度块。
第一步:定义实例对应的深度配置
直接按需求给不同实例配置深度值,数值越大划分越精细:
- 精细划分实例(泰迪熊、植物):配置最大深度为8~10,数值越大块越小,可按需调整
- 粗糙划分实例(背景、桌子、花瓶、椅子):配置最大深度为3~4即可
在代码开头添加配置字典:
# 实例像素值对应最大划分深度配置,值越大划分越精细 INSTANCE_MAX_DEPTH = { 0.0: 3, # background 粗粒度 1.0: 3, # chair 粗粒度 2.0: 3, # table 粗粒度 3.0: 9, # teddy bear 细粒度 4.0: 9, # plant 细粒度 5.0: 3 # vase 粗粒度 } # 全局默认最大深度兜底 GLOBAL_MAX_DEPTH = 10
第二步:修改QuadTree类的初始化和分裂逻辑
原有类没有记录当前节点深度,也没有深度判定逻辑,需要调整3处:
- 初始化方法增加
depth参数,记录当前节点所处的深度层级,根节点深度从0开始 - 分裂生成子节点时,子节点深度为当前节点深度+1,同时把实例深度配置透传给子节点
- 插入点触发分裂前,先判断当前节点深度是否已经小于节点内所有实例要求的最大深度,若当前深度已经达到所有实例的深度上限,就不再分裂直接存点
注意:Point类和Rectangle类完全不需要修改,保持原有代码即可
修改后的QuadTree类完整代码:
class QuadTree: def __init__(self, boundary, capacity = 1, depth=0, instance_depth_config=None): self.boundary = boundary self.capacity = capacity self.depth = depth # 透传实例深度配置,根节点未传参时使用默认配置 self.instance_depth_config = instance_depth_config if instance_depth_config is not None else INSTANCE_MAX_DEPTH self.points = [] self.instances = [] self.unique_instances = [] self.divided = False def insert(self, point, instance): if not self.boundary.containsPoint(point): return False if self.divided: if self.nw.insert(point, instance): return True elif self.ne.insert(point, instance): return True elif self.sw.insert(point, instance): return True elif self.se.insert(point, instance): return True else: assert "Should never happen" else: instance_already_seen = instance in self.unique_instances # 分裂判定:当前深度未达到节点内所有实例的最大允许深度、且唯一实例数超容量时才分裂 can_divide = True # 即将插入的实例也纳入深度判定范围 all_instances_in_node = self.unique_instances + [instance] for inst in all_instances_in_node: inst_max_depth = self.instance_depth_config.get(inst, GLOBAL_MAX_DEPTH) if self.depth >= inst_max_depth: can_divide = False break if instance_already_seen or len(self.unique_instances) < self.capacity or not can_divide: self.points.append(point) self.instances.append(instance) if not instance_already_seen: self.unique_instances.append(instance) return True self.divide() assert self.insert(point, instance) return True def queryRange(self, range): found_points = [] if not self.boundary.intersects(range): return [] if self.divided: found_points.extend(self.nw.queryRange(range)) found_points.extend(self.ne.queryRange(range)) found_points.extend(self.sw.queryRange(range)) found_points.extend(self.se.queryRange(range)) else: for point in self.points: if range.containsPoint(point): found_points.append(point) return found_points def queryRadius(self, range, center): if not self.boundary.intersects(range): return [] found_points = [] if self.divided: found_points.extend(self.nw.queryRadius(range, center)) found_points.extend(self.ne.queryRadius(range, center)) found_points.extend(self.sw.queryRadius(range, center)) found_points.extend(self.se.queryRadius(range, center)) else: for point in self.points: if range.containsPoint(point) and point.distanceToCenter(center) <= range.width: found_points.append(point) return found_points def divide(self): center_x = self.boundary.center.x center_y = self.boundary.center.y new_width = self.boundary.width / 2 new_height = self.boundary.height / 2 # 子节点深度+1,透传深度配置 nw = Rectangle(Point(center_x - new_width, center_y - new_height), new_width, new_height) self.nw = QuadTree(nw, depth=self.depth+1, instance_depth_config=self.instance_depth_config) ne = Rectangle(Point(center_x + new_width, center_y - new_height), new_width, new_height) self.ne = QuadTree(ne, depth=self.depth+1, instance_depth_config=self.instance_depth_config) sw = Rectangle(Point(center_x - new_width, center_y + new_height), new_width, new_height) self.sw = QuadTree(sw, depth=self.depth+1, instance_depth_config=self.instance_depth_config) se = Rectangle(Point(center_x + new_width, center_y + new_height), new_width, new_height) self.se = QuadTree(se, depth=self.depth+1, instance_depth_config=self.instance_depth_config) self.divided = True self.unique_instances = [] for (point, instance) in zip(self.points, self.instances): assert self.insert(point, instance) self.points = [] self.instances = [] def __len__(self): if self.divided: return len(self.nw) + len(self.ne) + len(self.sw) + len(self.se) else: return len(self.points) def draw(self, ax): if self.divided: self.nw.draw(ax) self.ne.draw(ax) self.se.draw(ax) self.sw.draw(ax) else: self.boundary.draw(ax)
第三步:运行说明
原有运行脚本不需要修改逻辑,初始化QuadTree时不用传额外参数,根节点默认深度为0,会自动加载配置的实例深度规则,直接运行即可得到目标效果:
- 泰迪熊、植物区域会分裂到更深层级,网格块更小更精细
- 背景、桌子、花瓶、椅子区域到达配置的浅深度就停止分裂,网格块更大更粗糙
如果粒度不符合预期,直接调整INSTANCE_MAX_DEPTH里的数值即可:想让某个实例划分更细就调大对应值,想更粗就调小对应值。跨实例的边界混合块会自动按块内要求最精细的实例深度分裂,不会出现精细实例被粗块截断的问题。
内容的提问来源于stack exchange,提问作者ashah
相关产品推荐
相关产品推荐

