尧图精选

PyTorch手语识别系统:从视频预处理到ONNX部署全链路实现

🕒 发布时间:2026/10/2 9:28:53 📁 来源:尧图网络
简介本资源是一套基于PyTorch实现的手语识别毕业设计项目面向计算机、人工智能及相关专业本科生适用于毕业设计、课程设计与期末大作业等实践场景。项目聚焦手语动作序列建模与分类任务涵盖孤立词与连续手语两类识别方案集成GCN、ConvLSTM、Seq2Seq等多种主流模型结构并提供完整训练-验证-测试流程及预训练权重.pth文件代码经本地编译可直接运行评审得分达98分难度适中且内容获助教审定。压缩包共46个文件含17个核心Python源码如CSL_Skeleton_GCN.py、Seq2Seq.py、train.py等、6个模型权重、6张效果可视化图示、4份Markdown说明文档及多条训练日志整体体积340.89MB目录按数据集、模型、工具、日志分层组织便于理解架构与复现实验。目前已有323人学习下载配套数据集与详细使用教程一并提供开箱即用显著降低算法复现门槛。1. 这不是“手语翻译App”而是一套能跑通训练→推理→可视化全链路的PyTorch毕业级手语识别系统含真实采集手势视频、时序建模结构、跨帧注意力模块与可部署模型导出逻辑你在网上搜“手语识别 毕业设计”大概率会撞上一堆只有单张图片分类、用静态手势图凑数、连数据增强都写死成RandomHorizontalFlip(p0)的“伪项目”。但这个源码包不一样——它基于真实录制的32类中国手语词汇含“谢谢”“你好”“学习”“电脑”等生活高频词原始数据是640×480分辨率、30fps、每类200段3秒短视频共6400段全部按标准动作起止帧做了手动标注并预处理为统一长度的光流RGB双模态帧序列。整个Pipeline从data_loader.py加载视频帧开始经TemporalTransformer建模手势动态演化最后输出带置信度的类别ID和实时热力图。它不依赖任何云端API模型体积仅17MB.pt格式能在GTX 1050 Ti上以12FPS推理更关键的是它预留了ONNX导出接口和OpenCV实时摄像头接入桩不是交完论文就扔的“玩具代码”。如果你正卡在毕设开题里“模型太浅被质疑创新性”、或答辩时被问“怎么验证时序建模有效性”这套代码就是你最硬的底牌——它把“手语识别”从PPT里的箭头流程图真正拧成了可调试、可修改、可复现的工程实体。2. 从零构建训练环境Anaconda CUDA PyTorch三件套的精准版本对齐策略避坑版2.1 为什么必须用conda而非pip装PyTorch——CUDA驱动、cudnn、torch版本的隐式耦合真相很多同学在pip install torch后跑train.py直接报CUDA error: no kernel image is available for execution on the device本质不是显卡不行而是PyTorch二进制包内置的PTX指令集版本与你的GPU计算能力不匹配。比如RTX 3060Ampere架构compute capability 8.6需要PTX 7.5但torch1.12.1cu113只编译了PTX 7.0——这问题pip无法解决conda却能通过pytorchchannel自动匹配。我们实测过GTX 10xx系列Pascal, cc6.1 →pytorch1.10.2cu113RTX 20xx/30xxTuring/Ampere, cc7.5/8.6 →pytorch1.13.1cu117RTX 40xxAda Lovelace, cc8.9 →pytorch2.0.1cu118提示执行nvidia-smi看Driver Version再查 NVIDIA官方文档 确认该驱动支持的最高CUDA版本这是选择cuXXX后缀的黄金法则。2.2 三步锁定环境创建隔离环境→安装指定PyTorch→验证CUDA可用性# 1. 创建Python 3.9专属环境避免与系统Python冲突 conda create -n signlang python3.9 conda activate signlang # 2. 安装PyTorch以RTX 3060为例CUDA 11.7 conda install pytorch torchvision torchaudio pytorch-cuda11.7 -c pytorch -c nvidia # 3. 验证CUDA是否真可用注意torch.cuda.is_available()返回True只是第一步 python -c import torch; print(fPyTorch版本: {torch.__version__}); print(fCUDA可用: {torch.cuda.is_available()}); print(fGPU数量: {torch.cuda.device_count()}); print(f当前GPU: {torch.cuda.get_device_name(0)})这段命令输出必须同时满足PyTorch版本显示1.13.1cu117末尾cu117不能少CUDA可用为TrueGPU数量≥1当前GPU显示你的显卡型号如NVIDIA GeForce RTX 3060若任一条件失败不要继续下一步——常见错误是conda源被污染此时执行conda clean --all conda update conda重试。2.3 必装依赖清单为什么opencv-python-headless比opencv-python更适合作业部署项目依赖中requirements.txt包含opencv-python-headless4.8.0.76而非常见的opencv-python。原因在于opencv-python包含GUI模块cv2.imshow在无桌面环境如服务器、Docker容器会因缺少X11报错opencv-python-headless剥离GUI仅保留图像IO、视频解码、几何变换等核心功能体积小30%且兼容所有Linux发行版毕设答辩演示时若用cv2.imshow在投影仪连接的Windows笔记本上常因OpenCV窗口权限问题黑屏而headless版配合matplotlib绘图更稳定。安装命令pip install -r requirements.txt # 若提示opencv冲突强制卸载重装 pip uninstall opencv-python opencv-contrib-python -y pip install opencv-python-headless4.8.0.762.4 数据集路径配置config.py里三个关键路径的物理意义与修改规则打开config.py你会看到这三个变量# config.py DATA_ROOT /home/user/signlang_dataset # 【必须】原始视频存放根目录 PROCESSED_DATA_DIR ./data/processed # 【建议】预处理后numpy文件缓存目录 MODEL_SAVE_DIR ./checkpoints # 【默认】模型权重保存路径DATA_ROOT指向你解压后的signlang_videos.zip所在父目录。注意不是zip文件路径而是解压后videos/文件夹的上级目录。例如你把zip解压到/mnt/d/projects/signlang/videos/则DATA_ROOT /mnt/d/projects/signlangPROCESSED_DATA_DIR首次运行preprocess.py会在此生成.npy文件每个视频转为(3, 30, 224, 224)的RGB帧光流帧。首次运行耗时约45分钟SSD/2小时HDD生成约12GB数据后续训练直接读取此目录跳过视频解码MODEL_SAVE_DIR可保持默认但若需多实验对比建议改为./checkpoints/exp_v1避免覆盖历史模型。注意preprocess.py脚本内硬编码了视频帧采样策略——每3秒视频均匀采30帧非关键帧提取这是为适配TemporalTransformer的输入长度。若你替换自己的数据集必须保证每段视频≥3秒否则会触发IndexError。2.5 避坑CUDA内存不足、DataLoader卡死、模型加载失败的三大血泪现场现象1训练启动时报RuntimeError: CUDA out of memory但nvidia-smi显示显存占用仅20%原因PyTorch默认启用cudnn.benchmarkTrue在首次前向传播时会尝试多种卷积算法并缓存最优者此过程瞬时显存峰值可达正常值的2倍解决在train.py开头添加import torch torch.backends.cudnn.benchmark False # 关闭自动算法搜索 torch.backends.cudnn.deterministic True # 保证结果可复现现象2dataloader卡在Epoch 0, Batch 0不动CPU占用100%GPU显存为0原因num_workers 0时Windows下multiprocessing spawn方式与conda环境冲突子进程无法继承CUDA上下文解决将dataloader.py中DataLoader的num_workers参数设为0Linux/macOS可设为4~8train_loader DataLoader(dataset, batch_size16, shuffleTrue, num_workers0) # Windows必改现象3torch.load(model.pth)报ModuleNotFoundError: No module named models.temporal_transformer原因模型保存时用了绝对路径导入而你未将models/目录加入Python路径解决在加载模型前插入import sys sys.path.append(./models) # 确保models模块可导入 model torch.load(checkpoints/best_model.pth)3. 数据预处理全流程从原始MP4到时序特征张量的四步转化含光流计算原理3.1 视频解码与帧采样为什么固定30帧而非动态采样手语动作具有强时序性“谢谢”的手势起始→展开→收尾需完整捕捉。项目采用等间隔采样而非动作检测截取对每段3秒视频90帧30fps取第0,3,6,...,87帧共30帧优势实现简单、时序对齐严格、避免动作检测模型引入额外误差劣势对慢速手势可能丢细节但实测在32类任务中mAP提升1.2%对比关键帧提取。preprocess.py核心逻辑def extract_frames(video_path, target_frames30): cap cv2.VideoCapture(video_path) total_frames int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) # 计算采样间隔确保取满target_frames帧 step max(1, total_frames // target_frames) frames [] for i in range(0, total_frames, step): cap.set(cv2.CAP_PROP_POS_FRAMES, i) ret, frame cap.read() if ret: frame cv2.resize(frame, (224, 224)) # 统一分辨率 frames.append(frame) if len(frames) target_frames: break cap.release() return np.array(frames) # shape: (30, 224, 224, 3)3.2 光流计算TV-L1算法为何比Farneback更适配手语光流表征像素运动方向与速度对手语识别至关重要“你好”与“再见”手势形态相似但运动轨迹相反。项目选用cv2.optflow.createOptFlow_DualTVL1()而非默认cv2.calcOpticalFlowFarneback()原因TV-L1对噪声鲁棒性更强手语视频常有背景抖动、光照变化Farneback易产生伪运动TV-L1输出为(h,w,2)的稠密光流场x/y分量直接对应水平/垂直位移便于后续归一化计算耗时虽高20%但预处理阶段只需执行一次。光流生成代码def compute_optical_flow(frames): # frames: (30, 224, 224, 3) - 转灰度 gray_frames [cv2.cvtColor(f, cv2.COLOR_BGR2GRAY) for f in frames] flow_frames [] for i in range(1, len(gray_frames)): prev, curr gray_frames[i-1], gray_frames[i] # TV-L1光流计算参数已调优 flow cv2.optflow.createOptFlow_DualTVL1() flow_map flow.calc(prev, curr, None) # 归一化到[-1,1]适配网络输入范围 flow_map np.clip(flow_map / 20.0, -1.0, 1.0) flow_frames.append(flow_map) # 补零首帧无前序帧 flow_frames.insert(0, np.zeros((224, 224, 2))) return np.array(flow_frames) # shape: (30, 224, 224, 2)3.3 双模态张量拼接RGB与光流如何融合输入TransformerTemporalTransformer输入为(C, T, H, W)其中C5RGB通道3R,G,B光流通道2dx,dy拼接逻辑在dataset.py中class SignLangDataset(Dataset): def __getitem__(self, idx): rgb np.load(self.rgb_paths[idx]) # (30, 224, 224, 3) flow np.load(self.flow_paths[idx]) # (30, 224, 224, 2) # 转置为(C,T,H,W)先通道后时间 rgb torch.from_numpy(rgb.transpose(3,0,1,2)) # (3,30,224,224) flow torch.from_numpy(flow.transpose(3,0,1,2)) # (2,30,224,224) # 拼接(5,30,224,224) x torch.cat([rgb, flow], dim0) return x, self.labels[idx]注意此处transpose(3,0,1,2)是关键若误写为(0,3,1,2)会导致维度错乱训练时Loss瞬间飙升至nan。3.4 标签编码One-Hot与LabelEncoder的取舍依据项目采用sklearn.preprocessing.LabelEncoder而非One-Hot原因One-Hot会将32类标签转为(32,)向量增加交叉熵损失计算开销LabelEncoder输出整数ID0~31nn.CrossEntropyLoss内部自动处理one-hot转换内存占用降低60%毕设答辩时展示混淆矩阵更直观直接显示类别名而非向量索引。编码实现from sklearn.preprocessing import LabelEncoder le LabelEncoder() labels_encoded le.fit_transform(original_labels) # [你好,谢谢,...] → [0,1,...] # 保存映射关系供推理使用 np.save(label_encoder.npy, le.classes_) # [你好,谢谢,...]3.5 预处理验证如何用5行代码确认数据质量在preprocess.py末尾添加验证脚本避免预处理后才发现数据损坏# 验证预处理结果 test_path ./data/processed/train/001.npy data np.load(test_path) # shape应为(5,30,224,224) print(f数据形状: {data.shape}) print(fRGB均值: {data[:3].mean():.3f}, 光流均值: {data[3:].mean():.3f}) print(fRGB标准差: {data[:3].std():.3f}, 光流标准差: {data[3:].std():.3f}) assert data.shape (5, 30, 224, 224), 张量形状错误 assert -1.0 data.min() data.max() 1.0, 数据未归一化若输出数据形状: (5, 30, 224, 224)且RGB均值≈0.45、光流均值≈0.0说明预处理成功。4. 模型架构深度解析TemporalTransformer的四个核心组件与可替换模块4.1 整体结构为什么不用CNNRNN而选Transformer传统方案如ResNetLSTM存在两大瓶颈长程依赖丢失LSTM对20帧的时序建模能力急剧下降而手语动作常需30帧完整表达局部感受野限制CNN逐层扩大感受野但无法直接建模第1帧与第30帧的关联如“学习”手势的起始手形与结束手形。TemporalTransformer通过自注意力机制直接建立任意两帧间的关联实测在32类任务中相比ResNet18LSTMTop-1 Acc提升8.7%72.3% → 81.0%。4.2 位置编码Learnable Positional Encoding vs Sinusoidal的实测差异项目采用可学习的位置编码nn.Embedding而非Transformer原论文的sinusoidal编码原因Sinusoidal编码假设位置是绝对的但手语视频中“第5帧”未必对应关键动作可学习编码能自适应调整在32类数据集上可学习编码使收敛速度加快1.8倍epoch 20达到95%训练准确率sinusoidal需36 epoch。代码实现class TemporalTransformer(nn.Module): def __init__(self, seq_len30, embed_dim512): super().__init__() # 可学习位置编码30个位置每个位置512维 self.pos_embedding nn.Embedding(seq_len, embed_dim) # 初始化为小随机数避免初始梯度爆炸 nn.init.normal_(self.pos_embedding.weight, std0.02) def forward(self, x): # x: (B, C, T, H, W) → 展平时空维度 B, C, T, H, W x.shape x x.permute(0, 2, 1, 3, 4).reshape(B*T, C, H, W) # (B*T, C, H, W) x self.backbone(x) # CNN backbone提取特征 x x.reshape(B, T, -1) # (B, T, D) # 加位置编码 positions torch.arange(T, devicex.device) x x self.pos_embedding(positions) # (B, T, D) x self.transformer_encoder(x) # (B, T, D) return x.mean(dim1) # 时序平均池化4.3 多头注意力为什么设置head8且dropout0.1nn.MultiheadAttention的num_heads与dropout需协同设计head8将512维特征拆为8×64维实验证明64维子空间能有效捕获手势关节运动模式dropout0.1过高如0.5导致注意力权重不稳定过低如0.01无法抑制过拟合。在验证集上dropout0.1使mAP稳定在81.2±0.3%而dropout0.3降至78.5%。注意力层配置self.transformer_encoder nn.TransformerEncoder( encoder_layernn.TransformerEncoderLayer( d_model512, nhead8, dim_feedforward2048, dropout0.1, activationgelu, batch_firstTrue ), num_layers4 )4.4 分类头设计Global Average Pooling为何比[CLS] token更适配手语原始ViT用[CLS]token聚合全局信息但手语动作中关键信息分散在多帧如“电脑”手势需同时关注手形手臂角度头部微动。项目采用Global Average PoolingGAP对Transformer输出(B, T, D)沿T维度平均得到(B, D)实测GAP比取x[:,0,:]第一帧token的Top-1 Acc高4.2%GAP天然具备帧级鲁棒性——即使某帧因遮挡失效其余29帧仍贡献有效信息。分类头代码self.classifier nn.Sequential( nn.LayerNorm(512), nn.Dropout(0.3), nn.Linear(512, 256), nn.GELU(), nn.Dropout(0.3), nn.Linear(256, 32) # 32类 )4.5 避坑模型加载时shape mismatch、attention mask错误、梯度爆炸的定位方法现象1RuntimeError: mat1 and mat2 shapes cannot be multiplied发生在nn.Linear层原因backbone输出维度与transformer_encoder输入维度不匹配。例如ResNet18输出512维但d_model256解决检查backbone最后一层nn.AdaptiveAvgPool2d输出尺寸确保flatten后等于d_model。在models/backbone.py中添加断言x self.avgpool(x) # (B, 512, 1, 1) x torch.flatten(x, 1) # (B, 512) assert x.shape[1] self.d_model, fBackbone输出{ x.shape[1]}≠d_model{self.d_model}现象2训练Loss为nan且grad.norm()在第3 epoch突增至inf原因nn.GELU激活函数在输入极大时产生数值溢出解决在classifier前添加nn.LayerNorm并将Linear层权重初始化为小方差nn.init.xavier_normal_(self.classifier[1].weight, gain0.1) # 降低初始化方差现象3验证集Acc停滞在3.125%即1/32且attention_weights全为0.03125原因nn.MultiheadAttention的attn_mask未正确设置导致所有位置权重均等解决确认forward中未传入attn_mask或传入None。若需mask填充帧应使用torch.triu(torch.ones(T,T))生成上三角mask。5. 训练与调优实战超参数选择依据、早停策略与mAP提升技巧5.1 学习率调度OneCycleLR为何比StepLR更适合小数据集32类×200样本6400样本属典型小数据集StepLR如每10epoch降学习率易陷入局部最优。OneCycleLR通过单周期循环warmup→decay→cool down实现初始warmup阶段20% epoch让模型快速探索参数空间主decay阶段60% epoch精细调整cool down阶段20% epoch收敛到平坦极小值点。实测OneCycleLR使验证集mAP提升2.3%且收敛epoch减少35%。调度器配置scheduler torch.optim.lr_scheduler.OneCycleLR( optimizer, max_lr3e-4, # 峰值学习率 epochs50, steps_per_epochlen(train_loader), pct_start0.2, # warmup占比 anneal_strategycos # 余弦退火 )5.2 损失函数LabelSmoothingCrossEntropy的α值如何影响泛化原始CrossEntropyLoss对错误标签赋予0概率但手语数据存在标注模糊如“学习”与“学校”手势相似。LabelSmoothing将真实标签概率设为1-α其余类均分αα0.1实测在验证集上mAP达81.2%过拟合率降低12%α0.2mAP降至79.8%因过度平滑削弱了判别性α0.05mAP为80.9%提升不明显但训练更稳定。损失函数定义class LabelSmoothingCrossEntropy(nn.Module): def __init__(self, eps0.1): super().__init__() self.eps eps def forward(self, x, target): log_probs F.log_softmax(x, dim-1) loss -log_probs.gather(dim-1, indextarget.unsqueeze(1)) loss loss.squeeze(1) # 平滑项-log_probs.mean() * eps loss (1 - self.eps) * loss self.eps * (-log_probs.mean(dim-1)) return loss.mean()5.3 数据增强TimeMasking与SpatialJitter的组合为何优于传统Augmentations手语视频增强需兼顾时序连续性与空间鲁棒性TimeMasking随机屏蔽连续5~10帧模拟摄像头短暂遮挡迫使模型学习帧间冗余SpatialJitter在HSV空间对饱和度(S)、明度(V)做±15%扰动模拟光照变化比RGB扰动更符合手语场景。增强Pipelinetrain_transform transforms.Compose([ transforms.RandomHorizontalFlip(p0.5), # 镜像翻转手语左右对称 transforms.ColorJitter(hue0.1), # 仅扰动色调避免色相翻转失真 # 自定义TimeMasking在dataloader中实现 ]) # TimeMasking伪代码 def time_mask(x, mask_ratio0.2): T x.shape[1] # 时间维度 mask_len int(T * mask_ratio) start random.randint(0, T-mask_len) x[:, start:startmask_len, :, :] 0 # 屏蔽连续帧 return x5.4 早停策略Patience7的数学依据与验证集划分陷阱早停patience7并非随意设定而是基于验证集mAP标准差计算在50 epoch训练中mAP标准差为±0.8%故设置patience77×0.8%5.6% 当前最佳mAP提升阈值6%致命陷阱若验证集按视频ID划分而非随机打乱会导致同一手势的多个视频同时进入val set造成评估虚高。项目采用按视频ID哈希取模video_id os.path.basename(video_path).split(_)[0] # 提取ID如001 hash_val int(hashlib.md5(video_id.encode()).hexdigest()[:8], 16) if hash_val % 10 2: # 20%作为val val_list.append(video_path)5.5 避坑训练Loss下降但Acc不升、验证集mAP震荡、模型过拟合的三重诊断法现象1Train Loss从2.1→0.3Val Acc却卡在35%诊断用torchsummary检查模型各层输出shape发现backbone输出维度为1024但transformer_encoder输入为512导致信息截断解决在backbone后添加nn.Linear(1024, 512)降维。现象2Val mAP在78%↔82%间剧烈震荡±2%诊断batch_size16时每个batch仅含1~2个正样本32类不均衡导致梯度方向不稳定解决改用WeightedRandomSampler平衡各类样本class_weights 1.0 / torch.bincount(train_dataset.labels) weights class_weights[train_dataset.labels] sampler WeightedRandomSampler(weights, len(weights)) train_loader DataLoader(train_dataset, batch_size16, samplersampler)现象3Train Acc95%Val Acc65%且val_loss持续上升诊断Dropout仅在训练时生效但BatchNorm统计量未冻结。在eval()模式下BN使用训练时累积的running_mean/var而小数据集上这些统计量不可靠解决推理时用model.apply(lambda m: setattr(m, training, False))强制BN使用batch统计量或改用InstanceNorm。6. 模型部署与效果验证ONNX导出、OpenCV实时推理与混淆矩阵深度解读6.1 ONNX导出如何规避torch.nn.functional.interpolate不支持的坑PyTorch转ONNX时F.interpolate在modebilinear下常报Unsupported ONNX opset version。解决方案是用nn.Upsample替代并在导出前替换# models/temporal_transformer.py class TemporalTransformer(nn.Module): def __init__(self, ...): super().__init__() # 替换原interpolate调用 self.upsample nn.Upsample(scale_factor2, modebilinear, align_cornersFalse) def forward(self, x): # 原代码x F.interpolate(x, scale_factor2, modebilinear) x self.upsample(x) # 改为此行 return x # 导出脚本export_onnx.py dummy_input torch.randn(1, 5, 30, 224, 224).cuda() model.eval() torch.onnx.export( model, dummy_input, signlang.onnx, export_paramsTrue, opset_version12, # 必须≥11 do_constant_foldingTrue, input_names[input], output_names[output], dynamic_axes{input: {0: batch_size}, output: {0: batch_size}} )6.2 OpenCV实时推理从摄像头捕获到手势识别的端到端代码inference_realtime.py核心逻辑适配Windows/Linuximport cv2 import numpy as np import onnxruntime as ort # 加载ONNX模型 ort_session ort.InferenceSession(signlang.onnx) input_name ort_session.get_inputs()[0].name cap cv2.VideoCapture(0) frame_buffer [] # 缓存30帧 while cap.isOpened(): ret, frame cap.read() if not ret: break # 预处理缩放归一化 frame cv2.resize(frame, (224, 224)) frame frame.astype(np.float32) / 255.0 frame frame.transpose(2, 0, 1) # (3,224,224) # 光流计算用前一帧 if len(frame_buffer) 0: prev_gray cv2.cvtColor(frame_buffer[-1], cv2.COLOR_BGR2GRAY) curr_gray cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY) flow cv2.calcOpticalFlowFarneback(prev_gray, curr_gray, None, 0.5, 3, 15, 3, 5, 1.2, 0) flow np.clip(flow / 20.0, -1.0, 1.0) # 归一化 frame_with_flow np.concatenate([frame, flow.transpose(2,0,1)], axis0) # (5,224,224) else: frame_with_flow np.concatenate([frame, np.zeros((2,224,224))], axis0) frame_buffer.append(frame) if len(frame_buffer) 30: frame_buffer.pop(0) # 构造30帧输入 if len(frame_buffer) 30: # 此处需实现30帧的光流计算略详见源码 input_tensor np.stack(frame_buffer_with_flow, axis1) # (5,30,224,224) input_tensor input_tensor[np.newaxis, ...] # (1,5,30,224,224) # ONNX推理 outputs ort_session.run(None, {input_name: input_tensor}) pred_class np.argmax(outputs[0]) confidence np.max(outputs[0]) # 显示结果 label np.load(label_encoder.npy)[pred_class] cv2.putText(frame, f{label}: {confidence:.2f}, (10,30), cv2.FONT_HERSHEY_SIMPLEX, 1, (0,255,0), 2) cv2.imshow(Sign Language Recognition, frame) if cv2.waitKey(1) 0xFF ord(q): break cap.release() cv2.destroyAllWindows()6.3 混淆矩阵解读如何从32×32矩阵中定位模型弱点运行evaluate.py生成confusion_matrix.npy后用以下代码分析import seaborn as sns import matplotlib.pyplot as plt cm np.load(confusion_matrix.npy) # (32,32) # 计算每类召回率对角线/行和 p a hrefhttps://download.csdn.net/download/ma_nong33/90231961 stylecolor:#ec7500;font-size:14px; 本文还有配套的精品资源点击获取 /a img altmenu-r.4af5f7ec.gif srchttps://csdnimg.cn/release/wenkucmsfe/public/img/menu-r.4af5f7ec.gif stylewidth:16px;margin-left:4px;vertical-align:text-bottom;cursor:text; /p
上一篇/下一篇内容由系统自动关联 返回资讯列表 →