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 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189
| """ mmWave 隐私保护全身建模框架 从稀疏点云生成完整人体网格 """
import torch import torch.nn as nn import torch.nn.functional as F import numpy as np from typing import Optional
class PointCloudEncoder(nn.Module): """ 点云编码器 基于 PointNet++ 变体 """ def __init__(self, config: dict): super().__init__() in_dim = config.get('in_dim', 3) feat_dim = config.get('feat_dim', 256) self.mlp = nn.Sequential( nn.Linear(in_dim, 64), nn.ReLU(inplace=True), nn.Linear(64, 128), nn.ReLU(inplace=True), nn.Linear(128, 256), nn.ReLU(inplace=True), ) self.global_feat = nn.Sequential( nn.Linear(256, feat_dim), nn.ReLU(inplace=True), nn.Linear(feat_dim, feat_dim) ) def forward(self, points: torch.Tensor) -> torch.Tensor: """ Args: points: (B, N, 3) 点云坐标 Returns: global_feature: (B, feat_dim) """ point_feat = self.mlp(points) global_feat = point_feat.max(dim=1)[0] global_feat = self.global_feat(global_feat) return global_feat
class MeshDecoder(nn.Module): """ 网格解码器 从全局特征生成人体网格顶点 """ def __init__(self, config: dict): super().__init__() feat_dim = config.get('feat_dim', 256) n_vertices = config.get('n_vertices', 6890) self.pos_embed = nn.Parameter(torch.randn(1, n_vertices, feat_dim) * 0.02) decoder_layer = nn.TransformerDecoderLayer( d_model=feat_dim, nhead=4, dim_feedforward=feat_dim * 4, dropout=0.1, batch_first=True ) self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=4) self.vertex_head = nn.Sequential( nn.LayerNorm(feat_dim), nn.Linear(feat_dim, feat_dim // 2), nn.GELU(), nn.Linear(feat_dim // 2, 3) ) def forward(self, global_feat: torch.Tensor) -> torch.Tensor: """ Args: global_feat: (B, feat_dim) Returns: vertices: (B, n_vertices, 3) """ B = global_feat.shape[0] n_verts = self.pos_embed.shape[1] memory = global_feat.unsqueeze(1) query = self.pos_embed.expand(B, -1, -1) decoded = self.decoder(query, memory) vertices = self.vertex_head(decoded) return vertices
class mmWaveBodyMesh(nn.Module): """ 完整的 mmWave 人体网格生成模型 """ def __init__(self, config: dict): super().__init__() self.encoder = PointCloudEncoder(config) self.decoder = MeshDecoder(config) self.n_vertices = config.get('n_vertices', 6890) def forward(self, points: torch.Tensor) -> dict: """ Args: points: (B, N, 3) mmWave 点云 Returns: { 'vertices': (B, V, 3), # 顶点坐标 'global_feat': (B, feat_dim) # 全局特征 } """ global_feat = self.encoder(points) vertices = self.decoder(global_feat) return { 'vertices': vertices, 'global_feat': global_feat } def extract_pose_keypoints(self, vertices: torch.Tensor) -> torch.Tensor: """ 从网格顶点提取关键姿态关键点 简化版:从6890个顶点中选取17个关键点 """ joint_indices = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16] joints = vertices[:, joint_indices] return joints
if __name__ == "__main__": config = { 'in_dim': 3, 'feat_dim': 256, 'n_vertices': 6890, } model = mmWaveBodyMesh(config) points = torch.randn(2, 6, 3) * 2 output = model(points) print(f"输入点云: {points.shape}") print(f"输出网格: {output['vertices'].shape}") print(f"全局特征: {output['global_feat'].shape}") joints = model.extract_pose_keypoints(output['vertices']) print(f"姿态关键点: {joints.shape}") n_params = sum(p.numel() for p in model.parameters()) print(f"模型参数量: {n_params:,} ({n_params * 4 / 1024 / 1024:.1f} MB)")
|