三条压缩边界:stem、深度、分辨率

Posted by Closure on August 15, 2026

做消融做的比较少,遇到了很多遇不到的问题….其实是api/网页端会员同时到期了当时又懒得续(已经续上了),所以这篇代码和古法cot含量很高

因为是3D nnUNet加上一次训练上千epoch,消融和主实验抢同一个槽位,我每次实验都是一次完整的1000epoch 3D fullres训练+备份+三重验证,在这种预算下可想而知….

而且-p参数搞混覆盖过别人的实验/旧版会清空输出目录/plans被污染到要求codex写修复脚本,大量低智力操作发生中(

变量还控不干净,14层15层还差了参数量和一层容量,Dice掉了只能归因于stem特殊,但是严格说这是stem特殊or多一层压缩ANYWHERE都会掉,要做参数对齐的对照实验数直接翻倍,烧不起

详见这篇,是在脑出血CT分割任务上验证SignedGrid,每层C_out个自由核→P 个正核+N个负核,物理核数只需P+N,输出通道数P×N不变,物理动机来自光子硬件。

消融1

A Prompt-driven Universal Segmentation Model as well as A Strong Representation Learner

在哪可以复制粘贴抄到呢

在做这两个消融之前

  • 基线 ResEnc-L 140.8M 1×
  • 宽度压缩(65 层) 46.7M 3.02×
  • 深度压缩(14 层,直接训练) 20.19M 6.97×
  • 深度压缩(14 层,蒸馏) 20.6M 6.84×

宽度压缩叠加深度压缩后参数量从基线的约1/3压到不足1/7,精度与基线在统计上无差异,14层直接训练是当时最优的一组。

14层这个最优版本里有两个从未被单独验证过的设计选择:

  1. 跳过了 stem 层。”跳过 stem 是对的”从未被测过,它只是从上游配置继承下来的默认值(shallow_7layer.yaml 中 skip_first: true)。
  2. 分辨率层级数始终是 7。此前所有压缩动的都是每层多宽和每层几个block,从未动过网络有几个分辨率层级。

所以两个消融分别检验这两个假设,消融一检验stem是否该被跳过,消融二检验分辨率层级数能否削减。

为什么stem值得单独对待

要用一个模型接收多种模态/多种任务的输入,就必须专门处理网络最前端如何接收原始输入通道这件事,不同任务的输入通道数和强度分布各不相同,而网络更深的层处理的都已经是被 stem 编码过的特征。UniSeg 这类通用分割模型的架构实践(首 stage 保持全分辨率、第一层卷积专门负责接收多模态原始输入通道)为网络第一层与内部层在功能上属于两类不同的东西提供了先验支持。

在nnUNet的ResEnc-L架构中,stem有两个内部层不具备的特殊性:

  1. 它直接接收原始输入通道。本任务的输入是 2 通道(CT + noNorm),即未经任何抽象的原始 CT 信号;而内部层接收的是已被前面层编码过的 32、64、128…… 维抽象特征。原始信号与抽象特征在信息密度和数值分布上完全不同。
  2. 它是唯一保持全分辨率的层。stem 的 stride 为 [1,1,1],不做下采样;内部各 stage 的第一个卷积几乎都带 stride-2 下采样,3D 各向异性下部分 stage 在 z 方向不下采样,但平面内仍下采样。

因此消融一的假设是既然stem处理的是未经抽象的原始信号,且工作在全分辨率上,它与内部层性质不同,那么把适用于内部层的正负核压缩约束同样施加于它很可能并不合适。如果没有UniSeg这一先验,消融一就只是一次把被跳过的那层也加回来试试尝试。

实验设计

两个config的唯一区别就是一个开关:

# shallow_7layer.yaml(14 层,跳过 stem)
skip_first: true

# shallow_15layer_with_stem.yaml(15 层,替换 stem)
skip_first: false

递归遍历时用路径判断子模块是否属于encoder,并用 id() 去重以避免重复替换 nnUNet 中被多处引用的同一个卷积:

if _is_eligible_conv(child, config):
    if in_encoder and config.skip_first and not state["encoder_first_done"]:
        state["encoder_first_done"] = True      # 遇到 stem,标记并跳过
        stats["skipped"] += 1
        skipped_original_by_id[child_id] = child
        continue
    replacement = _replace_module(parent, child_name, child, config)  # 否则替换为 SignedGridConvNd
    replacement_by_original_id[child_id] = (child, replacement)
    stats["replaced"] += 1
    if in_encoder and not state["encoder_first_done"]:
        state["encoder_first_done"] = True      # skip_first=False 时,stem 被替换后同样标记
    continue

除该开关外两组实验的架构plans/数据划分/训练超参数完全一致

结果

15层 Dice 0.6674,相比14层的0.7126是下降的

所以越靠近原始输入的层越特殊,它需要在全分辨率上保留低层原始特征,对表达自由度的要求更高,因此不适合施加与内部层相同的激进压缩约束。

14层版本跳过stem这一继承来的默认值

消融2 STU-Net

STU-Net: Scalable and Transferable Medical Image Segmentation Models Empowered by Large-Scale Supervised Pre-training

STU-Net研究如何把nnUNet架构放大/缩小到不同规模,并做大规模预训练/怎么系统地调整网络规模,但是这个过程中它遇到nnUNet的网络结构不是随便设的,很多超参是由数据自动推导的。

它和消融二的关系是改stage数是牵一发动全身,nnUNet的stage数是由patch size和spacing 自动推导出来的,不是可以随便乱设的独立超参。

↑,的原理是网络要能把输入图像一路下采样到瓶颈,再上采样回来,每一次下采样都要求当前的特征图尺寸能被stride整除。stage数/每个stage的stride\patch size三者之间存在一个自洽约束,它们必须互相匹配,否则要么某个维度下采样到0,要么特征图尺寸对不上。

所以我要从7个stage削减到6 5个,等于主动打破了nnUNet自动推导出来的stage 数 ↔ patch/spacing平衡,如果只改 stage 数不同步改其他联动参数,网络就会建不起来或下采样冲突。

我的数据patch_size 是 [40, 320, 320],其中 z 方向只有 40,本来能做的下采样次数就有限(z 方向每下采样一次减半,40→20→10→5,最多几次),削减 stage 直接影响这个下采样链条,风险最高的就是z方向。

消融二的设计里加了两条防护:

  • 削减 stage 必须协调修改一整组联动数组,而不是只改一个数字。
  • 需要同步裁剪的字段(都在 plans 的 arch_kwargs 下)
  1. n_stages(层级总数)
  2. features_per_stage(每层通道数,长度 = n_stages)
  3. kernel_sizes(每层卷积核,长度 = n_stages)
  4. strides(每层下采样步长,长度 = n_stages)
  5. n_blocks_per_stage(每层 block 数,长度 = n_stages)
  6. n_conv_per_stage_decoder(decoder,长度 = n_stages − 1)

训练前必须做构建+一次前向的校验

因为即使各数组长度对上了,下采样倍数和 patch 尺寸是否冲突,手算算不全(手算只能验证”能不能整除”,验证不了卷积对齐、decoder 上采样等深层约束)。唯一可靠的方式是让网络真的构建一次、跑一次前向。

第一次发现5stage的问题其实不是stage/patch冲突,是显存被lab的同学占了(删了)(?

按STU-Net的提示用脚本统一裁剪所有联动字段,每步打印旧值→新值

# 长度 == n_stages 的数组字段(沿分辨率 stage 排列,浅 -> 深)
STAGE_LEN_FIELDS = ["features_per_stage", "kernel_sizes", "strides", "n_blocks_per_stage"]
# 长度 == n_stages - 1 的 decoder 数组字段
DECODER_LEN_FIELDS = ["n_conv_per_stage_decoder", "n_conv_per_stage"]

def _trim_field(container, key, new_len, expected_old_len, where, changes):
    """裁剪一个数组字段到 new_len,并做长度一致性校验。"""
    if key not in container:
        return
    old = container[key]
    if len(old) != expected_old_len:
        # 长度和预期不符 → 停止,不盲改(防止 plans 结构和假设不一致时误改)
        raise RuntimeError(f"{where}.{key} 长度 {len(old)} != 预期 {expected_old_len},停止")
    container[key] = old[:new_len]   # 从深端(列表末尾)裁剪
    changes.append(f"  {where}.{key}: {old} -> {container[key]}")

def make_reduced_plans(base_plans, base_name, target_stages, configuration):
    arch = base_plans["configurations"][configuration]["architecture"]["arch_kwargs"]
    old_n = int(arch["n_stages"])
    # 从深端移除、且不允许过于激进
    if target_stages >= old_n:  raise RuntimeError("目标 >= 当前,无需削减")
    if target_stages < 3:       raise RuntimeError("目标 < 3,过于激进,停止")

    arch["n_stages"] = target_stages
    # encoder 各数组裁到 target_stages
    for key in STAGE_LEN_FIELDS:
        _trim_field(arch, key, target_stages, old_n, "arch_kwargs", changes)
    # decoder 各数组裁到 target_stages - 1
    for key in DECODER_LEN_FIELDS:
        _trim_field(arch, key, target_stages - 1, old_n - 1, "arch_kwargs", changes)

    # 兜底:扫描 arch_kwargs 里还有没有长度恰好 == 旧 stage 数的未知数组字段
    for key, value in arch.items():
        if isinstance(value, list) and len(value) in (old_n, old_n - 1):
            if key not in known_fields:
                changes.append(f"  !! 警告: {key} (长度 {len(value)}) 未被处理,请人工核对")

    print("\n".join(changes))   # 打印所有改动供人工核对
    return base_plans

实际运行输出从深端裁 各数组同步

  • [nnUNetResEncUNetLPlans_shallow_6stage] 7 -> 6 stages
  • features_per_stage: [32,64,128,256,320,320,320] -> [32,64,128,256,320,320]
  • kernel_sizes: (7个) -> (6个)
  • strides: [[1,1,1],[1,2,2],[1,2,2],[2,2,2],[2,2,2],[2,2,2],[1,2,2]] -> (去掉末尾一个)
  • n_blocks_per_stage: [1,1,1,1,1,1,1] -> [1,1,1,1,1,1]
  • n_conv_per_stage_decoder: [1,1,1,1,1,1] -> [1,1,1,1,1] ← decoder 保持 n_stages-1
  • 保持不变: patch_size=[40,320,320], batch_size=2 ← 复用预处理数据

训练前构建校验

def validate_one(dataset, configuration, trainer, plans_name, config_path, device):
    # 1) 只读构建网络(不实例化 trainer、不调 initialize)
    build_info = report_cost.build_network(dataset, configuration, 0, trainer, plans_name, config_path)
    network, patch_size = build_info.network, build_info.patch_size

    # 2) 用 patch_size 大小的假数据跑一次前向
    x = torch.randn((1, input_channels, *patch_size), device=device)
    network.train()   # train 模式返回深监督各尺度输出
    with torch.no_grad():
        outputs = network(x)

    # 3) 检查主输出空间尺寸是否 == patch_size(对不上说明下采样/上采样链条有问题)
    main_shape = outputs[0].shape
    if tuple(main_shape[2:]) != tuple(patch_size):
        print(f"FAIL: 主输出 {main_shape[2:]} != patch_size {patch_size}")
        return False
    print(f"PASS: {plans_name} 构建与 forward 正常")
    return True

校验结果

**6stage: 12 层, 13.624M, 312 物理核, 深监督 5 尺度 5stage: 10 层, 6.648M, 240 物理核, 深监督 4 尺度 **

消融2 Homography / ResNet stage

Reproducibility, Replicability, and Repeatability: A survey of reproducible research with a focus on high performance computing

它能给消融2给出一个可对照的预期,这篇工作对编码器用几个 stage做了消融,结论是少stage会改变瓶颈分辨率,用较少stage会让分辨率减半。

所以这个结论为消融二提供了一个预期基线,让消融二带着一个明确的预期去验证,减stage 应该会掉点,问题是掉多少 什么时候掉到不可接受,有了预期结果才有参照系。

消融二是从深端逐级削减得到一条完整的砍几层vs精度曲线,这条曲线同时做了两件事,在 6 stage处偏离了文献预期,在 5stage处回归了文献预期。

砍第1层:偏离文献预期几乎无损, 6stage病例口径 Dice0.6796,和基线0.6802几乎完全重合(p=0.98,是所有实验里与基线差异最小的一次)。

文献说减stage会掉精度,但是砍掉最深那一层几乎没有代价。

砍第2层:回归文献预期开始掉点,5stage病例口径Dice0.6225vs基线p从6stage的0.98骤降到0.1266,虽然0.1266严格说仍未跌破0.05显著线,但是趋势明确了。

拐点在5stage和6stage之间,最深1层是纯冗余,砍掉纯赚参数,砍到5stage 开始付出可观测的代价。

IVH的机制

5 stage相对基线的精度损失集中在IVH

  • Class 3 (IVH) SG=0.8158 BL=0.8817 t-p=0.0626 W-p=0.0582
  • Class 1 (EDH) SG=0.4333 BL=0.4577 t-p=0.0855 W-p=0.0464
  • Class 2 (IPH) SG=0.6790 BL=0.6671
  • Class 4 (SAH) SG=0.6689 BL=0.6539
  • Class 5 (SDH) SG=0.5337 BL=0.5236

让claude解释了一下,虽然我也看不懂()IVH是脑室内出血,它的分布往往贯穿整个脑室系统,是相对大范围、全局性的结构。而砍掉最深的分辨率层级,损失的恰恰是网络”捕捉大范围上下文的能力,最深层看的是缩到 5×5 的全局图,负责把握整体结构。所以5 stage 损失的正是看大范围结构的能力,也就最伤 IVH 这种大范围病灶。

对这个脑出血CT分割任务,从深端逐级削减分辨率stage的效果是非线性的,砍第1个几乎零成本,最深的分辨率层级是纯冗余,砍掉后病例口径与基线几乎完全重合,参数从20降到 13。砍第2个触及下限,第2个分辨率层级开始承载有用的全局信息,砍掉后出现可观测的精度下降。