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

为TensorFlow PoseNet模型添加关键点提取后处理层

将PoseNet后处理逻辑集成到TensorFlow模型的解决方案

要把你的后处理逻辑直接集成到PoseNet模型中,让模型直接输出关键点,核心是把原numpy实现的后处理转换成TensorFlow原生图操作,然后包装成自定义层,和原模型拼接成新的完整模型。以下是具体步骤和代码:

1. 将numpy后处理代码转换为TensorFlow操作

原代码依赖numpy,无法直接嵌入TensorFlow计算图,需要用tf的API替换所有numpy操作,同时兼容批量输入(模型通常处理批量数据):

import tensorflow as tf

def tf_parse_output(heatmap, offset):
    # 获取关节数量,PoseNet为17
    joint_num = heatmap.shape[-1]
    batch_size = tf.shape(heatmap)[0]
    
    # 1. 对每个关节的heatmap取最大值和对应坐标
    # 取每个关节heatmap的最大值(置信度)
    max_prob = tf.reduce_max(heatmap, axis=[1,2], keepdims=True)
    # 找到最大值在heatmap中的位置(y, x),注意tf.argmax的轴顺序
    max_pos_indices = tf.argmax(tf.reshape(heatmap, [batch_size, -1, joint_num]), axis=1)
    max_pos_y = tf.cast(max_pos_indices // heatmap.shape[2], tf.float32)
    max_pos_x = tf.cast(max_pos_indices % heatmap.shape[2], tf.float32)
    
    # 2. 将heatmap坐标映射到原输入图像尺寸(假设输入图像是257x257)
    remap_x = max_pos_x / 8 * 257
    remap_y = max_pos_y / 8 * 257
    
    # 3. 从offset中取出对应位置的偏移量
    # 把坐标转换为整数索引,用于索引offset张量
    max_pos_y_int = tf.cast(max_pos_y, tf.int32)
    max_pos_x_int = tf.cast(max_pos_x, tf.int32)
    # 构造batch索引
    batch_indices = tf.tile(tf.expand_dims(tf.range(batch_size), axis=1), [1, joint_num])
    
    # 取出x方向偏移量(前17个通道)
    offset_x = tf.gather_nd(offset, tf.stack([batch_indices, max_pos_y_int, max_pos_x_int, tf.range(joint_num)], axis=-1))
    # 取出y方向偏移量(后17个通道)
    offset_y = tf.gather_nd(offset, tf.stack([batch_indices, max_pos_y_int, max_pos_x_int, tf.range(joint_num, joint_num*2)], axis=-1))
    
    # 4. 计算最终关键点坐标和置信度
    kp_x = remap_x + offset_x
    kp_y = remap_y + offset_y
    kp_conf = tf.squeeze(max_prob, axis=[1,2])
    
    # 拼接成[batch_size, 17, 3]的输出格式(x, y, conf)
    pose_kps = tf.stack([kp_x, kp_y, kp_conf], axis=-1)
    return pose_kps

2. 创建自定义后处理层

把上面的tf函数包装成Keras自定义层,方便和原模型拼接:

class PoseNetPostProcess(tf.keras.layers.Layer):
    def __init__(self, **kwargs):
        super().__init__(**kwargs)
    
    def call(self, inputs):
        heatmap, offset = inputs
        return tf_parse_output(heatmap, offset)
    
    # 必须定义get_config,否则模型保存后无法加载
    def get_config(self):
        return super().get_config()

3. 加载原PoseNet模型并拼接后处理层

假设你的原模型是SavedModel格式,加载后把后处理层接在输出端:

# 加载预训练的PoseNet模型
original_posenet = tf.keras.models.load_model("path/to/your/posenet/savedmodel")

# 获取原模型的输入和输出
input_layer = original_posenet.input
heatmap_output, offset_output = original_posenet.output

# 添加后处理层
post_process_layer = PoseNetPostProcess()
keypoints_output = post_process_layer([heatmap_output, offset_output])

# 构建新的完整模型
new_posenet = tf.keras.Model(inputs=input_layer, outputs=keypoints_output)

# 测试模型输出
test_image = tf.random.normal([1, 257, 257, 3])  # 模拟输入
keypoints = new_posenet.predict(test_image)
print(keypoints.shape)  # 应该输出(1, 17, 3)

# 保存新模型
new_posenet.save("path/to/save/new_posenet_with_postprocess")

关键注意事项

  • 批量兼容:原numpy代码是单样本处理,转换后的tf代码支持批量输入(batch维度),符合模型部署的常规需求。
  • 索引处理:tf中用tf.gather_nd替代numpy的索引方式,确保能正确从offset张量中取出对应位置的偏移量。
  • 模型可保存性:自定义层必须实现get_config方法,否则SavedModel无法正确序列化。
  • 坐标映射:原代码中的max_val_pos/8*257是假设PoseNet的输入图像尺寸为257x257,若你的模型输入尺寸不同,需要调整这个比例。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 04:10:31