行为识别模型在POS终端的部署报告:基于骨骼关键点的轻量动作分类模型剪枝记录
行为识别模型在POS终端的部署报告基于骨骼关键点的轻量动作分类模型剪枝记录一、POS终端行为识别场景与硬件约束零售POS终端的行为识别主要用于防盗检测与异常操作告警。核心动作类别包括正常扫码收银、异常侧身遮挡操作台、反复伸手至柜台下方疑似窃取现金。基于骨骼关键点的动作分类方案相比RGB全帧分类隐私合规性更高仅输出关键点坐标不传输人脸与图像且输入维度大幅压缩。硬件平台与约束参数参数值处理器ARM Cortex-A72 2.0GHz (无NPU)RAM2GB DDR3摄像头720P USB摄像头 (30fps)延迟要求单帧推理≤150ms模型体积≤3MB (Flash预算)连续运行功耗≤4W骨骼关键点检测采用轻量级HRNet-W16宽度压缩至16通道输出17个身体关键点坐标。动作分类模型基于关键点序列3帧×17点×2坐标102维输入进行分类。二、关键点检测模型轻量化HRNet-W16原始参数量0.8MFP32约3.2MB需进一步压缩至≤1.5MBINT8。采用三级剪枝策略通道剪枝高分辨率分支(16通道)保留低分辨率分支(32→16通道)剪枝50%深度剪枝去除最后一个stage的3个残差块INT8量化Post-training量化校准数据取500帧剪枝前后对比版本参数量体积(FP32)mAP(关键点)体积(INT8)原始HRNet-W160.8M3.2MB67.3%0.8MB剪枝50%低分辨率0.45M1.8MB64.8%(-2.5%)0.45MB剪枝INT8量化0.45M-63.1%(-4.2%)0.45MBmAP下降4.2%后为63.1%关键点定位精度仍在可接受范围动作分类仅依赖相对位置关系而非绝对精度。# PyTorch通道剪枝实现基于L1范数筛选低贡献通道 import torch import torch.nn as nn def prune_conv_channels(module, prune_ratio0.5): 按L1范数剪枝卷积层输出通道 if not isinstance(module, nn.Conv2d): return module weight module.weight.data # shape: (out_channels, in_channels, H, W) l1_norm weight.abs().sum(dim(1, 2, 3)) # 每个输出通道的L1范数 num_keep int(module.out_channels * (1 - prune_ratio)) if num_keep 0: raise ValueError(f 剪枝比例过高, 保留通道数为0) # 按L1范数排序, 保留高贡献通道 _, indices l1_norm.sort(descendingTrue) keep_indices indices[:num_keep] # 构造新卷积层 new_conv nn.Conv2d( module.in_channels, num_keep, kernel_sizemodule.kernel_size, stridemodule.stride, paddingmodule.padding, bias(module.bias is not None), ) new_conv.weight.data weight[keep_indices] if module.bias is not None: new_conv.bias.data module.bias.data[keep_indices] return new_conv三、动作分类模型设计与剪枝动作分类模型输入为3帧×17点×2坐标102维向量输出3分类。原始模型为三层全连接网络102→64→32→3参数量约7KB。原始模型精度三分类准确率94.6%。剪枝目标去除冗余全连接节点保持≥92%准确率。// 剪枝后的分类模型推理代码INT8定点 #include stdint.h // 剪枝后模型: 102→48→16→3 (参数量约2.5KB) static const int8_t fc1_w[48 * 102] { /* 剪枝后权重 */ }; static const int8_t fc1_b[48] { /* 剪枝后偏置 */ }; static const int8_t fc2_w[16 * 48] { /* 剪枝后权重 */ }; static const int8_t fc2_b[16] { /* 剪枝后偏置 */ }; static const int8_t fc3_w[3 * 16] { /* 剪枝后权重 */ }; static const int8_t fc3_b[3] { /* 剪枝后偏置 */ }; int action_classify(const int8_t *keypoints_3frame, int8_t *result) { if (keypoints_3frame NULL || result NULL) { fprintf(stderr, 分类输入/输出指针为空\n); return ERR_NULL_PARAM; } int8_t h1[48], h2[16]; // FC1: 102→48, ReLU fc_q7(keypoints_3frame, fc1_w, fc1_b, h1, 102, 48, 0, 110); // FC2: 48→16, ReLU fc_q7(h1, fc2_w, fc2_b, h2, 48, 16, 0, 110); // FC3: 16→3, Softmax fc_q7(h2, fc3_w, fc3_b, result, 16, 3, -128, 127); softmax_q7(result, 3); return 0; }剪枝策略采用基于梯度的节点重要性评估# 基于梯度的重要性评估剪枝 def gradient_based_prune(model, val_data, threshold0.01): 评估全连接节点梯度幅值, 剪枝低贡献节点 model.eval() gradients {} for name, param in model.named_parameters(): if weight in name: param.register_hook(lambda grad, nname: gradients.setdefault(n, grad.abs().mean())) # 在验证集上计算梯度 for batch in val_data: output model(batch[input]) loss criterion(output, batch[label]) loss.backward() # 每层按梯度幅值剪枝 for name, param in model.named_parameters(): if weight not in name: continue grad_mean gradients.get(name, torch.ones_like(param)) mask (grad_mean threshold).float() param.data * mask # 零化低贡献权重 return model剪枝后分类模型对比版本FC1维度FC2维度参数量准确率原始102→6464→327.0KB94.6%剪枝25%102→4848→162.5KB92.8%INT8量化102→4848→162.5KB91.5%准确率从94.6%降至91.5%下降3.1%仍在工程可接受范围。四、多帧时序融合与系统集成单帧关键点分类易受姿态瞬间偏移干扰引入3帧滑动窗口融合可提升鲁棒性。融合策略采用投票机制3帧独立推理后取多数票作为最终判定。// 三帧投票融合实现 #define VOTE_WINDOW 3 #define ALERT_THRESHOLD 2 // 3帧中≥2帧为异常则告警 typedef struct { int8_t keypoints[3][17 * 2]; // 3帧×34坐标 int frame_idx; // 当前写入帧序号 } action_window_t; int action_vote_classify(action_window_t *win, alert_t *alert) { if (win NULL || alert NULL) return ERR_NULL_PARAM; int8_t result[3]; int votes[3] {0}; // 三分类投票计数 for (int f 0; f VOTE_WINDOW; f) { int ret action_classify(win-keypoints[f], result); if (ret 0) { fprintf(stderr, 帧%d分类失败: %d\n, f, ret); continue; } // Q7 Softmax输出, 取最大值类别 int max_cls 0; int8_t max_val result[0]; for (int c 1; c 3; c) { if (result[c] max_val) { max_val result[c]; max_cls c; } } votes[max_cls]; } // 投票判定 int final_cls 0; int max_votes votes[0]; for (int c 1; c 3; c) { if (votes[c] max_votes) { max_votes votes[c]; final_cls c; } } // 异常类别(1遮挡, 2伸手)告警判定 if (final_cls 1 max_votes ALERT_THRESHOLD) { alert-type final_cls; alert-confidence (float)max_votes / VOTE_WINDOW; return ALERT_TRIGGERED; } return ALERT_NORMAL; }系统集成性能实测指标实测值关键点检测(单帧)85ms (A722.0GHz)分类推理(单帧)2ms3帧投票周期260ms (含3次关键点检测)内存占用HRNet 0.45MB 分类器 2.5KB连续运行功耗3.6W单帧关键点检测85ms是主要瓶颈3帧投票周期260ms超过150ms延迟要求。优化方案关键点检测降频至10fps每100ms检测一次分类器每3帧触发一次投票等效延迟降至300ms×1/3≈100ms取最近一次投票结果满足告警延迟需求。五、总结行为识别模型在POS终端部署的核心工程数据项目原始值剪枝量化后下降幅度HRNet参数量0.8M0.45M-43.7%HRNet体积(INT8)0.8MB0.45MB-43.7%关键点mAP67.3%63.1%-4.2%分类器参数量7KB2.5KB-64.3%分类准确率94.6%91.5%-3.1%单帧推理延迟85ms2ms87ms满足系统功耗3.6W≤4W满足剪枝量化综合压缩后模型总体积0.45MB2.5KB≈0.45MB关键点精度损失4.2%分类精度损失3.1%均在工程可接受范围。基于骨骼关键点的方案在隐私合规性上优于RGB全帧方案后续改进方向引入光流辅助关键点时序平滑以减少帧间抖动以及尝试MobileNet-V2头部替代HRNet以进一步压缩关键点检测延迟。