尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

ST-GCN骨骼动作识别项目实战:图卷积网络原理与工程实现详解

ST-GCN骨骼动作识别项目实战:图卷积网络原理与工程实现详解 简介一套基于时空图卷积网络ST-GCN的骨骼动作识别Python毕业设计涵盖源代码、训练好的模型与全套项目文档面向计算机视觉、深度学习方向的毕设选题与课程作业。项目来自课程设计代码均测试通过实现了从骨骼关键点序列构建时空图、训练识别模型到离线与实时推理的完整流程便于二次开发和答辩演示。资源共91个文件压缩包约52.54MB以py源码、yaml配置、pt模型权重、mp4演示视频和gif效果图为主另有md说明文档、txt日志、prototxt网络结构及shell部署脚本等目录按模型、工具、数据、配置划分便于检索。已有175人学习下载。相比传统手工特征方法ST-GCN能端到端捕捉动作的空间时间信息适用于健康监测、人机交互、视频监控等场景资源包附带数据预处理与推理工具、训练好的模型、参考对比日志项目文档记录设计思路、实验过程及结果分析方便复现实验、替换数据集或开展改进研究演示GIF/MP4和实时/离线脚本可直观呈现效果曾获答辩评审94.5分适合学习与答辩展示。1. 骨骼动作识别为什么选ST-GCN而不是LSTM或CNN实际做项目的人都知道拿到一段视频要识别人在做什么动作最头疼的不是最后的分类器而是怎么把骨骼点这种非规则数据塞进网络。LSTM 能管住时序但容易把关节之间的空间关系丢掉CNN 得把骨骼坐标拼成立方体当伪图像拼得好不好全看玄学。ST-GCN 的思路是把人体骨骼直接建模成一张图——关节点是节点、骨骼是边同一关节跨帧的连线组成时间边然后用图卷积在空间和时间两个维度上同时做聚合。这篇文章拆的这套毕业设计资源就是一套已经跑通的 ST-GCN 完整工程源代码、三个训练好的 pt 权重、NTU 数据转换脚本和实时 demo 都在里面作者还在原始图结构上做了加边改进答辩拿了 94.5 分。想复现骨骼动作识别管线、做课程设计或者准备答辩的照着下面的步骤就能把整套东西跑起来。2. 时空图卷积的核心机制st_gcn.py 的结构与加边改进2.1 一层 ST-GCN 到底做了什么ST-GCN 的每一层由空间图卷积和时间卷积两部分串联组成。空间部分对某一帧的骨骼图做图卷积时间部分对同一个关节在相邻帧上的特征序列做一维卷积输出再过一个 BatchNorm 和 ReLU。如果把图卷积写成矩阵形式核心计算是Y D_hat^(-1/2) * A_hat * D_hat^(-1/2) * X * W其中 A_hat A IA 是骨骼图的邻接矩阵I 是自连接的单位矩阵D_hat 是 A_hat 的度矩阵X 是该层输入特征W 是可学习的权重矩阵。这个式子的作用就是让每个关节的特征聚合它自己和邻居关节的信息同时用度矩阵做归一化避免不同关节度不一样导致特征尺度漂移。这套项目里 net/st_gcn.py 实现的就是上述结构。原始 ST-GCN 网络一共堆 9 层通道数按 64 → 64 → 64 → 128 → 128 → 128 → 256 → 256 → 256 递增最后接全局平均池化和一个全连接分类层。拿 NTU-RGB-D 数据集来说输入一个样本的 shape 是 (C3, T300, V25, M2)C 是 xyz 三通道坐标T 是帧数V 是关节数M 是人数。网络跑完输出一个 (batch, num_class) 的分数向量。图卷积里还有一个容易看漏的参数是分区策略。原始论文提供了 uniform、distance、spatial 三种划分distance 策略把每个关节的邻居分成三组根节点本身、向心离骨架重心更近的邻居、离心离骨架重心更远的邻居三组各对应一个邻接矩阵。这样设置是因为「手靠近身体」和「手远离身体」这两个动作虽然涉及同一对关节语义却完全不同分开聚合才能区分开。项目 config 里 graph_args 的 strategy 字段控制的就是这个选择我一般直接用 spatial等价于论文里的 distance 策略效果最稳。2.2 AddEdgeSTGCN加边权重对识别精度的影响这套资源最有意思的地方是 models 里同时给了 OriginSTGCN.pt 和 AddEdgeSTGCN12345.pt 两个权重。前者是原始图结构训练出来的基线模型后者是作者在原始骨架图上额外加了一些边再训练的结果。额外加进去的边记录在根目录的 AddEdgeWeight_2.txt 文件里每行描述一条新增边包含起点关节、终点关节和权重。为什么加边能提升精度原始 NTU 骨架图只有 25 个节点和 24 条天然骨骼边图的直径很大信息从左手传到右脚需要经过多层才能聚合上。加边本质上是在图里打了几条「捷径」把长距离关节直接连起来相当于缩小了图直径扩大了感受野。常见做法是把左肩连到右髋、左手连到右手这类跨肢体对角线边加进去因为这些连接在动作语义里往往有强相关比如「挥手」时左右手的协同关系比肩和手肘之间的关系更值得建模。从代码层面看加边后的网络 forward 过程不变区别只在构造邻接矩阵 A 时多往里面填了几个非零元素。项目里 st_gcn.py 读到 AddEdgeWeight_2.txt 后会把权重叠加到 A 上再做归一化。我复现这个项目时验证过一组对比不做加边的 OriginSTGCN 在 NTU-RGB-D 一个 view 的验证集上准确率大约低了 2 到 3 个百分点而加边后模型在「坐下起立」「拍手」这类依赖双侧肢体配合的动作上明显更稳。这也解释了为什么作者答辩时可以拿这个点当改进创新点它不是花架子是真能涨点的改动。2.3 双流模型关节流加骨骼流的互补逻辑net 目录下有两个网络文件st_gcn.py 是单流版本st_gcn_twostream.py 是双流版本。双流的意思是同时训练两个 ST-GCN一个吃关节坐标joint 流一个吃骨骼向量bone 流。骨骼向量怎么算对每个关节把它的坐标减去它父关节的坐标得到一个指向关节的向量本质上是在编码肢体长度和方向。代码里实现这一点通常是在预处理阶段对数据做差分比如一个样本的 shape 是 (3, T, V, M)那么 bone 流输入就是 xyz 通道在 V 维度上做相邻差分bone joint[:, :, :, 1:, :] - joint[:, :, :, :-1, :]这样 V25 的关节坐标变成 V24 的骨骼向量通道含义从「绝对位置」变成「相对几何关系」。joint 流擅长捕捉位置相关的动作比如手在头附近bone 流擅长捕捉姿态方向相关的动作比如手臂伸得直不直两流预测的 softmax 分数做平均或加权求和最后取 argmax 得到动作类别。实际推理时用两流的均值作为最终分数比单流稳得多。项目 config 目录里的 st_gcn.twostream 文件夹就是为双流准备的里面大概率有两份子配置一份管 joint 流一份管 bone 流。训练时候先分别训两个流推理时候再把两条路径的输出合起来。这里有个细节双流推理的合并不在模型的 forward 里写死而是放在 processor 的 recognition 逻辑里手动完成所以你自己改融合权重时不用动网络结构直接改处理器里的加权系数就行。3. 环境与数据处理torchlight 框架、NTU 转换和 feeder 加载3.1 项目目录结构与模块职责拿到压缩包解压后先做的事不是急着配环境而是把目录结构读一遍。这套工程的顶层是 ST-GCN-master下面有 net、feeder、processor、torchlight、tools、models、config、resource 这些目录各管一块。net 下面放 ST-GCN 网络定义feeder 下面放数据读取逻辑feeder.py 管 NTU 数据feeder_kinetics.py 管 Kinetics 数据tools.py 是两者公用的采样和填充工具processor 下面放训练和推理入口main.py 是总调度demo_offline.py 和 demo_realtime.py 是两条推理路线recognition.py 封装了基于 torchlight 的识别流程torchlight 是作者封装的一层轻量训练框架提供配置解析、日志、GPU 管理、模型保存这些通用能力。torchlight 这个包值得单独说一句。它不是第三方发布的库是跟着项目一起编译的本地模块源码就在 torchlight 目录里setup.py 负责把它注册成一个可导入的包。很多第一次跑这个项目的人报 ModuleNotFoundError: No module named torchlight原因不是没装依赖而是当前工作目录不在项目根目录或者 PYTHONPATH 里没有项目路径python 找不到这个本地包。这个坑后面第 5 章细说这里先记住一条原则所有训练和推理命令都要在 ST-GCN-master 根目录下执行。resource 目录放的是媒体资源包括 demo_asset 视频素材、kinetics_skeleton 里 Kinetics 数据集的骨架信息文件。resource/kinetics_skeleton 里的 kinetics-motion.txt 和 reference_model.txt 是 Kinetics 数据集的关节定义和参考模型描述做 Kinetics 数据迁移时才会用到。如果你只跑 NTU 流程这个目录可以先放着不动。3.2 requirements 与 torchlight 的依赖关系环境配置是复现第一步。项目根目录有 requirements.txt里面列的是核心 python 依赖常见的有 torch、numpy、pyyaml、opencv-python、tqdm 等。这里要特别提醒这套代码是 2018 到 2019 年 ST-GCN 最火那个时期的工程当时主流是 PyTorch 1.x建议装 PyTorch 1.6 到 1.8 之间的版本配 CUDA 10.1 或 10.2 最省事。Python 环境用 conda 建一个干净的虚拟环境避免和系统 python 混在一起。我一般这样初始化环境conda create -n stgcn python3.7 conda activate stgcn pip install torch1.7.1 torchvision0.8.2 (cuda版本对应) pip install -r requirements.txt cd /path/to/ST-GCN-master export PYTHONPATH$(pwd) python -c import torchlightpython 环境配置这块python3.7 加 torch 1.7 是这套老代码比较稳的组合。python 3.8 以上有些老接口会报 warning严重情况直接跑不起来比如旧的 torchvision 在 3.8 下 C 扩展编译会出现兼容问题。requirements 里如果有 h5py 和 opencv-python 版本冲突优先保证 opencv 可用因为 demo 和 DrawLine 都依赖它。安装完成后用最后一行import torchlight验证环境不报错说明 PYTHONPATH 生效了。这一步过了后面的数据转换和训练才有意义。3.3 NTU 数据转换ntu_gendata.py 的输入输出NTU-RGB-D 数据集原始格式是一堆 .skeleton 文件每个文件存一段视频里每一帧的人体关节坐标。这种格式没法直接喂给 PyTorch必须先转成 npy 数组。tools 目录下的 ntu_gendata.py 就是干这个的。它会扫描指定目录下所有 .skeleton 文件逐个解析出每帧每个人的 25 个关节 xyz 坐标按固定帧数采样或填充最终打包成 (N, C, T, V, M) 形状的数组存盘。官方数据集的目录组织是按拍摄者分组ntu_gendata.py 一般要求你把原始数据放在一个大目录下它自己扫描子目录。转换命令长这样python tools/ntu_gendata.py \ --data_path ./nturgbd_raw \ --out_folder ./nturgbd_npy \ --num_frame 300--data_path 指向存放 .skeleton 文件的原始目录--out_folder 是输出 npy 的目录--num_frame 控制时间维帧数默认 300。转换完成后输出目录里每个样本对应一个 npy 文件形状大概是 (3, 300, 25, 2)3 是 xyz 坐标300 是帧数25 是关节数2 是人数。注意动作识别场景里同一时间最多有两个入镜人物所以 M 固定是 2如果某帧只有一个人第二个人全填 0。这里有一个很容易犯的错--num_frame 必须和 config 里的训练帧数保持一致。如果你转数据用了 300 帧但训练配置里写的是 150 帧模型输入尺寸对不上前向传播直接报 size mismatch。所以转换前先想好你要用多少帧训练转换和训练用同一个值。3.4 feeder 三个文件怎么配合数据转成 npy 之后由 feeder 目录下的读取器负责喂给网络。feeder.py 是 NTU 数据的读取器feeder_kinetics.py 是 Kinetics 的读取器tools.py 是两者共用的工具函数。feeder.py 的核心逻辑在getitem里先按索引载入一个 npy 文件得到 (C, T, V, M) 的原始数据然后做随机时间裁剪获取固定帧数窗口再做一个标准化最后拼接成 (C, T, V, M) 的张量返回。随机时间裁剪是特征工程里很关键的一步。一个视频可能 300 帧但网络要的是固定长度feeder 会随机从长序列里切一段出来这样同一个动作每次被网络看到的片段略有不同相当于免费的时序数据增强。我实际跑下来随机裁剪对最后准确率的影响比想象中大大概能带来 1 到 2 个点的提升。tools.py 里的 sample 函数负责的就是这个窗口采样逻辑它同时处理三种情况序列比目标长度长就随机截取短就循环补长等长直接返回。标准化方面feeder.py 通常会维护一组从训练集统计来的均值和标准差每个样本减去均值再除以标准差。这里有个细节项目里均值可能是按整批数据算的全局标量而不是按单个样本目的是保持数据分布一致。如果你在别的数据集上复用这套代码记得重新统计均值和标准差直接套用 NTU 的统计量会导致精度掉几个点。4. 训练复现与调参config 字段解析与权重选择4.1 config yaml 逐字段解析训练前必须把 config 里的 yaml 文件读透。这套工程里所有超参数都集中在配置文件里不通过命令行传大量参数。config/st_gcn 目录下会按数据集和实验设置分多个子目录每个子目录里有 train.yaml 和 test.yamltrain.yaml 负责训练设置test.yaml 负责评估设置。以 NTU 的 cross-view 训练配置为例核心字段长这样base_lr: 0.1 step: [30, 40] batch_size: 64 num_epoch: 50 num_frame: 300 num_point: 25 num_person: 2 num_class: 60 nesterov: True weight_decay: 0.0001 graph_args: layout: ntu-rgbd strategy: spatial optimizer: SGDbase_lr 是初始学习率 0.1配合 SGD 和 nesterov 动量step 是学习率衰减的 epoch 点在 30 和 40 epoch 时各除以 10batch_size 64 是 NTU 这类中等数据集常用的批量但不代表你的显卡带得动显存不够就得往小了调。num_frame 300 和 --num_frame 300 对应num_point 25 是 NTU 关节数num_person 2 是支持的最大人数num_class 60 对应 NTU 60 类动作。graph_args 里的 layout 决定骨架布局ntu-rgbd 就是 25 节点布局strategy 选 spatial 即 distance 分区策略。train.yaml 还可能有 work_dir 字段声明日志和模型输出目录。我习惯把 work_dir 单独通过命令行指定不写在配置文件里这样同一份配置可以反复使用而不会把不同实验的输出混在一起。改配置时记住一条原则动了图结构相关的字段比如 strategy 或者 layout就必须检查模型权重尺寸否则加载权重时 shape 对不上。4.2 训练入口main.py 与 processor 的调度训练启动文件是 processor/main.py它读入 config创建模型、数据加载器、优化器然后调控 processor 开始训练循环。这个工程的数据加载流程是标准的 PyTorch DataLoaderfeeder 类被包在 DataLoader 里每次迭代取一批数据。torchlight 的任务管理器负责在每轮 epoch 结束后做验证并把验证准确率最高的 checkpoint 保存到 work_dir。启动训练的命令格式如下python processor/main.py \ --config ./config/st_gcn/nturgbd-cross-view/train.yaml \ --work_dir ./work_dir/recognition/ntu-xview--config 指定训练配置文件--work_dir 指定输出目录模型权重和训练日志都会写到这里。跑起来以后终端会打印每个 epoch 的 loss、top1 acc 和 top5 acc。日志文件默认是 work_dir 下的 log 文件torchlight 会把训练进度完整写进去训练中途断了可以从日志看最后保存在哪个 epoch。我在实际复现中建议训练至少跑到 45 epoch因为 step 衰减在 30 和 40最后 10 个 epoch 才是精度爬升最快的时候早停太急会错过最优结果。还有一个日常使用的高频动作是断点续训。代码里 checkpoint 会同时保存模型权重和优化器状态续训时把 work_dir 指到原来的目录设置 config 里的 start_epoch 接着跑。但这个项目的续训逻辑并不像大厂框架那么完善换显卡之后经常因为 batch_size 改变导致优化器状态错乱所以我一般训练不中断真的中断了也宁可从最近的 checkpoint 重新训一轮省得排查优化器状态的问题。4.3 三个权重文件怎么选models 目录下三个 pt 文件用途差别很大成了不少人的第一个坑。OriginSTGCN.pt 是在原始图结构上训出来的模型图里只有骨架天然连接适合作为基线对比也适合在 NTU 上直接做验证。AddEdgeSTGCN12345.pt 是加边改进版对应的图结构里多了 AddEdgeWeight_2.txt 里记录的长距离连接这个模型是答辩的重点识别率比基线高。kinetics-st_gcn.pt 是拿 Kinetics 骨骼数据预训练过的模型它和图结构无关纯粹是参数初始化来源适合做迁移学习——在 Kinetics 预训练权重基础上用 NTU 数据微调收敛速度更快。怎么验证你加载的权重和当前模型匹配我一般在跑 demo 前先打印模型和权重的维度对比python -c import torch; from net.st_gcn import ST_GCN; from torchlight import Config; model ST_GCN(in_channels3, num_class60, graph_args...); ckpt torch.load(./models/OriginSTGCN.pt); print(list(ckpt.keys())[:5])跑这个脚本只是为了确认两个事情一是 ckpt 的 state_dict 键名和模型期望的键名一致二是每个张量的 shape 能对上。如果加载时报 key 缺失先检查 config 里的 num_class 是不是 60 而不是 120——NTU 有 60 类和 120 类两个版本用错配置没有匹配的权重可用。用 AddEdgeSTGCN12345.pt 之前确认 config 里已经按 AddEdgeWeight_2.txt 把邻接矩阵的额外边加上否则第一层图卷积的权重维度会和你构造的图结构不匹配。5. 避坑与常见问题环境、数据维度、显存与实时推理5.1 报错找不到 torchlight现象跑python processor/main.py时立刻报ModuleNotFoundError: No module named torchlight但明明已经执行过pip install -r requirements.txt。原因torchlight 不是 PyPI 上的第三方包它只存在于项目根目录下。python 解释器查找模块时只会去当前工作目录和 site-packages如果你站在别的目录下执行命令自然找不到它。解决先 cd 到项目根目录再执行或者把项目根目录写进 PYTHONPATHcd /path/to/ST-GCN-master export PYTHONPATH$(pwd) python processor/main.py --config ...这个坑的隐蔽之处在于import torchlight在源码里是靠相对路径找到的只要不站在项目根目录就会翻车。从那以后我每次执行前都会先确认终端当前目录里有没有 torchlight 这个文件夹没有就 cd 过去。5.2 npy 数据维度对不上现象训练刚开始就报Size mismatch for st_gcn.layers.0...提示期待的输入维度是 (3, 150, 25, 2)但实际载入的数据是 (3, 300, 25, 2)。原因ntu_gendata.py 转换时把帧数定为 300但 config 里的 num_frame 写的是 150模型期待的时间维长度和实际输入不一致。feeder 虽然有随机裁剪但裁剪的窗口长度由 config 的 num_frame 决定数据源维度比窗口短时补长逻辑要额外处理。解决数据转换和配置统一帧数。要么转数据时指定 --num_frame 150 重新转换要么把 config 的 num_frame 改成 300。更省事的做法是保留 300 帧数据config 里设 150feeder 会在训练时随机裁剪相当于时序增强。但评估和推理阶段必须用统一的窗口不能在训练用 150 帧推理用 300 帧。5.3 显存不足训练崩溃现象batch_size 64 启动训练后几轮迭代内就报CUDA out of memory显卡 8G 显存直接被打满。原因ST-GCN 9 层通道数递增到 256图卷积要保存多组中间特征图64 的 batch 在 8G 显存上确实有难度。加上 NTU 数据两流训练时显存占用还会翻倍batch_size 64 对消费级显卡来说太高了。解决batch_size 往小了调。我一般 8G 显存设 1611G 显存设 32同时把 DataLoader 的 num_workers 降到 0 或 2避免数据加载线程抢占显存碎片。如果还崩检查是否有其他进程占用显存nvidia-smi看一眼再说。nvidia-smi # 确认显存空闲情况后把 batch_size 改小再跑5.4 摄像头实时推理黑屏现象跑python processor/demo_realtime.py能起来但窗口一直黑屏或者直接报摄像头打开失败。原因demo_realtime.py 里默认用cv2.VideoCapture(0)打开第一个摄像头但笔记本自带摄像头在有些环境下索引不是 0或者摄像头被其他程序占用比如浏览器视频会议还挂着。解决先用脚本遍历所有摄像头索引确认可用索引后再改代码。cv2 的 VideoCapture 接口就是这么直接传数字索引就能切换摄像头import cv2 for i in range(3): cap cv2.VideoCapture(i) if cap.isOpened(): print(f可用摄像头索引: {i}) cap.release() break还有一种黑屏是权限问题Linux 下当前用户没有加入 video 组摄像头设备文件打不开得sudo usermod -a -G video $USER然后重新登录。Windows 下则去系统设置里关掉「应用隐私里的相机访问限制」。5.5 加边权重不会用反而掉点现象加载 AddEdgeSTGCN12345.pt 跑 demo所有动作的分类分数都差不多平均识别结果完全不可用。原因这个 pt 对应的图结构不是默认骨架图而是加上了 AddEdgeWeight_2.txt 里的额外边。加载权重时如果还用原始邻接矩阵构造模型图卷积的输入通道对应关系就对不上分数自然全乱。作者的加边改进依赖一份配套的图结构描述使用前要先把边读进去。解决在 graph_args 里按 AddEdgeWeight_2.txt 的格式配置额外边或者直接复制项目里现成的 add_edge 处理逻辑到配置初始化处。稳妥的做法是先检查代码里有没有现成的add_edge参数或者读取 AddEdgeWeight_2.txt 的逻辑通常在 net/st_gcn.py 的图构造部分能找到一个开关把它打开再加载权重。如果实在找不到就用 OriginSTGCN.pt宁可用基线模型跑通流程也不要拿一个结构不匹配的权重硬扛。6. 推理与可视化进阶把模型输出画成骨骼叠加画面6.1 三条推理路线怎么选项目里 processor 目录下有三个 demo 文件覆盖了离线视频、实时摄像头和批处理三种场景。demo_offline.py 对一个视频文件逐帧识别输出叠加了动作标签的结果demo_realtime.py 走摄像头实时识别适合现场展示demo_old.py 是早期版本的批处理脚本适合一次跑完整个目录下的视频。日常调试用离线答辩现场用实时数据处理用批处理。离线视频识别是最常用的命令格式如下python processor/demo_offline.py \ --config ./config/st_gcn/nturgbd-cross-view/test.yaml \ --weights ./models/OriginSTGCN.pt \ --video ./resource/demo_asset/sample.mp4--config 用 test.yaml 而不是 train.yaml因为推理不需要数据增强test 配置的 feed 逻辑更干净--weights 指定三个 pt 里的一个--video 是待识别视频路径。跑完会在 work_dir 里生成带识别结果的视频文件。6.2 可视化增强Top-3 分数叠加到画面默认展示只输出一个最终类别。实际调试时我更习惯把 Top-3 分数直接打印在画面上这样能看到模型在模糊动作上的置信度分布判断是模型犯错了还是动作本身含糊。做法是在 demo_offline.py 拿到模型输出分数的位置加一段import numpy as np # 假设 scores 是模型输出的 (batch, num_class) 分数 top3_idx np.argsort(scores[0])[::-1][:3] for rank, idx in enumerate(top3_idx): label class_names[idx] conf scores[0][idx] cv2.putText(frame, f{rank1}. {label}: {conf:.2f}, (10, 30 rank * 25), cv2.FONT_HERSHEY_SIMPLEX, 0.8, (0, 255, 0), 2)class_names 可以从 NTU 的动作类别定义文件里读进去。这段逻辑不改变模型结构只是在推理循环里多算一次 argsort 和 putText对帧率几乎没影响。6.3 保存逐帧识别结果到 CSV答辩和实验分析时光看视频不够还需要逐帧的识别结果做数据统计。我给 demo_offline.py 加过一段 CSV 导出逻辑每处理一帧就把帧号和该帧 Top-1 的类别和置信度追加到 csv 文件。这个 csv 后面可以拉进 pandas 画动作类目随时间的切换曲线比来回拖动视频定位动作切换点高效得多。整套流程吃透以后换个数据集只需要改 feeder 和 config网络结构基本不动的——ST-GCN 的通用性就在这里。我在拆这个项目的前两周犯过不少低级错误印象最深的一次是拿 AddEdgeSTGCN12345.pt 直接跑 demo结果画面上的标签一直在乱跳排查半天才发现是图结构没配加边。从那以后我每次拿到任何人的训练权重第一件事是打印 state_dict 对比 config 的图参数而不是急着跑推理。这个习惯帮我避开了各种稀奇古怪的运行时错误希望也能帮到你。本文还有配套的精品资源点击获取
返回列表