1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114
| class ChairPose(nn.Module): """ ChairPose: 基于压力分布图的坐姿估计系统 架构: 1. 压力特征提取(CNN) 2. 椅子形态编码(MLP) 3. 特征融合(拼接) 4. 姿态回归(MLP) 输入: - 压力分布图:(B, 1, 64, 64) - 椅子参数:(B, 10) 输出: - 3D姿态:(B, 17, 3) """ def __init__(self, config: dict = None): super().__init__() config = config or {} self.pressure_encoder = PressureFeatureExtractor( in_channels=config.get('pressure_channels', 1), out_dim=config.get('feature_dim', 512) ) self.chair_encoder = ChairMorphologyEncoder( chair_feature_dim=config.get('chair_dim', 128) ) fusion_dim = 512 + 128 self.fusion = nn.Sequential( nn.Linear(fusion_dim, 512), nn.ReLU(), nn.Dropout(0.3) ) self.pose_regressor = PoseRegressor( feature_dim=512, num_joints=config.get('num_joints', 17) ) def forward(self, pressure_map: torch.Tensor, chair_params: torch.Tensor) -> torch.Tensor: """ 前向传播 Args: pressure_map: (B, 1, 64, 64) 压力分布图 chair_params: (B, 10) 椅子参数 Returns: pose_3d: (B, 17, 3) 3D姿态 """ pressure_feat = self.pressure_encoder(pressure_map) chair_feat = self.chair_encoder(chair_params) fused_feat = torch.cat([pressure_feat, chair_feat], dim=1) fused_feat = self.fusion(fused_feat) pose_3d = self.pose_regressor(fused_feat) return pose_3d
if __name__ == "__main__": """ 测试ChairPose模型 """ device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = ChairPose().to(device) print(f"模型参数量: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M") batch_size = 4 pressure_map = torch.randn(batch_size, 1, 64, 64).to(device) chair_params = torch.randn(batch_size, 10).to(device) model.eval() with torch.no_grad(): pose_3d = model(pressure_map, chair_params) print(f"\n输入:") print(f" 压力分布图: {pressure_map.shape}") print(f" 椅子参数: {chair_params.shape}") print(f"\n输出:") print(f" 3D姿态: {pose_3d.shape}") print(f"\n关节点坐标范围:") print(f" X: [{pose_3d[:, :, 0].min():.2f}, {pose_3d[:, :, 0].max():.2f}]") print(f" Y: [{pose_3d[:, :, 1].min():.2f}, {pose_3d[:, :, 1].max():.2f}]") print(f" Z: [{pose_3d[:, :, 2].min():.2f}, {pose_3d[:, :, 2].max():.2f}]") gt_pose = torch.randn(batch_size, 17, 3).to(device) mpjpe = torch.mean(torch.norm(pose_3d - gt_pose, dim=2)) print(f"\nMPJPE(示例): {mpjpe:.2f} mm")
|