You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在TIAToolBox中使用自定义PyTorch HoVer-Net模型进行细胞核分割

解决TIAToolBox中自定义HoVer-Net模型的细胞核检测问题

你遇到的ValueError是因为TIAToolBox的NucleusInstanceSegmentor需要明确知道模型的输入输出patch尺寸,或者提供对应的IO配置。以下是具体的解决步骤和修改后的代码:

核心问题说明

NucleusInstanceSegmentor无法自动推断自定义模型的输入输出patch规格,必须显式指定patch_input_shape和patch_output_shape(或提供ioconfig)。对于HoVer-Net,输入输出尺寸的对应关系由模型结构决定——比如常用训练配置中,输入patch为(256,256)时,输出patch为(164,164)(因模型的卷积/池化操作会缩小输出尺寸),你需要根据自己训练时的实际patch尺寸调整这两个参数。

修改后的完整代码

import torch
from tiatoolbox.models.architecture.hovernet import HoVerNet
from tiatoolbox.models.engine.nucleus_instance_segmentor import NucleusInstanceSegmentor

n_classes_pannuke = 6
ON_GPU = torch.cuda.is_available()

# 加载自定义HoVer-Net模型
hovernet = HoVerNet(n_classes=n_classes_pannuke)

# 加载训练好的权重(处理可能的DataParallel前缀问题)
state_dict = torch.load("hovernet_best_perf.pt")
# 如果训练时未使用DataParallel,需要移除权重中的module.前缀
# from collections import OrderedDict
# new_state_dict = OrderedDict()
# for k, v in state_dict.items():
#     name = k[7:] if k.startswith('module.') else k
#     new_state_dict[name] = v
# hovernet.load_state_dict(new_state_dict)
# 如果训练时用了DataParallel,直接加载即可
hovernet.load_state_dict(state_dict)

# 多GPU包装(如果需要)
if ON_GPU and torch.cuda.device_count() > 1:
    hovernet = torch.nn.DataParallel(hovernet)

hovernet.eval()

# 初始化分割器,关键是指定patch输入输出尺寸
inst_segmentor = NucleusInstanceSegmentor(
    model=hovernet,
    num_loader_workers=2,
    num_postproc_workers=2,
    batch_size=4,
    patch_input_shape=(256, 256),  # 替换为你训练时的输入patch尺寸
    patch_output_shape=(164, 164),  # 替换为对应模型输出的patch尺寸
    num_classes=n_classes_pannuke,
)

# 执行预测
tile_output = inst_segmentor.predict(
    imgs_to_predict,
    save_dir="sample_tile_results/",
    mode="tile",
    on_gpu=ON_GPU,
    crash_on_exception=True,
)

关键注意事项

  • patch尺寸匹配:必须确保patch_input_shape和patch_output_shape与你训练模型时使用的完全一致。如果不确定,可以通过测试模型输入输出获取:
    test_input = torch.randn(1, 3, 256, 256)  # 假设输入是3通道256x256
    with torch.no_grad():
        output = hovernet(test_input)
    print(output[0].shape)  # 查看输出尺寸,对应patch_output_shape
    
  • 权重加载问题:如果训练时没有使用DataParallel,直接加载会报错,此时需要移除权重字典中的module.前缀(代码中已注释相关处理逻辑)。
  • 模型输出格式:自定义HoVer-Net的输出必须符合TIAToolBox预期——返回包含核类型图、实例分割图、水平梯度图、垂直梯度图的元组,这是HoVer-Net的标准输出结构,只要训练时沿用官方或TIAToolBox的模型结构就无需修改。

内容的提问来源于stack exchange,提问作者WilliWespe

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.10 16:10:28