如何将Mask R-CNN训练生成的.ckpt文件转为.h5或用于检测?
将Mask R-CNN的.ckpt模型转为.h5或直接用于检测的方法
方法一:把.ckpt转成.h5文件
情况1:.ckpt是PyTorch训练生成,检测代码用TensorFlow/Keras
由于两个框架的层参数格式、命名逻辑存在差异,需要手动映射权重:
- 加载PyTorck的ckpt模型
import torch # 导入你训练时用的基于resnet18的Mask R-CNN结构 from your_model_module import MaskRCNN_ResNet18 # 初始化模型,类别数需和训练时一致 model = MaskRCNN_ResNet18(num_classes=你的类别数) # 加载ckpt权重,部分ckpt会把权重存于state_dict键下 checkpoint = torch.load('your_model.ckpt') model.load_state_dict(checkpoint.get('state_dict', checkpoint)) model.eval()
- 构建对应结构的Keras模型,逐层迁移权重
比如卷积层权重,PyTorch格式为(out_channels, in_channels, h, w),TensorFlow为(h, w, in_channels, out_channels),需转置调整:
import tensorflow as tf from tensorflow.keras.models import Model # 定义和PyTorck完全对齐的Keras版Mask R-CNN结构 keras_model = build_your_keras_maskrcnn(num_classes=你的类别数) # 遍历层迁移权重 for (pt_name, pt_module), (tf_name, tf_module) in zip(model.named_modules(), keras_model.named_layers()): # 处理卷积层 if isinstance(pt_module, torch.nn.Conv2d) and isinstance(tf_module, tf.keras.layers.Conv2D): kernel = pt_module.weight.data.numpy().transpose(2, 3, 1, 0) bias = pt_module.bias.data.numpy() tf_module.set_weights([kernel, bias]) # 处理批量归一化层 elif isinstance(pt_module, torch.nn.BatchNorm2d) and isinstance(tf_module, tf.keras.layers.BatchNormalization): tf_module.set_weights([ pt_module.weight.data.numpy(), pt_module.bias.data.numpy(), pt_module.running_mean.data.numpy(), pt_module.running_var.data.numpy() ])
- 保存为.h5文件
keras_model.save('converted_model.h5')
情况2:.ckpt是TensorFlow训练生成(如用TF Object Detection API)
直接加载ckpt再转存即可:
import tensorflow as tf # 加载ckpt模型 model = tf.keras.models.load_model('path/to/ckpt_directory') # 保存为h5格式 model.save('converted_model.h5')
方法二:直接用.ckpt做检测,无需转格式
情况1:PyTorch的ckpt
修改原h5依赖代码为PyTorch推理逻辑:
import torch import cv2 import numpy as np from torchvision.transforms import functional as F from your_model_module import MaskRCNN_ResNet18 # 加载模型 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = MaskRCNN_ResNet18(num_classes=你的类别数) checkpoint = torch.load('your_model.ckpt') model.load_state_dict(checkpoint.get('state_dict', checkpoint)) model.to(device) model.eval() def detect(image_path): # 读取并预处理图像 img = cv2.imread(image_path) img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img_tensor = F.to_tensor(img_rgb).unsqueeze(0).to(device) # 推理 with torch.no_grad(): preds = model(img_tensor) # 解析结果,对齐原h5代码的输出格式 boxes = preds[0]['boxes'].cpu().numpy() scores = preds[0]['scores'].cpu().numpy() masks = preds[0]['masks'].cpu().numpy() labels = preds[0]['labels'].cpu().numpy() # 后续按原代码逻辑处理结果(如画框、掩码) return boxes, scores, masks, labels
情况2:TensorFlow的ckpt
直接加载ckpt进行推理:
import tensorflow as tf # 加载ckpt模型 model = tf.saved_model.load('path/to/ckpt_directory') infer_fn = model.signatures['serving_default'] def detect(image_path): # 读取并预处理图像,尺寸需和训练时一致 img = tf.io.read_file(image_path) img = tf.image.decode_image(img, channels=3) img = tf.expand_dims(img, axis=0) img = tf.image.resize(img, (512, 512)) # 推理 preds = infer_fn(img) # 解析结果,适配原代码逻辑 boxes = preds['detection_boxes'].numpy()[0] scores = preds['detection_scores'].numpy()[0] masks = preds['detection_masks'].numpy()[0] return boxes, scores, masks
内容的提问来源于stack exchange,提问作者Dhanraj Jain
相关产品推荐
相关产品推荐

