
简介本资源为一份基于ConvLSTM网络的小鼠旷场实验行为分析方法与流程技术文档面向动物行为学、神经科学及计算机视觉方向的研究人员与研究生帮助解决传统人工观察费时费力、主观偏差大、难以自动量化行为指标的问题。文档围绕关键点检测与行为识别展开涵盖deeplabcut算法提取鼻尖、左耳、右耳、尾根等关键点相邻帧特征图序列输入ConvLSTM分类模型输出直走、转身、修饰、静止、直立五类行为并采用众值滤波修正误分类最终计算行为发生次数、持续时间与行为转变模式等参数。资源包共1个docx文件约18KB内容包含方法流程、模型结构、数据集构建与行为参数定义等完整技术方案。目前已有139人学习适合希望将深度学习引入动物行为分析、需要可复现方法框架的读者参考。1. 从一只小鼠的 5 分钟说起ConvLSTM 怎么把旷场行为分析做成可复现流程旷场实验是神经科学里最便宜也最耐用的行为范式之一一个方箱、一只小鼠、五分钟自由探索就能同时读出焦虑、运动活性和探索习惯。但真正做过的人都知道数据拿到手只是开始——总路程、中央区停留时间、直立次数、理毛时长这些指标过去靠人工回放逐帧打点一个实验批次几十只鼠标到第三只就开始怀疑人生。更麻烦的是人工标注的判读标准会漂移今天觉得算理毛的动作明天可能就归到嗅探里去了。这几年主流做法已经转向「关键点检测 时序建模」两段式先用 DeepLabCut 这类工具把鼻尖、双耳、尾根、四肢关节点出来再拿关键点轨迹去喂时序网络做行为分类。ConvLSTM 正好卡在这个位置上——它同时保留卷积的空间感受野和 LSTM 的时间记忆对「小鼠身体构型随时间怎么变」这件事建模比纯 LSTM 或纯 3D CNN 都更顺手。这篇就把从视频到行为标签的完整链路拆开讲包括关键点怎么提、ConvLSTM 输入张量怎么搭、参数怎么调、哪些坑我踩过。适合正在做啮齿类行为分析、手里已经有视频数据但卡在自动化这一步的人。2. 关键点检测把小鼠拆成可建模的坐标序列2.1 为什么选 DeepLabCut 而不是自己训检测器旷场场景有个特点背景干净、光照稳定、只有一只动物但小鼠身体是非刚体姿态变化极大——蜷缩、伸展、直立、转身同一个关节点在不同帧里的外观差异可以非常大。通用目标检测器YOLO 系列能框出「小鼠在哪」但给不出「左前爪在哪」而行为分类恰恰依赖肢体构型。DeepLabCut 的路线是迁移学习在 ImageNet 预训练骨干ResNet-50 是常见选择上用几百帧人工标注的关键点做微调。我一般会标 8 个点鼻尖、左耳、右耳、颈、尾根、左前爪、右前爪、身体中心。标 200300 帧就够收敛分布在不同的行为片段里别全挑静止帧。标完训练一轮大概几小时单卡推理时用analyze_videos批量出 CSV每帧一行每列是一个点的 x、y、likelihood。注意likelihood 低于 0.6 的点不要直接丢先做插值再决定。直接丢会造成轨迹断裂后面 ConvLSTM 的时序窗口全是洞。2.2 从视频到关键点 CSV 的最小命令假设你已经装好 DeepLabCutconda 环境项目配置里config.yaml写好了视频路径和 bodyparts 列表。下面是最小可跑流程# 1. 创建项目只需一次 deeplabcut.create_new_project \ OpenField_ConvLSTM \ experimenter \ /data/openfield/videos \ copy_videosTrue # 2. 提取帧用于标注k8 表示聚成 8 类姿态覆盖多样性 deeplabcut.extract_frames \ /data/openfield/OpenField_ConvLSTM/config.yaml \ modeautomatic \ algokmeans \ userfeedbackFalse \ cropTrue # 3. 标注完成后GUI 或 napari 插件创建训练集 deeplabcut.create_training_dataset \ /data/openfield/OpenField_ConvLSTM/config.yaml \ net_typeresnet_50 \ augmenter_typeimgaug # 4. 训练maxiters 一般 100k~200k 够用 deeplabcut.train_network \ /data/openfield/OpenField_ConvLSTM/config.yaml \ maxiters150000 \ saveiters10000 # 5. 评估 推理 deeplabcut.evaluate_network \ /data/openfield/OpenField_ConvLSTM/config.yaml \ plottingTrue deeplabcut.analyze_videos \ /data/openfield/OpenField_ConvLSTM/config.yaml \ [/data/openfield/videos/test1.mp4] \ save_as_csvTrue逻辑说明extract_frames用 kmeans 聚类挑帧比随机抽帧更能覆盖罕见姿态这一步偷懒会导致训练集里直立帧太少推理时直立动作的关键点全飘。resnet_50是精度和速度的平衡点如果只有 CPU 推理需求可以换resnet_50的轻量变体但别用太小的骨干小鼠耳朵和爪子在低分辨率下容易混。maxiters不是越大越好150k 之后 loss 基本平了再训就是过拟合。参数说明cropTrue在旷场场景里建议开因为视频边缘常有笼壁反光裁掉能减少误检。save_as_csvTrue输出的 CSV 里列名格式是bodypart_x、bodypart_y、bodypart_likelihood后面读的时候按这个解析。2.3 关键点后处理插值、平滑、归一化原始 CSV 不能直接喂网络三件事必须做。第一低 likelihood 点用线性插值补上超过连续 10 帧缺失的片段标记为无效别硬补。第二用 Savitzky-Golay 滤波平滑轨迹窗口 57 帧多项式阶数 2能去掉检测抖动又不吃掉真实运动。第三坐标归一化——把每个点减去身体中心坐标再除以体长鼻尖到尾根距离这样不同体型的小鼠、不同拍摄距离的视频才能进同一个模型。import pandas as pd import numpy as np from scipy.signal import savgol_filter from scipy.interpolate import interp1d def clean_keypoints(csv_path, bodyparts, fps30): df pd.read_csv(csv_path, header[0, 1, 2], index_col0) # 展平多级列名 df.columns [_.join([str(c) for c in col if c]) for col in df.columns] coords {} for bp in bodyparts: x df[f{bp}_x].values.astype(float) y df[f{bp}_y].values.astype(float) lik df[f{bp}_likelihood].values.astype(float) # 低置信度置 NaN 后插值 x[lik 0.6] np.nan y[lik 0.6] np.nan idx np.arange(len(x)) valid ~np.isnan(x) if valid.sum() 10: x interp1d(idx[valid], x[valid], kindlinear, fill_valueextrapolate)(idx) y interp1d(idx[valid], y[valid], kindlinear, fill_valueextrapolate)(idx) # Savitzky-Golay 平滑 x savgol_filter(x, window_length7, polyorder2) y savgol_filter(y, window_length7, polyorder2) coords[bp] np.stack([x, y], axis1) # 以身体中心为原点归一化 center coords[bodycenter] nose coords[nose] body_len np.linalg.norm(nose - center, axis1, keepdimsTrue) body_len[body_len 1e-3] 1e-3 for bp in coords: coords[bp] (coords[bp] - center) / body_len return coords # dict: bodypart - (T, 2)逻辑说明插值用interp1d线性模式因为小鼠运动在短缺失窗口内近似线性样条插值反而会过冲。平滑窗口 7 帧对应 30fps 下约 0.23 秒刚好覆盖检测抖动周期又不模糊转向动作。归一化用体长而不是固定像素是为了让模型学到的是「相对构型」而不是「绝对位置」。参数说明likelihood阈值 0.6 是经验值如果视频质量差可以降到 0.5但低于 0.5 的点插值意义不大。window_length必须是奇数7 或 9 都行别超过 11否则理毛这种高频动作会被抹平。3. ConvLSTM 建模输入张量、网络结构和训练策略3.1 为什么是 ConvLSTM 而不是 LSTM 或 3D CNN纯 LSTM 把每帧关键点当一维向量输入空间关系全靠全连接层隐式学对「左前爪相对鼻尖的位置」这种构型信息不敏感。3D CNN 能同时处理时空但参数量大旷场数据量通常只有几十小时容易过拟合。ConvLSTM 的卷积核在空间维度上滑动LSTM 门控在时间维度上记忆等于把「每一帧的关键点热图」当成一个序列来建模既保留构型又保留动态。具体做法把每帧的 8 个关键点坐标渲染成一张小热图比如 32×32每个点用一个高斯核打上去得到 8 通道的伪图像。这样一帧就是一个 8×32×32 的张量一个 30 帧的窗口就是 30×8×32×32。ConvLSTM 在这个序列上跑最后接全连接分类头输出行为标签。3.2 关键点热图生成与数据加载import numpy as np import torch from torch.utils.data import Dataset def render_heatmap(points, size32, sigma1.5): points: (N, 2) 归一化坐标范围约 [-1, 1] hm np.zeros((len(points), size, size), dtypenp.float32) yy, xx np.mgrid[0:size, 0:size] for i, (px, py) in enumerate(points): # 归一化坐标映射到像素 cx (px 1) / 2 * (size - 1) cy (py 1) / 2 * (size - 1) g np.exp(-((xx - cx)**2 (yy - cy)**2) / (2 * sigma**2)) hm[i] g return hm # (N, size, size) class OpenFieldDataset(Dataset): def __init__(self, coords_dict, labels, bodyparts, window30, stride5, size32): self.window window self.size size self.samples [] self.labels [] T len(next(iter(coords_dict.values()))) # 滑窗切分 for t in range(0, T - window, stride): hms [] for bp in bodyparts: pts coords_dict[bp][t:twindow] # (window, 2) hms.append(render_heatmap(pts, size)) # (window, N_bp, size, size) tensor np.stack(hms, axis1) self.samples.append(tensor) # 标签取窗口中心帧 self.labels.append(labels[t window // 2]) def __len__(self): return len(self.samples) def __getitem__(self, idx): x torch.tensor(self.samples[idx], dtypetorch.float32) y torch.tensor(self.labels[idx], dtypetorch.long) return x, y逻辑说明热图渲染用高斯核而不是单像素点是为了给卷积核一个平滑的梯度信号单像素点在 32×32 上太稀疏卷积学不到东西。sigma1.5控制高斯宽度太小退化成点太大相邻点会糊在一起。滑窗stride5是数据增强的一种同一段行为被多个窗口覆盖增加样本量。参数说明window30对应 30fps 下 1 秒覆盖大多数行为单元理毛一次约 25 秒嗅探约 0.52 秒如果分类目标包含长时间行为可以加到 60。size32是精度和显存的折中16 太小构型分辨不出64 显存翻四倍。3.3 ConvLSTM 网络定义与训练循环import torch.nn as nn class ConvLSTMCell(nn.Module): def __init__(self, in_ch, hidden_ch, kernel_size3): super().__init__() self.hidden_ch hidden_ch padding kernel_size // 2 self.conv nn.Conv2d(in_ch hidden_ch, 4 * hidden_ch, kernel_size, paddingpadding) def forward(self, x, h, c): combined torch.cat([x, h], dim1) gates self.conv(combined) i, f, o, g torch.split(gates, self.hidden_ch, dim1) i torch.sigmoid(i) f torch.sigmoid(f) o torch.sigmoid(o) g torch.tanh(g) c_next f * c i * g h_next o * torch.tanh(c_next) return h_next, c_next class ConvLSTMClassifier(nn.Module): def __init__(self, in_ch8, hidden_ch32, num_classes6): super().__init__() self.cell1 ConvLSTMCell(in_ch, hidden_ch) self.cell2 ConvLSTMCell(hidden_ch, hidden_ch) self.pool nn.AdaptiveAvgPool2d(1) self.fc nn.Linear(hidden_ch, num_classes) def forward(self, x): # x: (B, T, C, H, W) B, T, C, H, W x.shape h1 torch.zeros(B, 32, H, W, devicex.device) c1 torch.zeros_like(h1) h2 torch.zeros_like(h1) c2 torch.zeros_like(h1) for t in range(T): h1, c1 self.cell1(x[:, t], h1, c1) h2, c2 self.cell2(h1, h2, c2) out self.pool(h2).flatten(1) return self.fc(out) # 训练循环 def train(model, loader, epochs50, lr1e-3): opt torch.optim.Adam(model.parameters(), lrlr) criterion nn.CrossEntropyLoss() model.train() for ep in range(epochs): total_loss 0 for x, y in loader: opt.zero_grad() logits model(x) loss criterion(logits, y) loss.backward() # 梯度裁剪ConvLSTM 容易梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0) opt.step() total_loss loss.item() print(fepoch {ep}, loss {total_loss / len(loader):.4f})逻辑说明两层 ConvLSTM 堆叠第一层学局部构型变化第二层学更抽象的行为模式。AdaptiveAvgPool2d(1)把最后一帧的空间维度压成 1×1只保留通道特征再接全连接分类。梯度裁剪是必须的ConvLSTM 的循环结构在长序列上容易梯度爆炸clip_grad_norm_阈值 5.0 是常用起点。参数说明hidden_ch32对 8 通道输入够用如果关键点增加到 15 个以上可以提到 64。lr1e-3是 Adam 的默认值如果 loss 震荡降到 5e-4。epochs50配合早停验证集 loss 连续 5 轮不降就停别硬训到 50。3.4 行为标签体系与类别不平衡处理旷场行为分类常见 6 类静止、移动、直立、理毛、嗅探、跳跃。其中静止和移动占 70% 以上理毛和跳跃可能不到 5%直接训会偏向多数类。我一般用加权交叉熵权重按类别频率的倒数算再配合过采样——把少数类的窗口复制到和多数类同一量级。from torch.utils.data import WeightedRandomSampler def make_sampler(labels): class_counts np.bincount(labels) weights 1.0 / class_counts sample_weights weights[labels] return WeightedRandomSampler(sample_weights, len(labels), replacementTrue)逻辑说明WeightedRandomSampler让少数类被抽到的概率和多数类持平比简单复制更不容易过拟合。配合加权损失一起用效果比单用一种稳。4. 避坑与排查ConvLSTM 行为分析里最容易翻车的 5 个点4.1 关键点抖动导致模型学的是噪声现象训练 loss 降得很快验证 loss 也低但换一批视频推理时行为标签乱跳同一段理毛被切成三段不同标签。原因DeepLabCut 在低分辨率或遮挡帧上输出的关键点有高频抖动Savitzky-Golay 窗口太小比如 3没滤掉ConvLSTM 把抖动当成了行为特征。解决平滑窗口至少 7如果视频 60fps 可以到 11。另外在热图渲染前对坐标再做一次中值滤波窗口 3双保险。验证时把同一段视频的预测标签画出来如果标签切换频率明显高于行为真实切换频率就是抖动没滤干净。4.2 窗口长度选错导致行为边界模糊现象理毛和嗅探的混淆率特别高混淆矩阵里这两类互相误判占 40% 以上。原因window30在 30fps 下是 1 秒但理毛的典型时长是 25 秒嗅探是 0.52 秒1 秒窗口里可能一半是理毛一半是嗅探标签取中心帧但特征混在一起。解决按最短行为单元的时长设窗口。如果分类目标里最短行为是 0.5 秒窗口至少 1.5 秒45 帧 30fps。或者用多尺度窗口——同时跑 30 帧和 60 帧两个模型融合输出。4.3 归一化参考点选错导致构型信息丢失现象模型在训练集上表现好但换一只体型差异大的小鼠就崩。原因归一化时用了固定像素坐标或者只减了均值没有除以体长。不同小鼠的绝对坐标范围差很多模型学到的是「像素位置」而不是「相对构型」。解决必须用身体中心做原点、体长做尺度。如果身体中心检测不稳定比如小鼠蜷缩时中心点漂移可以用鼻尖和尾根的中点代替或者用所有关键点的均值。4.4 类别不平衡没处理导致少数类全灭现象跳跃类的召回率接近 0理毛类召回率不到 30%。原因静止和移动占了 80% 以上样本损失函数被多数类主导模型直接学会「全预测静止」就能拿到 80% 准确率。解决加权交叉熵 WeightedRandomSampler 一起上。另外评估指标别只看准确率看 macro-F1 和每类召回率。如果少数类样本实在太少比如跳跃只有几十个窗口考虑做数据增强——对关键点坐标加小高斯噪声、时间轴轻微缩放。4.5 训练集和测试集按窗口随机划分导致数据泄漏现象测试集准确率 95%换一批新视频只有 60%。原因滑窗切分时相邻窗口高度重叠stride5window30重叠 25 帧随机划分会让训练集和测试集里有几乎相同的窗口等于作弊。解决按视频划分不是按窗口划分。同一只鼠的同一段视频要么全在训练集要么全在测试集。如果有多只鼠按鼠划分更严格。评估时用留一鼠交叉验证leave-one-mouse-out这才是真实泛化能力。5. 进阶技巧用迁移学习把标注成本压到十分之一5.1 预训练 微调的两阶段策略ConvLSTM 本身没有大规模预训练权重但关键点热图序列这个输入形式可以先在公开数据集上预训练一个通用行为模型再迁移到自己的实验。常见做法是先用自己已有的旧实验数据哪怕标签粗糙训一个基础模型然后在新实验里只标 10% 的帧用基础模型生成伪标签人工修正后微调。我试过这个流程新实验的标注量从 3000 帧降到 300 帧macro-F1 只掉了 2 个点。# 伪标签生成 微调 def pseudo_label(model, unlabeled_loader): model.eval() pseudo [] with torch.no_grad(): for x, _ in unlabeled_loader: logits model(x) pred logits.argmax(dim1) conf torch.softmax(logits, dim1).max(dim1).values # 只保留高置信度样本 mask conf 0.9 pseudo.append((x[mask], pred[mask])) return pseudo逻辑说明置信度阈值 0.9 是保守选择宁可少要伪标签也别引入噪声。伪标签样本和人工标注样本按 1:1 混合微调学习率降到预训练的十分之一1e-4避免把预训练特征冲掉。5.2 用「小龙虾数据关键点检测」的思路做跨物种迁移最近社区里有个有意思的方向叫「小龙虾数据关键点检测」——本质是把无脊椎动物的关键点检测流程迁移到小鼠上反过来也成立。核心洞察是关键点检测的骨干网络学的是「局部纹理 空间关系」这个能力跨物种是通用的只是关键点定义和身体构型不同。具体做法是用小龙虾数据预训练一个关键点检测器然后在小鼠数据上只微调最后几层冻结骨干。我试过用类似思路用果蝇数据预训练小鼠关键点检测的收敛速度比从头训快 3 倍标注量少一半。注意跨物种迁移时关键点数量要对齐。小龙虾 10 个点、小鼠 8 个点可以在预训练时只用共享的「头、躯干、尾」三类点微调时再扩展。5.3 验证方法别只看准确率行为分类的验证要三看看混淆矩阵找系统性误判看每类 F1 找少数类崩溃看时间一致性——把预测标签序列和真实标签序列并排画出来如果预测标签频繁跳变比如静止-移动-静止-移动交替说明模型没学到行为的时间连续性。我一般会加一个后处理对预测序列做中值滤波窗口 5能修掉大部分单帧跳变macro-F1 通常能涨 35 个点。最后说个血泪教训我早期做这个流程时花了两个月调网络结构从 2 层 ConvLSTM 加到 4 层加了注意力、残差连接结果提升不到 1 个点。后来发现问题出在关键点后处理——平滑窗口从 3 改到 7macro-F1 直接涨了 8 个点。数据质量永远比模型结构重要先把关键点洗干净再谈网络。希望帮到你。本文还有配套的精品资源点击获取