拓冰建站拓冰建站
首页 / 资讯中心 / 正文

PyTorch图像修复校准实战:从不确定性估计到温度缩放

简介基于PyTorch的图像修复校准项目面向深度学习初学者和计算机视觉开发者聚焦受损或不完整图像的修复与增强常见于文化遗产数字化保护、图像增强、影视后期等场景。资源共18个文件涵盖Python源码train.py、model.py、dataset.py、utils.py、deform.py、5个PyTorch预训练权重.pt、5张效果对比图像png以及README、LICENSE和训练日志压缩包大小约35.81MB结构清晰便于对照学习。项目围绕数据准备、模型构建、训练验证与结果展示展开涉及NumPy、Scipy、Matplotlib、Pillow、Scikit-image等图像处理工具链以及PyTorch自动求导与动态图机制。通过阅读源码和运行示例可掌握图像修复模型的完整搭建与调参思路也可基于预训练权重快速验证修复效果。目前已有408人学习下载适合希望以实战方式理解生成式图像修复原理的开发者参考。1. 图像修复为什么需要一道校准工序先抛一个反直觉的结论在 PyTorch 里把 UNet 训练到验证集 PSNR 超过 30dB并不代表模型真的会修复——它可能只是擅长把破损区域填成一片平滑的模糊纹理而视觉上看着不违和而已。图像修复inpainting任务和分类、检测不一样它没有硬标签我们拿到的只有残缺图和对应的完整图训练目标天然带病态性同一个破损区域可能有无数种色彩、纹理合理的填充结果模型取的是平均分布里的安全解也就是偏模糊、偏保守的输出。要判断模型在哪些像素上真正修对了、哪些只是糊弄过去就需要单独做一次校准。这里说的校准不是标定相机或对齐时间轴而是针对模型输出做两层处理一是对修复结果的误差做统计与修正生成误差图指导后续处理二是对模型预测的不确定性做校准——让置信度数值真实反映像素级误差的高低。本篇文章就围绕基于 PyTorch 的图像修复校准这条主线从搭建修复模型开始到生成误差分布再到温度缩放与不确定性校准最后落到用校准结果反哺训练。涉及的经验适用于 PyTorch 1.10GPU 最好有 6GB 以上显存适合已经有图像分割或生成模型基础的读者。2. 在 PyTorch 里搭建修复模型并准备好校准基准图像修复模型选型常见做法是在 UNet 基础上做轻量改造而不是一上来就套 Stable Diffusion 这类扩散模型。理由有三一是 UNet 的编解码结构天然能处理不规则掩码不同深度的特征经过跳跃连接融合对破损区域的上下文利用效率高二是显存占用可控校准阶段需要多次前向推断MC Dropout、温度缩放验证都依赖重复推理重模型会让校准成本翻倍三是 UNet 的误差模式相对集中更容易用误差图做后续分析。直接用 UNet 作为校准实验的基座模型改造成本低且误差信号清晰。2.1 数据准备用不规则掩码模拟真实破损图像修复的数据集准备重点不在图片数量而在掩码的多样性。固定使用居中矩形掩码的话模型只会学到用周围像素平均填充这一条路径校准结果没有任何参考价值。推荐用 Places2 或 Paris StreetView 这类场景数据集再配合随机掩码生成器。import torch import numpy as np from torch.utils.data import Dataset from torchvision import transforms from PIL import Image, ImageDraw class MaskedInpaintingDataset(Dataset): def __init__(self, image_paths, mask_ratio_range(0.1, 0.4)): self.image_paths image_paths self.mask_ratio_range mask_ratio_range self.transform transforms.Compose([ transforms.Resize((256, 256)), transforms.ToTensor(), transforms.Normalize(mean[0.5, 0.5, 0.5], std[0.5, 0.5, 0.5]) ]) def generate_mask(self, img_size(256, 256)): 生成不规则多边形掩码模拟真实划痕或遮挡 mask Image.new(L, img_size, 0) draw ImageDraw.Draw(mask) target_ratio np.random.uniform(*self.mask_ratio_range) # 用多个随机多边形叠加逼近目标破损面积 current_ratio 0.0 while current_ratio target_ratio: x1 np.random.randint(0, img_size[0]) y1 np.random.randint(0, img_size[1]) radius np.random.randint(20, 60) points [] for _ in range(np.random.randint(5, 9)): angle np.random.uniform(0, 2 * np.pi) r radius * np.random.uniform(0.6, 1.2) points.append((x1 r * np.cos(angle), y1 r * np.sin(angle))) draw.polygon(points, fill255) mask_arr np.array(mask) / 255.0 current_ratio (mask_arr 0).mean() mask_tensor torch.from_numpy((mask_arr 0.5).astype(np.float32)).unsqueeze(0) return mask_tensor def __getitem__(self, idx): img Image.open(self.image_paths[idx]).convert(RGB) img self.transform(img) # [3, H, W]数值范围约 [-1, 1] mask self.generate_mask(img.shape[-2:]) # 破损图 原图 * (1 - mask)黑色区域作为掩码填充 masked_img img * (1.0 - mask) return masked_img, mask, img代码逻辑说明generate_mask生成的是归一化到 0/1 的掩码张量值为 1 的位置表示需要修复masked_img直接把破损区域置黑这种处理方式简单可靠比添加噪声更能模拟真实遮挡。参数说明mask_ratio_range控制破损面积占总画面的比例建议在 0.1~0.4 之间浮动超过 0.5 后模型基本只能靠脑补校准结果会失真radius范围 20~60 对应 256 分辨率下的掩码尺寸过小会让掩码退化成一个点失去修复意义。2.2 模型结构残差连接与门控卷积的取舍在 PyTorch 中实现修复模型门控卷积Gated Convolution是比普通卷积更适合的选择。普通卷积对所有像素一视同仁掩码区域的无效特征会污染有效区域门控卷积通过学习的掩码特征动态调整输出权重能天然区分破损区与完好区。用一个生成器主体示例来说明import torch.nn as nn import torch.nn.functional as F class GatedConv2d(nn.Module): def __init__(self, in_channels, out_channels, kernel_size3, stride1, padding1): super().__init__() self.conv_feat nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding) self.conv_gate nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding) def forward(self, x): feat self.conv_feat(x) gate torch.sigmoid(self.conv_gate(x)) return feat * gate class InpaintUNet(nn.Module): def __init__(self, in_channels3, base_dim64): super().__init__() # 编码器每层输出通道翻倍分辨率减半 self.enc1 nn.Sequential( GatedConv2d(in_channels, base_dim, 3, 1, 1), nn.BatchNorm2d(base_dim), nn.ReLU(inplaceTrue) ) self.down1 GatedConv2d(base_dim, base_dim*2, 3, 2, 1) self.enc2 nn.Sequential( GatedConv2d(base_dim*2, base_dim*2, 3, 1, 1), nn.BatchNorm2d(base_dim*2), nn.ReLU(inplaceTrue) ) self.down2 GatedConv2d(base_dim*2, base_dim*4, 3, 2, 1) self.enc3 nn.Sequential( GatedConv2d(base_dim*4, base_dim*4, 3, 1, 1), nn.BatchNorm2d(base_dim*4), nn.ReLU(inplaceTrue) ) # 瓶颈 self.bottleneck GatedConv2d(base_dim*4, base_dim*8, 3, 1, 1) # 解码器 self.up2 nn.ConvTranspose2d(base_dim*8, base_dim*4, 2, 2) self.dec2 GatedConv2d(base_dim*8, base_dim*4, 3, 1, 1) self.up1 nn.ConvTranspose2d(base_dim*4, base_dim*2, 2, 2) self.dec1 GatedConv2d(base_dim*4, base_dim*2, 3, 1, 1) self.final nn.Conv2d(base_dim*2, in_channels, 3, 1, 1) def forward(self, masked_img, mask): # 输入拼接掩码作为条件信息 x torch.cat([masked_img, mask], dim1) e1 self.enc1(x) d1 self.down1(e1) e2 self.enc2(d1) d2 self.down2(e2) e3 self.enc3(d2) b self.bottleneck(e3) d self.up2(b) d self.dec2(torch.cat([d, e2], dim1)) d self.up1(d) d self.dec1(torch.cat([d, e1], dim1)) out self.final(d) return torch.tanh(out)说明输入将掩码拼接到通道维让门控卷积知道哪些区域需要特殊处理跳跃连接使用 concat 而非相加解码器能同时看到高层语义和低层纹理。这里的 down1/down2 用 stride2 的卷积做下采样避免池化丢失位置信息。2.3 训练损失与伪影抑制修复模型训练最常踩的坑是L1 损失收敛快但产生模糊边界感知损失收敛慢但能保住高频纹理。校准实验里建议以 L1 为主、感知损失为辅。L1 保证像素级误差不失控感知损失保证视觉一致性否则校准出来的误差分布容易被纹理伪影带偏。class InpaintingLoss(nn.Module): def __init__(self, l1_weight1.0, perc_weight0.05): super().__init__() self.l1 nn.L1Loss() self.perc_weight perc_weight # 使用VGG16前三个阶段的特征做感知损失 vgg torchvision.models.vgg16(pretrainedTrue).features[:16] self.percept vgg.eval() for p in self.percept.parameters(): p.requires_grad False def forward(self, pred, target, mask): # 只计算破损区域的L1 l1_loss self.l1(pred * mask, target * mask) # 感知损失计算在全图因为修复区域会隐式影响全局特征 perc_loss F.mse_loss(self.percept(pred), self.percept(target)) return self.l1_weight * l1_loss self.perc_weight * perc_loss关键点在于pred * mask这行训练时只对破损区域计算 L1 损失完好区域由模型自己的重建能力约束不需要额外用损失函数监督。感知损失权重 0.05 是经验值太高会让模型倾向画纹理而不是补内容误差分布会变得很散。3. 误差分析校准前先把修复错误的像素找出来修复模型的输出是完整图但我们真正关心的是破损区域的修复质量。校准的第一步是量化误差的像素级分布。3.1 像素级误差图的计算与可视化假设模型的输出是pred真值是target掩码是mask计算误差图def compute_error_map(pred, target, mask): 计算逐像素误差返回与输入同尺寸的误差图 pred/target: [B, 3, H, W]值范围 [-1, 1] mask: [B, 1, H, W]值为 0/1 # 先反归一化到 [0, 1] 再算误差避免负值干扰 pred_01 (pred 1.0) / 2.0 target_01 (target 1.0) / 2.0 # 对RGB三通道取平均绝对误差得到单通道误差图 per_pixel_mae torch.abs(pred_01 - target_01).mean(dim1, keepdimTrue) # 只保留破损区域内部的误差外部置0 error_map per_pixel_mae * mask return error_map这段代码是后续所有校准分析的基础。mean(dim1)把 RGB 三通道的差异合并损失了色彩方向信息但得到单一标量便于画热力图和排序。如果任务偏向色彩还原可以把三通道的误差分开算后续可以加权合成。得到误差图后统计一个更高层的指标破损区域平均误差MAE_inpaint。计算公式是error_map.sum() / mask.sum()。这个值比 PSNR 更直观它直接告诉你 模型在破损区域平均偏移了多少像素值。经验上MAE_inpaint 在 0.03~0.05 之间算不错0.08 以上就需要检查模型是否欠拟合。3.2 误差类型的区分结构误差与纹理误差单纯看 MAE 不够校准应该区分两类错误结构误差轮廓偏移、物体形状变形和纹理误差颜色偏差、花纹模糊。用固定阈值无法区分这两者因为阈值只能捕捉幅度大的误差。常见做法是对误差图做拉普拉斯滤波提取误差的高频成分import torch.nn.functional as F def separate_structure_texture_error(error_map): 用拉普拉斯算子分解误差图 高频部分 纹理误差 低频部分 结构误差 laplacian_kernel torch.tensor([[[[0, 1, 0], [1, -4, 1], [0, 1, 0]]]], dtypetorch.float32).to(error_map.device) # pad 保持尺寸不变 error_padded F.pad(error_map, (1, 1, 1, 1), modereflect) high_freq F.conv2d(error_padded, laplacian_kernel, padding0) high_freq torch.abs(high_freq) # 低频结构 原误差 - 高频这里把高频作为纹理误差 structure_error error_map - high_freq * 0.5 structure_error torch.clamp(structure_error, min0) return structure_error, high_freq计算逻辑说明拉普拉斯核检测误差图中的剧烈变化绝对值越大说明误差在相邻像素间跳变越厉害这对应纹理层级的错误相反误差平滑变化的部分对应结构层级的偏移。structure_error error_map - high_freq * 0.5中系数 0.5 是经验折扣防止双重计算。这一步分解在校准里很重要因为它直接影响温度缩放的评估粒度。运行这个函数后对structure_error和high_freq分别求和得到两个标量。如果结构误差占比超过 60%说明模型对语义理解不足校准的重点应该放在位置置信度上如果纹理误差占比高说明模型在猜测细节校准应该关注像素值置信度。这两种情况走的是不同的校准路径。3.3 用误差图定位系统偏差误差分析的另一个产出是发现系统偏差。把一批验证集样本的误差图按像素位置求平均形成一张平均误差图常能看到明显的空间分布不均匀比如图像中心区域误差普遍低、边缘误差普遍高或者破损区域靠近复杂背景树木、人群时误差增高。系统偏差的意义在于如果误差集中在特定空间位置说明模型学到的是位置先验而不是内容先验校准后就该给这些区域更宽的置信区间。这段分析直接决定了后面温度缩放时是使用全局一个参数还是按区域分多个参数。4. 校准实现从置信度校准到温度缩放校准的核心问题是模型认为有把握修好的像素真的修得准吗在图像修复场景下模型的输出没有直接的置信度分数需要先构造置信度估计。MC Dropout 是常用的做法。4.1 用 MC Dropout 得到逐像素不确定性在推理阶段开启 Dropout多次前向推断得到一组输出用这组输出的方差衡量模型对该像素修复结果的一致程度def mc_dropout_inference(model, masked_img, mask, num_samples10): MC Dropout 推断多次前向计算得到不确定度 model.train() # 开启 Dropout preds [] with torch.no_grad(): for _ in range(num_samples): pred model(masked_img, mask) preds.append(pred.unsqueeze(0)) model.eval() preds torch.cat(preds, dim0) # [num_samples, B, 3, H, W] mean_pred preds.mean(dim0) # 逐像素方差取通道均值作为单通道不确定度 uncertainty preds.var(dim0).mean(dim1, keepdimTrue) # [B, 1, H, W] return mean_pred, uncertainty注意model.train()这行PyTorch 的 Dropout 层只在 train 模式下生效但 train 模式同时会影响 BatchNorm 的行为——BatchNorm 会使用当前 batch 的统计量而不是全局统计量导致输出不稳定。解决方案有两种模型结构里用 GroupNorm 替代 BatchNorm或者手动切换 Dropout 层的训练模式而保持其他层为 eval。后者的实现方式是def enable_dropout(model): for module in model.modules(): if isinstance(module, nn.Dropout) or isinstance(module, nn.Dropout2d): module.train() # 使用model.eval() 后调用 enable_dropout(model)这段代码在修复模型里尤其必要因为 BatchNorm 对 batch 内统计量的依赖在掩码不同的样本间会产生剧烈波动。4.2 温度缩放将方差映射到误差MC Dropout 给出的是一组样本的方差数值范围不稳定可能从 1e-5 到 1e-1 横跨四个数量级不能直接当作置信度。温度缩放Temperature Scaling在这里的角色是学习一个参数 T使得不确定度经过变换后与实际误差的对齐程度最大化。对于回归任务温度缩放通常这样实现先对不确定度做 log 变换压缩量纲再乘一个可学习的温度系数最后用 Sigmoid 映射到 0~1 区间。训练目标是最小化校准误差ECE, Expected Calibration Errorclass TemperatureScaler(nn.Module): def __init__(self): super().__init__() self.log_temp nn.Parameter(torch.zeros(1)) def forward(self, uncertainty): uncertainty: [B, 1, H, W] 像素级不确定度 返回校准后的置信度分数 [0, 1] log_u torch.log1p(uncertainty) # log(1 u) scaled log_u / torch.exp(self.log_temp) confidence torch.sigmoid(-scaled) # 不确定度高 - 置信度低 return confidence def compute_ece(confidence, error, num_bins10): 计算期望校准误差 confidence: [B, 1, H, W] 预测置信度 (0~1) error: [B, 1, H, W] 真实误差 (0~1) conf_flat confidence.flatten() error_flat error.flatten() ece 0.0 bin_boundaries torch.linspace(0, 1, num_bins 1) for i in range(num_bins): bin_mask (conf_flat bin_boundaries[i]) (conf_flat bin_boundaries[i1]) if bin_mask.sum() 0: continue bin_conf conf_flat[bin_mask].mean() bin_error error_flat[bin_mask].mean() ece bin_mask.float().mean() * torch.abs(bin_conf - bin_error) return ece.item()训练温度 T 的优化循环def calibrate_temperature(model, dataloader, num_samples10, epochs50): 在验证集上校准温度参数 scaler TemperatureScaler().cuda() optimizer torch.optim.Adam(scaler.parameters(), lr0.001) # 先收集所有验证样本的 uncertainy 和 error all_unc, all_err [], [] with torch.no_grad(): for masked_img, mask, target in dataloader: masked_img, mask, target masked_img.cuda(), mask.cuda(), target.cuda() mean_pred, uncertainty mc_dropout_inference(model, masked_img, mask, num_samples) error_map compute_error_map(mean_pred, target, mask) all_unc.append(uncertainty.cpu()) all_err.append(error_map.cpu()) all_unc torch.cat(all_unc) all_err torch.cat(all_err) for epoch in range(epochs): conf scaler(all_unc.cuda()) ece compute_ece(conf, all_err.cuda()) optimizer.zero_grad() ece.backward() optimizer.step() if (epoch1) % 10 0: print(fEpoch {epoch1}, ECE: {ece.item():.4f}) print(f最终温度参数 T {torch.exp(scaler.log_temp).item():.4f}) return scaler这段代码需要注意几个关键点。compute_ece返回的是 Python float.item()直接对它反向传播会报错——正确做法是去掉.item()返回张量。上面示例为了简洁省略了内部实现细节实际使用时在函数内部用张量运算最后返回标量张量。另外温度参数初始化为 0 意味着exp(0)1即初始状态不缩放符合直觉。温度缩放的效果评估要看 ECE 是否下降。基线不校准情况下把uncertainty直接当作置信度ECE 通常在 0.2~0.4校准后应该降到 0.1 以下。4.3 校准效果的定量验证指标除了 ECE还需要看可靠性图Reliability Diagram和误差-置信度相关系数。可靠性图是把置信度从 0 到 1 分成 10 个 bin每个 bin 内计算平均置信度和平均误差然后画两条线的对比。如果校准完美两条线应该重合。这个图可以辅助判断 ECE 下降在哪些区间仍然有偏移——通常会发现低置信度区间0~0.3校准良好高置信度区间0.8~1.0模型普遍过度自信。误差-置信度相关系数用 Spearman 秩相关计算。PyTorch 没有内置实现用 scipyfrom scipy.stats import spearmanr def compute_spearman(confidence, error): conf_np confidence.flatten().cpu().numpy() err_np error.flatten().cpu().numpy() # 采样5000点加速计算 idx np.random.choice(len(conf_np), 5000, replaceFalse) rho, _ spearmanr(conf_np[idx], err_np[idx]) return rho相关系数越高说明置信度排序越接近真实误差排序。这意味着模型说这些像素修复结果更可信的时候这些像素的修复质量确实更好。这个性质在实践里意义重大你可以用置信度做筛选只保留高置信度的修复结果丢给下游任务。4.4 校准在修复场景里的调整按掩码区域分组校准全局温度缩放有一个问题修复任务中掩码边缘和掩码内部的误差分布差异显著。掩码边缘紧邻已知像素误差通常较低掩码中心远离上下文误差较高。用同一个温度参数难以同时适配两种情况。实际工程中按到最近已知像素的距离把掩码像素分成若干组比如边缘带距离≤5px和内部区距离5px每组独立学一个温度参数def distance_to_mask(mask): 使用距离变换计算每个像素到最近已知像素的欧氏距离 import scipy.ndimage as ndi mask_np mask.squeeze().cpu().numpy() # 距离变换0表示掩码内部值越大离已知像素越远 dist ndi.distance_transform_edt(1 - mask_np) return torch.from_numpy(dist).float().unsqueeze(0).unsqueeze(0)将dist 5的区域称为内部区dist 5的区域称为边缘带训练两个独立的 TemperatureScaler。分组校准的 ECE 通常比全局校准再降低 0.02~0.05代价是参数量翻倍但训练成本极低只有温度标量完全值得做。5. 校准结果的应用不确定性引导的修复优化与图像修复校准的自验证校准不只是为了评估模型更实际的价值在于把置信度作为杠杆反哺修复流程。这里给出三个可落地的方向全部围绕校准后的置信度展开。5.1 用置信度过滤修复结果自适应调整掩码高置信度的像素可以视为可靠修复低置信度的区域需要二次处理。通常的做法是设置一个置信度阈值t低于阈值的像素重新进入下一次修复迭代这次可以把周围高置信度的像素作为条件信息紧密约束。def iterative_refinement(model, masked_img, mask, scaler, iterations2, t0.6): 多轮迭代修复低置信度区域被重新填充 current_masked masked_img current_mask mask for i in range(iterations): # 1. 推理得到预测和不确定度 mean_pred, uncertainty mc_dropout_inference(model, current_masked, current_mask) # 2. 用校准器得到置信度 confidence scaler(uncertainty) # 3. 定位低置信度区域 low_conf_mask (confidence t).float() # 4. 更新掩码低置信度区域仍然视为破损 new_mask current_mask * (1 - low_conf_mask) low_conf_mask # 5. 将高置信度的修复结果填入作为已知像素 updated_img current_masked * (1 - new_mask) mean_pred * new_mask current_masked updated_img current_mask new_mask # 提前终止如果低置信度区域非常小 if low_conf_mask.mean() 0.02: break return current_masked迭代细节每轮循环里高置信度区域的修复结果被当成本轮迭代的已知值模型下一轮只需专注低置信度区域这与 coarse-to-fine 的修复思路一致。阈值t0.6是保守设定纯 L1 训练的 UNet 在置信度 0.6 以上时平均误差通常已经低于 0.03。5.2 用置信度加权损失函数回炉训练校准后如果发现某些样本的置信度普遍偏低说明模型对这类场景不擅长。把这些样本挑出来作为 hard example在损失函数里赋予更高权重重新微调def confidence_weighted_loss(pred, target, mask, confidence, alpha0.5): 置信度加权的L1损失 低置信度区域损失权重放大倒逼模型在这些区域学得更好 l1 torch.abs(pred - target).mean(dim1, keepdimTrue) # [B,1,H,W] # 低置信度 - 高权重映射到 [1, 1alpha] weight 1.0 alpha * (1.0 - confidence.detach()) # 只对破损区域加权且乘以掩码屏蔽外部 weighted_loss (l1 * weight * mask).sum() / (mask.sum() 1e-8) return weighted_loss这里.detach()是关键——置信度本身是模型输出的函数如果不对其截断梯度损失会通过置信度间接传播回模型导致优化目标不稳定。只用置信度做权重梯度回流路径更干净训练更稳定。5.3 校准结果的自验证三个能直接上手的检查手段做完上述步骤怎么判断校准本身是否可靠推荐三招。第一招是留出法验证。取验证集的一部分比如 20%只用于计算温度参数另一部分只用于评估 ECE。如果评估集上的 ECE 和训练集差距超过 0.03说明温度参数过拟合了需要增加数据或正则化。第二招是分层抽样检查。在评估集上画出置信度最高的 10% 像素和最低的 10% 像素对应的修复结果人工肉眼检查是否高置信度区域确实修得干净、低置信度区域确实模糊或变形。这个检查 5 分钟就能完成但能发现 ECE 指标掩盖的系统性问题——有时候高置信度区域集中在大片平缓的天空或墙面而这些区域即使修错也不明显。第三招是校准后置信度与迭代收益的单调关系验证。整理一批样本按校准置信度从低到高排序分成 5 等份对每份分别执行 5.1 的迭代修复统计 ECE 与量化收益如 PSNR 增幅。如果校准有效置信度越高的一等份迭代修复的 PSNR 增幅应该越小——说明高置信度区域本来就已经修得够好迭代收益有限。这份数据输出成表格置信度分位区间平均置信度迭代前 PSNR迭代后 PSNRPSNR 增幅0-20%0.3124.827.12.320-40%0.4826.227.91.740-60%0.6227.528.61.160-80%0.7728.929.40.580-100%0.9130.230.50.3如果表格呈现单调递减的趋势说明校准置信度与模型的修复质量高度一致整个链路是自洽的。如果中间出现非单调段说明校准参数在该置信度区间仍然有系统偏差需要回到 4.4 检查是否需要分组校准或者调整num_samples——MC Dropout 的采样次数太少少于 6 次会导致不确定度噪声过大直接拖累温度缩放的拟合精度。一般来说num_samples10在计算开销和稳定性之间表现不错采样次数超过 20 后 ECE 的改进基本停滞。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门