好的,没有问题。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, widthmax_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 是保存所有层输出通道的列表。ffrom 索引。如果 f = -1ch[-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 实例,而不是 nC2f 实例的序列。
        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)
  • 如果 mConcat (拼接):
    • 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: (其他所有情况):
    • 适用于不改变通道数的模块 (如 MaxPool2dnn.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 列表中保存哪些中间层的输出。
Logo

助力合肥开发者学习交流的技术社区,不定期举办线上线下活动,欢迎大家的加入

更多推荐