如何在build_targets函数中调用grids方法的gridsize与anchor_vec变量
调用方法
你已经在grids()方法中将gridsize通过register_buffer注册为模型的缓存变量,anchor_vec也绑定为了模型实例的属性,且build_targets()的第一个入参就是模型对象,直接从模型实例读取即可:
- 首先保证调用
build_targets()之前,已经执行过模型的grids()方法完成初始化,避免属性不存在报错 - 修改
build_targets()的取值逻辑即可,修改后代码如下:
def build_targets(model, targets): # 可提前加校验逻辑避免未初始化:assert hasattr(model, 'gridsize') and hasattr(model, 'anchor_vec'), "请先调用model.grids()完成初始化" for i in layers_list: ng, anchor_vec = model.gridsize, model.anchor_vec # 直接从传入的model对象读取属性
如果你的场景是多尺度检测(layers_list对应不同下采样倍率的检测层),则可以将grids()改为每个检测层的成员方法,每个层独立存储自己的gridsize和anchor_vec,遍历的时候取对应层的属性即可:
def build_targets(model, targets): for layer in model.detect_layers: # 遍历所有检测层 ng, anchor_vec = layer.gridsize, layer.anchor_vec
备选方案:若不想将参数绑定到模型实例,可修改
grids()方法返回两个参数,调用时传递给build_targets()即可:def grids(self, img_size=(608,608), gridsize=(19, 19), device='cpu', type=torch.float32): # 原有逻辑不变 return self.gridsize, self.anchor_vec # 调用侧代码 gridsize, anchor_vec = model.grids() build_targets(model, targets, gridsize, anchor_vec)
内容的提问来源于stack exchange,提问作者Heartbeat
相关产品推荐
相关产品推荐

