tasks.py里的parse_model
·
好的,没有问题。parse_model 是整个框架的基石,逐行理解它至关重要。
这是对 parse_model 函数的逐行详细解释:
def parse_model(d, ch, verbose=True):
"""
Parse a YOLO model.yaml dictionary into a PyTorch model.
Args:
d (dict): Model dictionary (加载自 YAML 文件).
ch (int): Input channels (输入通道数,通常为 3).
verbose (bool): Whether to print model details (是否打印模型构建的详细信息).
Returns:
model (torch.nn.Sequential): PyTorch model (构建好的 PyTorch 模型).
save (list): Sorted list of output layers (需要保存输出的层的索引列表).
"""
import ast # 导入“抽象语法树”模块
import ast: 导入ast(Abstract Syntax Trees) 库。这个库的唯一目的是为了安全地将 YAML 中的字符串参数(如'True','None','[1, 2]')转换为实际的 Python 对象(如True,None,[1, 2]),它比eval()安全得多。
1. 参数初始化 (Argument Initialization)
# Args
legacy = True # backward compatibility for v3/v5/v8/v9 models
max_channels = float("inf")
nc, act, scales = (d.get(x) for x in ("nc", "activation", "scales"))
depth, width, kpt_shape = (d.get(x, 1.0) for x in ("depth_multiple", "width_multiple", "kpt_shape"))
scale = d.get("scale")
legacy = True: 设置一个标志,用于处理旧版模型(v3/v5/v8/v9)的兼容性。如果解析器在 YAML 中遇到像C3k2这样的新模块,它会把这个标志设为False。max_channels = float("inf"): 初始化最大通道数。这主要用于 YOLOv9 等模型,它们通过scales字典来限制模型的最大宽度。nc, act, scales = ...: 这是一个生成器表达式,它从字典d(来自YAML) 中获取nc(类别数)、activation(激活函数) 和scales(缩放字典)。如果 YAML 中没有定义,这些值将是None。depth, width, kpt_shape = ...: 同样,获取depth_multiple(深度倍数) 和width_multiple(宽度倍数)。这是区分 n/s/m/l/x 模型的关键。- 例如,
yolov8n.yaml会设置depth_multiple: 0.33,width_multiple: 0.25。 kpt_shape用于姿态估计模型。.get(x, 1.0)确保如果 YAML 中未定义这些值,它们默认为1.0。
- 例如,
scale = d.get("scale"): 获取模型规模,如'n','s','x'。这个值是由yaml_model_load函数在调用parse_model之前 猜测 并注入到字典d中的。
if scales:
if not scale:
scale = next(iter(scales.keys()))
LOGGER.warning(f"no model scale passed. Assuming scale='{scale}'.")
depth, width, max_channels = scales[scale]
if scales:: 检查 YAML 是否使用了新的scales字典格式(如 YOLOv9)。if not scale:: 一个安全检查。如果scales字典存在,但scale没被设置,就自动选择scales字典中的第一个可用规模。depth, width, max_channels = scales[scale]: 核心。如果使用了scales字典,就用scale(如'n') 作为键,从中提取出特定规模的depth,width和max_channels,覆盖掉之前从d.get(...)获取的默认值。
2. 设置激活函数和日志
if act:
Conv.default_act = eval(act) # redefine default activation, i.e. Conv.default_act = torch.nn.SiLU()
if verbose:
LOGGER.info(f"{colorstr('activation:')} {act}") # print
if verbose:
LOGGER.info(f"\n{'':>3}{'from':>20}{'n':>3}{'params':>10} {'module':<45}{'arguments':<30}")
if act:: 检查 YAML 中是否指定了全局激活函数(如activation: 'torch.nn.SiLU')。Conv.default_act = eval(act): 非常重要。它会全局修改Conv模块(定义在ultralytics/nn/modules/conv.py中)的默认激活函数。eval()将字符串'torch.nn.SiLU'转换成实际的类torch.nn.SiLU。if verbose:: 如果verbose=True,就打印出模型结构表的表头(索引、来源、重复次数、参数量、模块名、参数)。
3. 循环变量初始化
ch = [ch]
layers, save, c2 = [], [], ch[-1] # layers, savelist, ch out
ch = [ch]: 这是一个巧妙的技巧。ch(初始通道,如 3) 被转换成一个列表[3]。在循环中,这个ch列表将保存每一层(每个索引i)的输出通道数。layers, save, c2 = [], [], ch[-1]:layers = []: 一个空列表,用于存放所有被实例化的 PyTorch 模块(如Conv,C2f)。save = []: 一个空列表,用于记录那些需要被后续层(如Concat,Detect)引用的层的索引。c2 = ch[-1]:c2是一个临时变量,代表“当前层的输出通道数”。它被初始化为列表ch中的最后一个值(目前只有输入的 3 通道)。
4. 定义模块集合 (为了快速查找)
base_modules = frozenset( ... )
repeat_modules = frozenset( ... )
frozenset: 这是一个不可变的、可哈希的集合。在这里使用frozenset是为了极快地进行成员资格检查(即if m in base_modules:)。base_modules: 包含所有“基础”模块,如Conv,C2f,SPPF。这些模块的共同点是它们都接受c1(in_channels) 和c2(out_channels) 作为它们的前两个参数。repeat_modules: 包含所有内部有重复结构的模块,如C2f,C3,BottleneckCSP。这些模块除了c1,c2之外,还需要一个n(重复次数) 参数。
5. 主循环 (The Main Loop)
for i, (f, n, m, args) in enumerate(d["backbone"] + d["head"]): # from, number, module, args
d["backbone"] + d["head"]: 这是 YAML 文件的核心。它将backbone列表和head列表拼接成一个长列表。enumerate(...): 遍历这个长列表,i是当前层的索引 (从 0 开始),(f, n, m, args)是从 YAML 中解包出来的四个值:f(from): 层的输入来源。-1表示上一层;[-1, 6]表示拼接上一层和第 6 层的输出。n(number): 模块的重复次数(受depth倍数影响)。m(module): 模块的字符串名称 (如'C2f')。args: 传递给模块的参数列表 (如[64, True])。
6. 循环内部:逐行解析
6.1. 解析模块名称 (Module Resolution)
m = (
getattr(torch.nn, m[3:])
if "nn." in m
else getattr(__import__("torchvision").ops, m[16:])
if "torchvision.ops." in m
else globals()[m]
) # get module
- 这行代码将字符串
m转换成一个实际的 Python 类。 if "nn." in m: 如果m是'nn.Conv2d',它会从torch.nn中获取Conv2d类。if "torchvision.ops." in m: 如果m是'torchvision.ops.DeformConv2d',它会导入torchvision.ops并获取该类。else globals()[m]: 这是最常用的情况。globals()返回一个包含当前文件中所有全局变量(包括导入的模块)的字典。当m是'C2f'时,globals()['C2f']会返回你在文件顶部导入的C2f类。
6.2. 解析参数 (Argument Parsing)
for j, a in enumerate(args):
if isinstance(a, str):
with contextlib.suppress(ValueError):
args[j] = locals()[a] if a in locals() else ast.literal_eval(a)
- 这个循环遍历
args列表(如[64, 'True', 'nc'])来处理字符串。 locals()[a] if a in locals(): 另一个巧妙的技巧。如果参数a是一个字符串(如'nc'),它会检查locals()(本地变量字典) 中是否存在nc。因为nc是在函数顶部定义的,所以这个检查为真,它会将字符串'nc'替换为变量nc的实际值(如80)。ast.literal_eval(a): 如果a不是本地变量(如'True'或'[1, 2]'),ast.literal_eval会安全地将其转换为True(布尔值) 或[1, 2](列表)。
6.3. 计算深度 (Depth Gain)
n = n_ = max(round(n * depth), 1) if n > 1 else n # depth gain
depth是我们之前获取的depth_multiple(如0.33)。if n > 1: 关键检查。只对那些n > 1的层(即C2f,C3等重复模块)应用深度缩放。n=1的层(如单个Conv)不受影响。max(round(n * depth), 1):n * depth: (例如3 * 0.33 = 0.99)。round(...): 四舍五入 (round(0.99) = 1)。max(..., 1): 确保重复次数至少为 1。
n = n_ = ...:n_保存这个计算出的新深度(用于日志打印),n稍后可能会被修改。
6.4. 核心逻辑:计算通道数 (Channel Calculation)
这是函数中最复杂的部分,它根据模块类型 m 来决定如何计算 c1 (输入通道) 和 c2 (输出通道)。
if m in base_modules:
c1, c2 = ch[f], args[0]
if c2 != nc: # if c2 not equal to number of classes (i.e. for Classify() output)
c2 = make_divisible(min(c2, max_channels) * width, 8)
args = [c1, c2, *args[1:]]
if m in repeat_modules:
args.insert(2, n) # number of repeats
n = 1
- 如果
m是一个“基础模块” (如Conv,C2f,SPPF):c1, c2 = ch[f], args[0]:c1 = ch[f]: 获取输入通道。ch是保存所有层输出通道的列表。f是from索引。如果f = -1,ch[-1]就会获取上一层的输出通道数,作为这一层的输入。c2 = args[0]: 获取基础输出通道,即args列表的第一个元素(如64)。
if c2 != nc:: 检查c2是不是nc(类别数)。如果是nc(例如Detect头的Conv层),我们不希望缩放它。c2 = make_divisible(min(c2, max_channels) * width, 8): 宽度缩放。c2 * width: 将基础通道数乘以width_multiple(如64 * 0.25 = 16)。min(..., max_channels): 确保不超过max_channels限制。make_divisible(..., 8): 将结果调整为 8 的倍数,这在 GPU 上计算效率更高。
args = [c1, c2, *args[1:]]: 重建参数列表。现在args变为[c1, c2, ...](如[16, 32, k=3, s=1, ...]),这正是Conv等模块构造函数所期望的。if m in repeat_modules:: 如果m还是一个“重复模块” (如C2f):args.insert(2, n): 将计算好的深度n(如1)插入到args列表的第 2 个位置,变为[c1, c2, n, ...]。C2f模块会用这个n来决定构建多少个内部的Bottleneck。n = 1: 将外部循环的n重置为 1。这是因为C2f模块自己处理重复,我们只需要创建一个C2f实例,而不是n个C2f实例的序列。
elif m is AIFI:
args = [ch[f], *args]
- 这是一个特定模块 (
AIFI) 的处理,它需要c1作为第一个参数。
elif m in frozenset({HGStem, HGBlock}):
# ... (特定于 HGNet (YOLOv9) 的复杂通道逻辑)
- 这是 YOLOv9 (HGNet) 独有的模块,有更复杂的通道计算。
elif m is Concat:
c2 = sum(ch[x] for x in f)
- 如果
m是Concat(拼接):f此时是一个列表,如[-1, 6]。sum(ch[x] for x in f): 它遍历f列表中的所有索引,从ch列表中查找到它们各自的输出通道数,然后求和。c2(输出通道数) 就是所有输入层通道数的总和。
elif m in frozenset(
{Detect, WorldDetect, YOLOEDetect, Segment, YOLOESegment, Pose, OBB, ImagePoolingAttn, v10Detect}
):
args.append([ch[x] for x in f])
# ... (处理 Detect/Segment 的特定参数, 如 proto-channels)
- 如果
m是一个“检测头”模块:f也是一个列表,如[15, 18, 21](代表 P3, P4, P5 特征图所在的层索引)。args.append([ch[x] for x in f]): 它会查找 P3, P4, P5 层的输出通道数 (如[256, 512, 1024]),并将这个通道列表追加到args中。Detect模块需要这个列表来构建其内部的卷积层。
# ... (其他 elif 块处理 RTDETRDecoder, CBLinear 等特殊模块) ...
else:
c2 = ch[f]
else:(其他所有情况):- 适用于不改变通道数的模块 (如
MaxPool2d或nn.Identity)。 c2 = ch[f]: 输出通道数c2直接等于输入通道数 (ch[f])。
- 适用于不改变通道数的模块 (如
6.5. 实例化模块 (Module Instantiation)
m_ = torch.nn.Sequential(*(m(*args) for _ in range(n))) if n > 1 else m(*args) # module
- 这是最终创建模块的地方。
m(*args): 使用我们精心构建的args列表来调用模块的构造函数(如C2f(c1, c2, n, True))。if n > 1: 这个n是外部循环的n。对于repeat_modules,它已经被重置为 1,所以会走else分支。这个if n > 1是为那些不在repeat_modules中、但又需要在 YAML 中被连续重复n次的模块准备的。m_:m_是新创建的nn.Module实例。
6.6. 附加信息并记录 (Attach Info & Log)
t = str(m)[8:-2].replace("__main__.", "") # module type
m_.np = sum(x.numel() for x in m_.parameters()) # number params
m_.i, m_.f, m_.type = i, f, t # attach index, 'from' index, type
if verbose:
LOGGER.info(f"{i:>3}{f!s:>20}{n_:>3}{m_.np:10.0f} {t:<45}{args!s:<30}") # print
t = ...: 获取模块的简洁类型字符串(用于日志)。m_.np = ...: 计算这个模块的总参数量 (np),并将其作为属性存储在模块m_上。m_.i, m_.f, m_.type = i, f, t: 至关重要! 它将当前层的索引i、来源f和类型t作为属性附加到模块m_上。BaseModel中的_predict_once方法会读取这些属性来决定如何执行前向传播。if verbose:: 打印这一行的所有信息(索引、来源、重复次数n_、参数量、模块名、参数)。
6.7. 更新状态 (Update State)
save.extend(x % i for x in ([f] if isinstance(f, int) else f) if x != -1) # append to savelist
layers.append(m_)
if i == 0:
ch = []
ch.append(c2)
save.extend(...):([f] if isinstance(f, int) else f): 确保f是一个列表(即使它只是单个int如-1)。if x != -1: 过滤掉-1(上一层),因为我们不需要显式保存上一层。save.extend(...): 将所有非-1的来源索引f添加到save列表中。
layers.append(m_): 将新创建的模块m_添加到layers列表中。if i == 0: ch = []: 在处理完第 0 层后,清空ch列表(它之前只包含[3])。ch.append(c2): 核心状态更新。将当前层i的输出通道数c2添加到ch列表的末尾。现在ch[i]就保存了第i层的输出通道数,供后续层在ch[f]中查询。
7. 返回结果 (Return)
return torch.nn.Sequential(*layers), sorted(save)
- 循环结束后:
torch.nn.Sequential(*layers): 使用*操作符将layers列表(包含所有m_模块)解包,并将它们传入nn.Sequential的构造函数。这就创建了最终的、完整的 PyTorch 模型。sorted(save): 返回一个排序后的save列表,DetectionModel会将其存储为self.save,供_predict_once方法使用,以决定在y列表中保存哪些中间层的输出。
更多推荐


所有评论(0)