交警手势识别实战:从CNN-LSTM建模到TensorRT端侧部署
简介本资源是一套基于Python与PyTorch框架实现的中国交通警察指挥手势识别系统专为本科毕业设计、高校课程设计及AI项目开发场景打造面向具备基础深度学习与计算机视觉知识的学习者解决交通手势图像分类与实时识别的实际问题。压缩包共37个文件含31个核心Python脚本涵盖数据预处理、模型训练、推理预测、可视化等全流程、2份Markdown文档中英文README说明项目结构与运行步骤、1个LICENSE授权文件、1个GIF演示动图及配套数据集加载与模型配置模块整体体积仅4.43MB轻量易部署。已有498人下载学习资源结构清晰包含aichallenger与pgdataset双数据集支持、models模型权重目录、pred预测接口及docs说明文档提供从环境配置、训练调参到结果可视化的完整闭环方案代码经严格测试可直接运行并支持二次开发与算法优化。1. 中国交通警察指挥手势识别不是“拍个照就能认”而是从数据采集到端侧部署的完整闭环你可能试过用 OpenCV HOG SVM 做手势分类结果在真实路口视频里一帧都对不上——光照突变、手臂遮挡、警服反光、手势起始帧模糊全崩。这个项目不是那种“跑通 demo 就交差”的毕业设计套壳它是一套实打实跑在交警实训视频、支持单帧推理与视频流持续跟踪、带标注规范模型剪枝轻量部署链路的完整工程。核心是 PyTorch 实现的时序增强型 CNN-LSTM 混合架构非纯静态图分类专为解决“同一手势在不同角度/速度下形态差异大”这一行业痛点设计数据集包含 8 类国标手势直行、停止、左转、右转、示意车辆靠边停车等共 12,643 张高质量标注图像 47 段原始训练视频含时间戳与关键帧标记源码里甚至预埋了 ONNX 导出脚本和 TensorRT 加速模板。适合需要真实交付能力的本科生毕设、高职课程设计或中小安防集成商快速验证算法可行性——它不教你什么是卷积但会告诉你为什么batch_size8在 RTX3060 上卡死以及怎么把.pth模型压到 15MB 还保持 92.3% mAP。2. 数据集结构与标注规范为什么 PGDataset 比公开手势库更适配中国场景这个项目的数据根基不是随便爬来的网络图而是基于《GB/T 23831-2009 道路交通信号灯设置与安装规范》和一线交警支队提供的教学视频逐帧拆解构建的。PGDatasetPolice Gesture Dataset是整个项目的“地基”理解它的组织逻辑才能避免后续训练时出现标签错位、尺度失真、类别漏标等血泪问题。2.1 PGDataset 目录结构与文件含义项目根目录下的pgdataset/是核心数据区其结构严格遵循工业级数据管理规范pgdataset/ ├── annotations/ # 所有标注文件JSON格式 │ ├── train.json # 训练集标注含 bbox gesture_class frame_id │ ├── val.json # 验证集标注同上独立于训练集 │ └── test.json # 测试集标注完全隔离用于最终评估 ├── images/ # 原始图像JPEG格式命名规则video_001_frame_00123.jpg │ ├── train/ # 训练图像按视频分组保留原始拍摄顺序 │ ├── val/ # 验证图像 │ └── test/ # 测试图像 ├── videos/ # 原始视频素材MP4格式用于数据增强与时序建模 │ ├── train_videos/ # 32段训练视频含交警制服、不同天气、早晚光线 │ └── test_videos/ # 15段测试视频含遮挡、远距离、运动模糊场景 └── README_pgdataset.md # 标注细则、手势定义表、坐标系说明必读提示不要直接用images/train/下所有图做训练——PGDataset 的train.json中每个样本都附带frame_id和video_id这是为后续时序建模如 LSTM 输入预留的索引。若忽略此字段LSTM 模块将失去时间连续性导致模型把“停止手势”的起始帧和结束帧当成两个独立样本处理精度暴跌 15%。2.2 标注字段详解不只是 bbox更是语义时空锚点打开annotations/train.json你会看到类似这样的条目{ image_id: video_007_frame_00892, file_name: images/train/video_007_frame_00892.jpg, width: 1280, height: 720, gesture_class: 3, bbox: [321.5, 187.2, 142.8, 215.6], video_id: video_007, frame_id: 892, keypoint: [[342,215],[367,241],[389,267],[412,293],[435,319]], gesture_speed: medium, light_condition: daylight }关键字段说明字段类型含义为什么重要gesture_classint0~7 八类手势 ID0直行1停止2左转预备3左转4右转预备5右转6靠边停车7示意车辆通行项目constants.py中硬编码映射修改此处必须同步改 constants.GESTURE_MAPbboxlist[float][x_min, y_min, width, height]COCO 格式检测模块输入基础若用 YOLOv5 训练需先转为xywh归一化格式keypointlist[list[int]]5 个关键点坐标手腕、肘、肩、指尖、掌心支持姿态估计分支用于解决“手臂被身体遮挡”时的手势判别gesture_speedstrslow/medium/fast用于加权损失函数快动作帧权重 ×1.2慢动作 ×0.8缓解动作速率不均导致的梯度偏移light_conditionstrdaylight/dusk/night训练时可作为 domain adaptation 的辅助标签提升夜间鲁棒性2.3 数据增强策略不是简单加高斯噪声而是模拟真实执法环境项目basic_tests/目录下提供了data_augmentation_demo.py它演示了针对交警手势的领域定制增强而非通用 torchvision.transforms# basic_tests/data_augmentation_demo.py import cv2 import numpy as np from albumentations import ( HorizontalFlip, RandomBrightnessContrast, MotionBlur, GaussianBlur, RandomShadow ) # 针对中国交警场景的增强组合已验证提升 val mAP 2.1% transform Compose([ HorizontalFlip(p0.5), # 模拟左右车道视角切换 RandomBrightnessContrast(brightness_limit0.3, contrast_limit0.3, p0.8), MotionBlur(blur_limit7, p0.3), # 模拟高速行驶中拍摄的运动模糊 GaussianBlur(blur_limit(3, 7), p0.3), # 模拟雨天镜头水汽 RandomShadow(num_shadows_lower1, num_shadows_upper3, shadow_dimension5, p0.4), # 模拟烈日下警帽投射阴影 ])注意RandomShadow的shadow_dimension5是经过实测调优的——值太小3阴影过于锐利像贴图太大8则覆盖整个手臂区域导致关键点丢失。这个参数在train/config.yaml中被固化若你替换自己的数据务必用basic_tests/visualize_aug.py可视化前 100 帧增强效果确认阴影未遮挡手掌区域。2.4 数据集加载器实现如何让 DataLoader 不丢帧、不错序、不爆显存项目ctpgr.py中的GestureDataset类不是简单继承torch.utils.data.Dataset它强制保证同一视频的帧按 frame_id 严格升序加载这对 LSTM 输入至关重要# ctpgr.py class GestureDataset(Dataset): def __init__(self, ann_file, img_dir, transformNone, is_video_modeFalse): self.anns json.load(open(ann_file))[annotations] # 关键按 video_id 分组再按 frame_id 排序 self.video_groups defaultdict(list) for ann in self.anns: self.video_groups[ann[video_id]].append(ann) for vid in self.video_groups: self.video_groups[vid].sort(keylambda x: x[frame_id]) # 强制时间连续 self.img_dir img_dir self.transform transform self.is_video_mode is_video_mode # True 时返回 (seq_len, C, H, W) tensor def __getitem__(self, idx): # 若启用 video_modeidx 指向 video_id而非单帧索引 if self.is_video_mode: video_id list(self.video_groups.keys())[idx] seq_anns self.video_groups[video_id][:32] # 截取最长32帧序列 frames [] for ann in seq_anns: img cv2.imread(os.path.join(self.img_dir, ann[file_name])) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) if self.transform: img self.transform(imageimg)[image] frames.append(torch.from_numpy(img).permute(2,0,1)) return torch.stack(frames), torch.tensor([ann[gesture_class] for ann in seq_anns]) else: ann self.anns[idx] # ... 单帧加载逻辑参数说明is_video_modeTrue启用时序模式__getitem__返回(32, 3, 224, 224)张量供LSTMGestureClassifier使用seq_len32由train/config.yaml中MAX_SEQ_LEN: 32控制过长48会导致 OOM过短16无法捕获手势起止过程video_groups分组机制避免 DataLoader 多进程打乱帧序这是很多初学者训练 LSTM 时精度上不去的根源。3. 模型架构与训练流程CNN-LSTM 混合不是噱头是解决动态手势的本质方案纯 CNN 做静态图分类在交警手势这种强时序动作上必然失效——“左转预备”和“左转”仅差一个手腕旋转角度单帧极易混淆。本项目采用CNN 提特征 LSTM 建模时序 Attention 加权关键帧的三级结构这才是能落地的真实方案。3.1 模型主干ResNet18 BiLSTM Temporal Attentionmodels/gesture_classifier.py中的LSTMGestureClassifier定义如下# models/gesture_classifier.py import torch import torch.nn as nn from torchvision.models import resnet18 class LSTMGestureClassifier(nn.Module): def __init__(self, num_classes8, hidden_size256, num_layers2, dropout0.3): super().__init__() self.cnn resnet18(pretrainedTrue) self.cnn.fc nn.Identity() # 移除原fc层取倒数第二层输出512维 self.lstm nn.LSTM( input_size512, hidden_sizehidden_size, num_layersnum_layers, batch_firstTrue, bidirectionalTrue, dropoutdropout if num_layers 1 else 0 ) self.attention nn.Sequential( nn.Linear(hidden_size * 2, 128), # BiLSTM 输出维度为 hidden_size*2 nn.Tanh(), nn.Linear(128, 1) ) self.classifier nn.Sequential( nn.Dropout(0.5), nn.Linear(hidden_size * 2, 128), nn.ReLU(), nn.Dropout(0.3), nn.Linear(128, num_classes) ) def forward(self, x): # x: (B, T, C, H, W) - (B*T, C, H, W) B, T, C, H, W x.shape x x.view(B*T, C, H, W) features self.cnn(x) # (B*T, 512) features features.view(B, T, -1) # (B, T, 512) lstm_out, _ self.lstm(features) # (B, T, hidden_size*2) # Temporal Attention attn_weights torch.softmax(self.attention(lstm_out), dim1) # (B, T, 1) context torch.sum(attn_weights * lstm_out, dim1) # (B, hidden_size*2) return self.classifier(context)关键设计点解析bidirectionalTrue让 LSTM 同时感知手势“起始→结束”和“结束→起始”两个方向的动态变化对“停止”这类对称手势尤其有效attention模块不是简单取最后时刻输出而是加权聚合所有时刻实验表明对“靠边停车”需多帧判断手臂下压幅度提升 3.7% 准确率cnn.fc nn.Identity()复用 ResNet18 的 backbone但冻结前 3 个 stage 的参数见train/train.py中freeze_cnn_layers()只微调最后 stage LSTM 全部参数防止小数据集过拟合。3.2 训练配置为什么 learning_rate1e-4 而不是 1e-3train/config.yaml是训练的中枢其中几个参数直接决定收敛质量# train/config.yaml TRAIN: batch_size: 8 num_workers: 4 epochs: 60 lr: 0.0001 # 关键ResNet18 backbone 已预训练过大lr导致特征坍塌 weight_decay: 1e-4 scheduler: StepLR # 每20 epoch *0.1 loss_fn: LabelSmoothingLoss # label_smoothing0.1缓解8类手势中“直行/停止”样本不均衡 MODEL: hidden_size: 256 num_layers: 2 dropout: 0.3 DATA: img_size: [224, 224] max_seq_len: 32 use_keypoints: true # 启用关键点分支见3.3节血泪经验曾用lr1e-3训练第5 epoch val loss 突然飙升至 5.2正常应 1.5查看grad_norm发现 CNN backbone 的梯度爆炸1000。原因预训练权重对新任务敏感需小步微调。lr1e-4是经 12 次消融实验确定的临界值——再小1e-5收敛太慢再大2e-4第3 epoch 就开始震荡。3.3 关键点引导分支用 5 个点解决“手臂遮挡”难题当车辆或身体遮挡部分手臂时纯 bbox 分类器会失效。本项目在models/keypoint_head.py中嵌入轻量级关键点回归头与主分类分支联合训练# models/keypoint_head.py class KeypointHead(nn.Module): def __init__(self, in_channels512, num_keypoints5): super().__init__() self.conv1 nn.Conv2d(in_channels, 256, kernel_size3, padding1) self.bn1 nn.BatchNorm2d(256) self.conv2 nn.Conv2d(256, 128, kernel_size3, padding1) self.bn2 nn.BatchNorm2d(128) self.conv3 nn.Conv2d(128, num_keypoints, kernel_size1) # 输出5通道热图 def forward(self, x): # x: (B, 512, 7, 7) 来自 ResNet18 layer4 输出 x F.relu(self.bn1(self.conv1(x))) x F.relu(self.bn2(self.conv2(x))) heatmaps self.conv3(x) # (B, 5, 7, 7) return heatmaps # 在 LSTMGestureClassifier.forward() 中调用 # keypoint_heatmaps self.keypoint_head(cnn_features.view(B*T, 512, 7, 7)) # kp_loss self.kp_criterion(keypoint_heatmaps, target_keypoints) # MSE Loss训练技巧关键点损失权重设为0.3train/config.yaml中kp_weight: 0.3过高会削弱主分类任务热图生成使用gaussian_kernel3见utils/keypoint_utils.py匹配标注中关键点 ±3 像素误差容忍度推理时仅用关键点热图做后处理对每个热图取 argmax 得到坐标若某点置信度0.3则该帧标记为“遮挡”触发 LSTM 时序补偿机制跳过该帧用前后帧加权。3.4 训练脚本执行从零开始跑通的 4 个命令所有训练逻辑封装在train/train.py无需修改代码即可启动# 步骤1准备数据软链接避免路径硬编码 ln -sf /path/to/your/pgdataset pgdataset # 步骤2检查数据完整性自动校验 JSON 结构、图像存在性、关键点范围 python train/check_dataset.py --dataset_root pgdataset # 步骤3启动训练默认使用 config.yaml可指定其他配置 python train/train.py --config train/config.yaml --device cuda:0 # 步骤4实时监控TensorBoard 日志在 logs/ 目录 tensorboard --logdir logs/ --bind_all参数说明--device cuda:0显存不足时可换cpu或cuda:1但batch_size需同步调小--config支持多配置如train/config_laptop.yaml专为 GTX16504GB优化batch_size4,hidden_size128check_dataset.py会输出缺失文件列表例如images/train/video_012_frame_00456.jpg not found这是新手最常踩的坑——务必先运行此脚本。4. 模型推理与部署ONNX TensorRT 不是炫技是让模型跑进嵌入式盒子训练完的.pth模型不能直接上车。本项目提供从 PyTorch → ONNX → TensorRT 的完整部署链路目标是让模型在 Jetson Nano2GB RAM上以 12 FPS 运行 720p 视频流。4.1 PyTorch 模型导出 ONNX避开 dynamic_axes 的三大陷阱pred/export_onnx.py是导出脚本但直接运行会报错——因为 LSTM 的batch_firstTrue与 ONNX 的动态轴要求冲突# pred/export_onnx.py def export_model_to_onnx(model_path, onnx_path, input_shape(1, 32, 3, 224, 224)): model LSTMGestureClassifier() model.load_state_dict(torch.load(model_path, map_locationcpu)) model.eval() # 关键构造 dummy_input 必须 match training input shape dummy_input torch.randn(input_shape) # (B1, T32, C3, H224, W224) torch.onnx.export( model, dummy_input, onnx_path, export_paramsTrue, opset_version12, do_constant_foldingTrue, input_names[input], output_names[output], # 动态轴声明B 和 T 必须动态否则 TensorRT 无法 batch 推理 dynamic_axes{ input: {0: batch_size, 1: seq_len}, # 注意这里 1 是 seq_len 维度 output: {0: batch_size} } ) print(fONNX exported to {onnx_path}) if __name__ __main__: export_model_to_onnx(models/best.pth, models/gesture.onnx)避坑清单现象原因解决RuntimeError: ONNX export failed: Couldnt export operator aten::lstmPyTorch 版本 1.12 且未指定opset_version12降级到torch1.12.1或显式设opset_version12ONNX checker failed: Node input input has incorrect rankdummy_inputshape 错误如写成(1,3,224,224,32)严格按(B,T,C,H,W)顺序T 必须是第2维TensorRT engine build failed: Unsupported ONNX data typeONNX 中存在int64类型张量常见于torch.arange在模型 forward 中强制torch.arange(..., dtypetorch.int32)4.2 TensorRT 引擎构建Jetson Nano 上的 12 FPS 是怎么榨出来的pred/build_engine.py将 ONNX 转为 TRT 引擎关键在于fp16_modeTrue和max_workspace_size的平衡# pred/build_engine.py import tensorrt as trt def build_engine(onnx_file_path, engine_file_path, fp16_modeTrue, max_batch_size1): TRT_LOGGER trt.Logger(trt.Logger.WARNING) builder trt.Builder(TRT_LOGGER) network builder.create_network(1 int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)) parser trt.OnnxParser(network, TRT_LOGGER) with open(onnx_file_path, rb) as model: if not parser.parse(model.read()): print(ERROR: Failed to parse the ONNX file.) for error in range(parser.num_errors): print(parser.get_error(error)) return None config builder.create_builder_config() config.max_workspace_size 1 30 # 1GB GPU memory if fp16_mode and builder.platform_has_fast_fp16: config.set_flag(trt.BuilderFlag.FP16) # 关键设置动态 batch 和 seq_len profile builder.create_optimization_profile() profile.set_shape(input, (1, 16, 3, 224, 224), (1, 32, 3, 224, 224), (1, 48, 3, 224, 224)) config.add_optimization_profile(profile) engine builder.build_engine(network, config) with open(engine_file_path, wb) as f: f.write(engine.serialize()) return engine参数说明max_workspace_size130Jetson Nano 仅 2GB 显存设为 1GB 是安全上限再大触发 OOMprofile.set_shape定义(min, opt, max)三元组opt32对应训练时MAX_SEQ_LEN这是性能拐点fp16_modeTrue开启后推理速度提升 1.8×但需确认 Nano 的 CUDA 版本 ≥10.2nvcc --version。4.3 实时视频流推理用 OpenCV TRT 实现 12 FPSpred/inference_trt.py是最终部署脚本它绕过 PyTorch直接调用 TRT 引擎# pred/inference_trt.py import cv2 import numpy as np import pycuda.autoinit import pycuda.driver as cuda import tensorrt as trt class TRTInference: def __init__(self, engine_path): self.engine self.load_engine(engine_path) self.context self.engine.create_execution_context() self.inputs, self.outputs, self.bindings, self.stream self.allocate_buffers() def load_engine(self, engine_path): with open(engine_path, rb) as f: runtime trt.Runtime(trt.Logger(trt.Logger.WARNING)) return runtime.deserialize_cuda_engine(f.read()) def allocate_buffers(self): inputs [] outputs [] bindings [] stream cuda.Stream() for binding in self.engine: size trt.volume(self.engine.get_binding_shape(binding)) * self.engine.max_batch_size dtype trt.nptype(self.engine.get_binding_dtype(binding)) host_mem cuda.pagelocked_empty(size, dtype) device_mem cuda.mem_alloc(host_mem.nbytes) bindings.append(int(device_mem)) if self.engine.binding_is_input(binding): inputs.append({host: host_mem, device: device_mem}) else: outputs.append({host: host_mem, device: device_mem}) return inputs, outputs, bindings, stream def infer(self, input_data): # input_data: (1, 32, 3, 224, 224) numpy array np.copyto(self.inputs[0][host], input_data.ravel()) cuda.memcpy_htod_async(self.inputs[0][device], self.inputs[0][host], self.stream) self.context.execute_async_v2(self.bindings, self.stream.handle) cuda.memcpy_dtoh_async(self.outputs[0][host], self.outputs[0][device], self.stream) self.stream.synchronize() return self.outputs[0][host].reshape(1, -1) # (1, 8) # 使用示例 trt_engine TRTInference(models/gesture.trt) cap cv2.VideoCapture(0) frame_buffer [] # 存储最近32帧 while True: ret, frame cap.read() if not ret: break frame cv2.resize(frame, (224, 224)) frame frame.transpose(2,0,1)[None] # (1,3,224,224) frame_buffer.append(frame) if len(frame_buffer) 32: frame_buffer.pop(0) if len(frame_buffer) 32: seq np.concatenate(frame_buffer, axis0)[None] # (1,32,3,224,224) pred trt_engine.infer(seq) gesture_id np.argmax(pred) print(fGesture: {constants.GESTURE_MAP[gesture_id]})关键细节frame_buffer维护滑动窗口确保每次输入严格 32 帧这是 TRT profile 的硬性要求transpose(2,0,1)[None]将(H,W,C)→(1,C,H,W)符合 TRT 输入格式context.execute_async_v2是异步执行配合stream.synchronize()保证时序实测比同步执行快 23%。5. 避坑指南训练/推理中 5 个真实翻车现场与后悔药这些坑我都亲手踩过有些导致重训 3 天有些让模型在测试集上掉点 8%列在这里省得你再走一遍弯路。5.1 现象训练 loss 从第1 epoch 就卡在 2.1 不下降原因pgdataset/annotations/train.json中gesture_class字段用了中文名如左转而非数字 ID3而GestureDataset.__getitem__()里ann[gesture_class]直接转int()报错但被try-except吞掉返回默认0导致所有样本标签都是0。解决运行python train/check_dataset.py --strict加--strict参数会强制校验类型修复 JSON 中所有gesture_class为整数或在constants.py中添加GESTURE_MAP_CHN_TO_INT {直行:0, 停止:1, ...}并修改数据加载逻辑。5.2 现象ONNX 模型在 TensorRT 中报错Assertion failed: scales.size() 4原因PyTorch 1.13 版本中torch.nn.functional.interpolate默认 mode 从bilinear改为nearest而 ResNet18 的 AdaptiveAvgPool2d 层在导出时触发了 interpolateONNX 不兼容新行为。解决在models/gesture_classifier.py的__init__中将self.cnn.avgpool nn.AdaptiveAvgPool2d((1,1))替换为self.cnn.avgpool nn.Sequential( nn.AdaptiveAvgPool2d((1,1)), nn.Flatten(1) # 显式 flatten避免 interpolate )5.3 现象Jetson Nano 上推理卡顿CPU 占用 95%原因OpenCV 的cv2.VideoCapture(0)默认使用 V4L2 后端但 Nano 的 CSI 摄像头需用cv2.CAP_GSTREAMER后端并指定 pipeline。解决替换cap cv2.VideoCapture(0)为gst_str (nvarguscamerasrc ! video/x-raw(memory:NVMM), width(int)1280, height(int)720, format(string)NV12, framerate(fraction)30/1 ! nvvidconv ! video/x-raw, format(string)BGRx ! videoconvert ! video/x-raw, format(string)BGR ! appsink) cap cv2.VideoCapture(gst_str, cv2.CAP_GSTREAMER)5.4 现象验证集 mAP 高达 95%但实拍视频准确率仅 63%原因pgdataset/videos/test_videos/中的测试视频是白天晴天场景而你的实拍是阴天逆光light_condition标签未参与训练模型未学习光照不变性。解决在train/config.yaml中启用 domain adaptationDATA: use_light_condition: true # 开启光照条件辅助标签 light_weight: 0.2 # 光照分类损失权重并在LSTMGestureClassifier.forward()中添加光照分类分支联合优化。5.5 现象pred/inference_trt.py运行时报错CUDA_ERROR_OUT_OF_MEMORY原因TRT 引擎构建时max_workspace_size设为1324GB但 Nano 仅 2GB 显存且系统占用约 0.5GB。解决重新构建引擎将max_workspace_size改为129512MB并确保nvidia-smi显示显存占用 1.2GB 再运行推理脚本。6. 模型轻量化实战把 87MB 的 .pth 压到 15MB精度只掉 0.7%毕业答辩时老师问“你这模型能跑在树莓派上吗”——别慌这不是玄学是可量化的工程动作。我用项目自带的models/prune_gesture.py3 步就把 ResNet18 backbone 剪掉 42% 参数同时保持 92.3% → 91.6% mAP这才是真正能交付的压缩。6.1 通道剪枝不是删 layer而是砍 feature map 维度prune_gesture.py基于torch.nn.utils.prune.l1_unstructured但做了关键改造只剪卷积层的输出通道out_channels且按每层 L1 norm 排序后剪 30%# models/prune_gesture.py def prune_resnet18(model, amount0.3): for name, module in model.named_modules(): if isinstance(module, nn.Conv2d) and layer in name: # 只剪 backbone 卷积层 # 计算每个输出通道的 L1 norm对 weight 的 dim[1,2,3] 求和 l1_norm torch.norm(module.weight.data, p1, dim[1,2,3]) # 获取要剪枝的通道索引norm 最小的 30% num_prune int(l1_norm.numel() * amount) _, indices torch.topk(l1_norm, knum_prune, largestFalse) # 创建 mask保留的通道设为1剪掉的设为0 mask torch.ones(module.out_channels, devicemodule.weight.device) mask[indices] 0 # 应用结构化剪枝删除整个通道 prune.CustomFromMask.apply(module, weight, maskmask.unsqueeze(1).unsqueeze(2).unsqueeze(3)) return model为什么有效ResNet18 的layer4最后一层有 512 个通道剪掉 30% 即 154 个直接减少后续全连接层输入维度比全局 unstructured 剪枝节省显存 37%。6.2 量化感知训练QAT用 fake quant 模拟 INT8避免部署后精度崩塌纯训练后量化PTQ会让精度掉 5%。本项目在train/qat_train.py中实现 QAT本文还有配套的精品资源点击获取