拓冰建站拓冰建站
首页 / 资讯中心 / 正文

多质心表征网络:解决行人重识别跨域泛化难题

1. 项目概述当行人重识别遇上域自适应如果你做过行人重识别Person Re-identification简称Re-ID项目大概率会遇到一个让人头疼的“拦路虎”辛辛苦苦在源域比如一个特定监控摄像头网络上训练好的模型一旦部署到目标域比如另一个光照、视角、背景完全不同的新摄像头网络性能往往会断崖式下跌。这就是典型的域偏移Domain Shift问题。我们这次要拆解的论文《Multi-Centroid Representation Network for Domain Adaptive Person Re-ID》就是针对这个痛点的一剂猛药。它没有走传统的对抗训练或者风格迁移的老路而是从一个非常直观的角度切入——用多个质心Multi-Centroid来更精细地刻画一个行人的特征分布。简单来说传统方法在训练时通常会给每个行人ID分配一个单一的特征向量或称为“类原型”作为其代表。但在跨域场景下同一个行人在不同摄像头下的外观变化可能极大比如从阳光下走到阴影里从正面走到侧面一个“点”状的原型很难覆盖所有这些变化。这篇论文的核心思想是为什么不用一组特征点即多个质心来代表一个行人呢这样模型就能学习到每个行人更丰富、更鲁棒的特征表示从而在面临新领域时拥有更强的泛化能力。这个想法听起来简单但背后的实现细节和理论支撑才是其精髓所在。接下来我们就深入代码和原理层面看看这个“多质心”网络是如何构建又是如何在实际中发挥威力的。2. 核心思路拆解从单一原型到多质心表征要理解这篇论文的价值我们得先看看它要解决什么问题以及它之前的主流方案存在什么局限。2.1 域自适应行人重识别的核心挑战行人重识别的目标是在非重叠摄像头视图下匹配同一个行人的不同图像。在有监督设定下我们拥有大量带标签的源域数据模型可以学得很好。但现实是目标域比如一个新部署的商场监控系统往往没有标签或者标注成本极高。域自适应Domain Adaptation, DA的目标就是利用有标签的源域数据让模型能直接在没有标签的目标域上取得好效果。域偏移主要体现在两个方面风格偏移不同摄像头由于硬件、参数、光照、色彩渲染不同导致图像的低级视觉特征颜色、纹理分布不同。内容偏移不同场景下的行人姿态、背景、遮挡情况、行人密度等高级语义特征分布也不同。早期的域自适应Re-ID方法比如直接使用生成对抗网络进行图像翻译如SPGAN、CycleGAN试图将目标域图像风格迁移到源域或者反之。这类方法能较好地缓解风格偏移但对内容偏移的解决能力有限且图像生成过程复杂容易引入伪影。另一大类方法是基于特征对齐的对抗学习如MMT、ECN在特征空间拉近源域和目标域的分布。这类方法更直接但训练不稳定且对齐整个特征分布有时会模糊掉对Re-ID任务至关重要的判别性细节。2.2 多质心表征的直觉与优势无论是图像翻译还是特征对齐大多还是在“全局分布”的层面做文章。而这篇论文的作者洞察到了一个更细粒度的关键点一个行人的特征在特征空间里不应该是一个点而应该是一个分布。想象一下同一个人的多张图片有正面、侧面、背面有穿外套、脱外套有戴帽子、不戴帽子。这些图片的特征在嵌入空间Feature Embedding Space中会形成一个“簇”Cluster。传统方法用这个簇的均值一个质心来代表这个人。但在跨域时目标域中这个人的图片可能只覆盖了这个簇的某一部分比如只有侧面照。此时用源域学到的那个“均值”质心去匹配就可能产生偏差。多质心表征的思路是主动地为每个行人ID学习K个质心用这K个点来共同刻画该行人特征簇的形状和范围。这样做有几个明显的好处更强的表征能力多个质心可以捕捉行人外观的多模态变化如不同视角、不同着装状态。更好的跨域鲁棒性即使目标域只出现了该行人的部分模态比如只看到了侧面只要有一个侧面相关的质心能匹配上就能实现正确识别。更精细的对比学习在计算对比损失时我们可以进行“质心-质心”之间的对比而不仅仅是“图像-质心”对比这能带来更丰富的监督信号。2.3 网络整体架构与工作流程论文提出的Multi-Centroid Representation Network整体架构清晰主要包含以下几个核心模块特征提取骨干网络通常是一个在ImageNet上预训练的ResNet或IBN-Net用于从输入图像中提取基础特征图。多质心生成模块这是论文的核心创新。它不是一个独立的子网络而是一种训练机制和表征形式。具体来说在训练过程中对于属于同一个行人ID的所有样本模型会动态地维护和更新K个质心向量。这些质心通过聚类或可学习的方式获得。域自适应模块为了处理域偏移论文通常会结合一个现有的、有效的域自适应方法例如基于记忆库的对比学习、对抗性判别器。多质心表征可以作为这个模块的更强、更稳定的输入。损失函数损失函数是多任务学习的组合通常包括源域有监督损失在源域上使用多质心表征计算交叉熵损失和三元组损失。目标域无监督损失在目标域上利用多质心进行伪标签生成和对比学习。例如为目标域样本分配伪标签时可以计算该样本特征与所有源域行人ID的多个质心之间的距离选择最近的那个质心所属的ID作为伪标签。这个过程因为有了多个质心而更加可靠。域对齐损失可选如果采用了对抗训练则包含一个域判别损失。整个训练流程是一个迭代优化的过程利用源域标签初始化多质心 - 在目标域上生成伪标签 - 用伪标签更新目标域特征并 refine 多质心 - 用更新后的多质心和特征进一步优化网络参数。3. 核心实现细节与实操要点理解了宏观思路我们深入到代码实现层面看看几个最关键的技术点是如何落地的。这里我会结合常见的PyTorch实现框架来讲解。3.1 质心数量K的选择与初始化质心数量K是一个超参数它控制着表征的细粒度。K太小退化成单一原型K太大则可能引入噪声并且增加计算开销也容易在小ID的类别上过拟合。实操选择论文中通常通过实验确定对于Market-1501、DukeMTMC-reID这类数据集K在3到5之间是一个较好的平衡点。对于更复杂、类内变化更大的数据集可以适当增大K。初始化策略这是一个关键细节。不能随机初始化因为那样质心可能没有意义。常见的策略有聚类初始化在训练初期先用源域数据对每个行人ID的所有样本特征进行K-Means聚类将聚类中心作为该ID的K个初始质心。可学习参数直接将每个ID的K个质心定义为可学习的参数与网络一起随机初始化并端到端训练。这种方式更灵活但需要更谨慎的学习率设置。# 伪代码示例聚类初始化质心 def init_centroids_by_kmeans(features, labels, K): features: 源域所有样本的特征向量 [N, D] labels: 对应的行人ID [N] K: 每个ID的质心数 returns: centroids_dict {label: [K, D]} centroids_dict {} unique_labels torch.unique(labels) for label in unique_labels: idx (labels label) id_features features[idx] # 获取该ID的所有特征 # 使用K-Means聚类 kmeans KMeans(n_clustersK, random_state0).fit(id_features.cpu().numpy()) centroids_dict[label] torch.from_numpy(kmeans.cluster_centers_).to(features.device) return centroids_dict注意聚类初始化通常在第一个训练周期epoch开始前进行并且只做一次。在后续训练中质心会随着网络参数的更新而动态更新例如通过移动平均。3.2 质心的动态更新机制在训练过程中随着网络参数更新特征也在不断变化因此质心也需要同步更新。最直接的方式是在每个batch后用该batch中属于同一ID的样本特征来重新计算质心。但这样做计算量大且由于mini-batch的采样偏差质心会非常不稳定。论文普遍采用“动量更新”策略这借鉴了无监督对比学习如MoCo的思想。为每个行人ID维护一个队列或直接维护K个质心向量每次前向传播得到样本特征后用该特征以动量方式更新对应的质心。# 伪代码示例动量更新质心 class MomentumCentroidUpdater: def __init__(self, momentum0.999): self.momentum momentum # 假设我们已经有一个centroids张量形状为 [num_classes, K, feature_dim] self.centroids ... # 初始化好的质心 def update(self, features, labels, centroid_indices): features: 当前batch的特征 [B, D] labels: 当前batch的ID [B] centroid_indices: 每个特征对应其ID下的第几个质心 [B]需要通过最近邻匹配得到 with torch.no_grad(): for feat, label, c_idx in zip(features, labels, centroid_indices): # 找到对应的旧质心 old_centroid self.centroids[label, c_idx] # 动量更新: new m * old (1 - m) * feat new_centroid self.momentum * old_centroid (1 - self.momentum) * feat self.centroids[label, c_idx] new_centroid这里有一个关键步骤如何为当前的特征feat分配它应该更新哪个质心c_idx这通常通过计算该特征与其所属ID的所有K个质心的余弦相似度或欧氏距离选择最接近的那个质心的索引作为c_idx。3.3 基于多质心的损失函数设计损失函数是驱动模型学习多质心判别性表征的关键。主要包括两部分1. 源域有监督损失多质心交叉熵损失对于源域样本我们需要计算它与对应ID的所有K个质心的相似度。一种做法是将K个质心视为该ID的K个“子类”然后计算一个多类交叉熵损失。但更常见的、也是论文中的做法是将K个质心融合成一个代表向量例如取平均或加权平均然后用这个融合后的向量与样本特征计算余弦相似度再送入交叉熵损失。在反向传播时梯度会流向参与融合的所有质心。# 伪代码基于融合质心的交叉熵损失 def fused_centroid_ce_loss(features, labels, centroids_dict): batch_size features.size(0) loss 0 for i in range(batch_size): feat features[i] label labels[i] # 获取该label的K个质心 [K, D] cents centroids_dict[label] # 融合质心例如简单平均 fused_cent torch.mean(cents, dim0) # [D] # 计算余弦相似度作为logit logit F.cosine_similarity(feat.unsqueeze(0), fused_cent.unsqueeze(0)) # 这里需要将logit整合进一个所有类别的logits向量中计算CE Loss # 简化表示计算该样本与融合质心的距离损失 # 实际实现会更复杂需要构建分类权重矩阵 return loss多质心三元组损失传统的三元组损失是在样本级别进行的。引入多质心后我们可以构造“样本-质心-质心”的三元组。例如对于一个锚点样本选择其对应ID中最近的一个质心作为正样本选择其他ID中最近的一个质心作为负样本。这样拉近了样本与其类内质心的距离同时推远了类间质心的距离。2. 目标域无监督损失目标域没有真实标签核心是利用多质心生成高质量的伪标签并进行对比学习。伪标签生成对于目标域的一个样本特征计算它与所有源域ID的所有质心共num_source_ids * K个的距离。找到最近的质心则该质心所属的源域ID就被赋予该目标样本作为伪标签。由于有K个质心匹配成功的概率和可靠性比单一原型高得多。对比学习获得伪标签后就可以像在源域一样为目标域数据计算基于多质心的交叉熵损失和三元组损失。同时也可以进行目标域内部的特征对比拉近伪标签相同的样本特征。3.4 与现有域自适应方法的结合多质心表征是一个表征学习框架它可以与多种现有的域自适应范式结合作为它们的“增强组件”。论文中常见的结合方式有与基于记忆库的方法结合例如MMT方法维护了一个目标域特征的内存库。在多质心版本中我们可以维护两个内存库一个是目标域样本特征库另一个是源域质心库。目标域伪标签的生成通过查询源域质心库来完成而对比学习则在目标域特征库内部进行。与对抗性方法结合可以在特征提取器后添加一个域分类器进行对抗训练。此时输入域分类器的特征可以是样本特征也可以是该样本所属的融合质心特征。多质心提供的更稳定、更具判别性的特征有助于对抗训练更聚焦于域不变特征的学习。4. 实验配置与复现指南想要复现或借鉴这篇论文的工作你需要一个标准的行人重识别研究环境。以下是详细的步骤和避坑点。4.1 环境搭建与数据准备硬件与软件基础GPU至少需要一块显存11GB以上的GPU如RTX 2080Ti, RTX 3080, RTX 4090。因为Re-ID模型和图像尺寸较大且batch size不能太小。深度学习框架PyTorch 1.7 CUDA版本需要与你的GPU驱动匹配。关键Python库torch,torchvision,numpy,scikit-learn(用于K-Means),opencv-python,PIL,tqdm,tensorboard(用于可视化)。数据集下载与预处理 你需要准备至少一个源域数据集和一个目标域数据集。经典组合包括源域 - 目标域 Market-1501 - DukeMTMC-reID, DukeMTMC-reID - Market-1501, MSMT17 - Market-1501。数据下载从学术网站如Zheng等人主页、Duke大学主页或开源项目如fast-reid提供的链接下载。预处理这是Re-ID实验最繁琐但最重要的一步。必须严格按照数据集的官方说明或主流代码库如fast-reid, Torchreid的预处理脚本进行操作。通常包括解压文件到特定目录结构。运行提供的Python脚本生成训练/查询/画廊gallery的图像列表文件。确保每个人物ID一个文件夹或者列表文件中包含正确的图像路径和ID、摄像头ID信息。常见坑点DukeMTMC-reID有多个版本注意使用“DukeMTMC-reID”这个重识别专用版本而不是原始的多目标跟踪版本。MSMT17数据量很大预处理和加载较慢需要耐心。4.2 模型实现关键代码片段这里提供一个高度简化的、聚焦于多质心核心逻辑的代码框架import torch import torch.nn as nn import torch.nn.functional as F from sklearn.cluster import KMeans class MultiCentroidReIDNet(nn.Module): def __init__(self, backbone, num_classes, feature_dim2048, K4, momentum0.999): super().__init__() self.backbone backbone # 例如 ResNet50 self.pooling nn.AdaptiveAvgPool2d((1, 1)) self.bottleneck nn.BatchNorm1d(feature_dim) # BNNeck是一种常用技巧 self.bottleneck.bias.requires_grad_(False) # 多质心相关参数 self.K K self.momentum momentum # 为每个源域类别维护K个质心 [num_classes, K, feature_dim] self.register_buffer(source_centroids, torch.zeros(num_classes, K, feature_dim)) # 标记质心是否已初始化 self.centroids_initialized False # 分类器用于源域监督损失 self.classifier nn.Linear(feature_dim, num_classes, biasFalse) def init_centroids(self, source_dataloader): 在训练开始前用源域数据初始化质心 print(Initializing centroids with K-Means...) self.eval() all_features [] all_labels [] with torch.no_grad(): for imgs, labels, _ in source_dataloader: imgs imgs.cuda() feats self.backbone(imgs) feats self.pooling(feats).flatten(1) feats self.bottleneck(feats) all_features.append(feats.cpu()) all_labels.append(labels) all_features torch.cat(all_features, dim0) all_labels torch.cat(all_labels, dim0) for label in torch.unique(all_labels): label_mask (all_labels label) id_features all_features[label_mask].numpy() if len(id_features) self.K: # 如果某个ID的样本数少于K复制样本或减少该ID的K值 # 这里简单复制直到满足K repeat_times self.K // len(id_features) 1 id_features np.tile(id_features, (repeat_times, 1))[:self.K] else: kmeans KMeans(n_clustersself.K, random_state0, n_init10).fit(id_features) self.source_centroids[label] torch.from_numpy(kmeans.cluster_centers_) self.centroids_initialized True self.train() print(Centroids initialization done.) def forward(self, x, labelNone, domainsource): # 提取特征 feat self.backbone(x) feat self.pooling(feat).flatten(1) bn_feat self.bottleneck(feat) # 用于度量学习计算损失的特征 if domain source and self.training: # 源域训练计算分类logits和质心相关损失 cls_logits self.classifier(bn_feat) # 用于交叉熵损失 # 同时需要更新质心 self._update_centroids_momentum(bn_feat, label) return cls_logits, bn_feat else: # 目标域推理或测试直接返回特征 return bn_feat def _update_centroids_momentum(self, features, labels): 动量更新源域质心 with torch.no_grad(): for feat, label in zip(features, labels): # 找到该特征对应ID的K个质心 centroids self.source_centroids[label] # [K, D] # 计算与所有质心的距离选择最近的 distances 1 - F.cosine_similarity(feat.unsqueeze(0), centroids, dim1) # [K] nearest_idx distances.argmin() # 动量更新 old_centroid self.source_centroids[label, nearest_idx] new_centroid self.momentum * old_centroid (1 - self.momentum) * feat self.source_centroids[label, nearest_idx] new_centroid def assign_pseudo_label(self, target_features): 为目标域特征分配伪标签基于最近质心 with torch.no_grad(): pseudo_labels [] for feat in target_features: # feat: [D] # 计算与所有源域质心的距离 [num_classes * K] # 这里需要将source_centroids展平 flat_centroids self.source_centroids.view(-1, self.source_centroids.size(-1)) # [num_classes*K, D] distances 1 - F.cosine_similarity(feat.unsqueeze(0), flat_centroids, dim1) nearest_flat_idx distances.argmin() # 将展平的索引映射回 [class_idx, centroid_idx] class_idx nearest_flat_idx // self.K pseudo_labels.append(class_idx) return torch.tensor(pseudo_labels).to(target_features.device)4.3 训练策略与超参数调优训练一个多质心域自适应Re-ID模型通常分为两个阶段或多个交替迭代的阶段。第一阶段源域预训练目的在源域数据上训练一个强大的基线模型并初始化多质心。操作使用标准的Re-ID损失交叉熵损失 三元组损失训练网络。关键步骤在第一个epoch开始前调用init_centroids函数用K-Means初始化质心。超参数学习率如0.00035 Batch Size如64 优化器如Adam SGD with momentum。训练足够多的epoch如60直到收敛。第二阶段域自适应训练目的利用源域质心和目标域数据进行无监督域自适应。流程 a.前向传播分别通过源域batch和目标域batch。 b.源域分支计算有监督损失交叉熵三元组并动量更新源域质心。 c.目标域分支用assign_pseudo_label函数为目标域特征生成伪标签。这里有一个重要技巧对伪标签进行筛选。只保留那些与最近质心距离小于某个阈值高置信度的样本参与损失计算以减少噪声标签的影响。 d.目标域损失对高置信度的目标域样本用其伪标签计算交叉熵损失和三元组损失此时将伪标签视为真实标签。同时可以加入目标域样本之间的对比学习损失如Instance Contrastive Loss。 e.总损失总损失 源域监督损失 λ * 目标域伪标签损失。λ是一个平衡权重通常从0开始随着训练逐渐增大课程学习策略。超参数调优重点K尝试3, 4, 5。通常4是一个不错的起点。momentum质心动量更新系数通常很高如0.999或0.99。伪标签阈值用于筛选高置信度目标样本。需要根据特征距离的分布来调整例如只选择距离最小的前50%的样本。λ目标域损失权重。可以采用线性增长策略λ current_epoch / total_epochs * max_lambda。学习率调整在域自适应阶段通常使用更小的学习率如预训练阶段的1/10并配合余弦退火调度器。5. 常见问题、排查技巧与效果分析在实际复现和应用过程中你肯定会遇到各种问题。下面是我在实验过程中踩过的一些坑和总结的经验。5.1 训练不稳定与发散问题现象损失值出现NaN或者mAP/Rank-1指标在域自适应阶段不升反降剧烈震荡。可能原因1伪标签噪声过大。初期模型在目标域上性能很差生成的伪标签错误率极高用这些错误标签进行监督学习会导致模型崩溃。排查与解决可视化特征使用t-SNE或PCA将目标域特征和源域质心可视化。如果目标域特征一团糟且与任何质心都不靠近说明伪标签不可信。实施严格的伪标签筛选提高置信度阈值只使用最可靠的样本例如与最近质心余弦相似度大于0.8的。在训练初期可以只使用非常少的“高置信度”样本随着模型变好逐步放宽阈值。采用软标签或标签平滑不要使用硬性的one-hot伪标签可以基于与多个最近质心的距离分布分配一个软标签概率分布。可能原因2源域与目标域损失不平衡。λ参数设置过大导致目标域的噪声损失主导了训练。排查与解决监控两个损失的数值。在训练初期源域损失应远大于目标域损失。确保λ从一个很小的值如0.1开始并缓慢增长。使用梯度裁剪torch.nn.utils.clip_grad_norm_防止梯度爆炸。可能原因3质心更新过于激进。动量系数momentum设置得太小导致质心被当前batch的噪声样本带偏。排查与解决将momentum设置为一个非常接近1的值如0.999。这保证了质心的更新是平滑、缓慢的继承了历史信息的“惯性”。5.2 性能提升不明显问题现象相比不使用多质心的基线方法如直接使用源域模型或简单伪标签方法mAP和Rank-1提升有限2%。可能原因1K值设置不当。K太小表征能力不足K太大对于样本数少的ID容易过拟合且增加了匹配的模糊性。排查与解决在验证集可以从目标域划分一小部分上尝试不同的K值2,3,4,5,6。观察性能曲线。通常存在一个最优值。可能原因2质心初始化效果差。如果源域某个ID的样本本身多样性不足K-Means初始化出的质心可能没有意义。排查与解决检查源域数据每个ID的图片数量。对于图片数少于K的ID可以采用数据增强如复制、镜像来生成更多样本或者直接减少该ID的K值。也可以尝试用可学习参数初始化让网络自己学习质心。可能原因3特征提取骨干网络能力不足。多质心是“锦上添花”如果基础特征提取能力弱效果也有限。排查与解决确保源域预训练模型已经达到该数据集的SOTA或接近SOTA的水平。可以考虑使用更强的骨干网络如ResNet-101, IBN-Net-a或引入非局部注意力Non-local、通道注意力SE等模块提升特征质量。5.3 计算与内存开销问题现象训练速度明显变慢GPU内存占用激增。可能原因1质心匹配计算量大。在为目标域样本分配伪标签时需要计算该样本与所有num_classes * K个质心的距离。优化策略批次计算利用矩阵运算一次性计算一个batch的目标特征与所有质心的距离避免循环。近似最近邻搜索当源域类别数很大时如MSMT有1000多类可以考虑使用FAISS库进行高效的最近邻搜索大幅加速。缓存质心质心在训练过程中变化缓慢可以每隔几个iteration才计算一次伪标签而不是每个iteration都计算。可能原因2多质心存储开销。质心张量大小为[C, K, D]如果C700, K4, D2048存储为float32则占用内存约700*4*2048*4 ≈ 22 MB可以接受。但如果C很大如MSMT则需要留意。5.4 实际部署考量当模型训练好后部署到实际摄像头系统中还需要考虑推理速度在推理时我们只需要使用训练好的特征提取器。多质心只是在训练阶段用于提供监督信号的“工具”不会增加推理时的计算负担。这是该方法的一大优势。特征库构建对于目标域新摄像头我们需要构建一个画廊gallery特征库。直接用训练好的模型提取所有画廊图片的特征即可。查询匹配当有一个查询query图片时提取其特征然后计算与画廊特征库中所有特征的余弦相似度排序返回最相似的结果。这里完全不需要用到“多质心”。因为多质心已经将其知识蒸馏到了网络参数中网络提取的特征本身就具备了强大的跨域判别能力。从我多次实验的经验来看成功应用多质心方法的关键在于“稳”。初期务必通过严格的伪标签筛选和保守的超参数小λ高动量保证训练稳定性。在模型对目标域有了一定适应性后再逐步引入更强的无监督信号。多质心不是银弹它需要与高质量的数据预处理、强大的骨干网络、以及精心设计的训练策略相结合才能发挥出最大威力。它更像是一个“特征表征的增强器”让模型学会用一组点而不是一个点去思考一个人的身份这种思维方式上的改变正是其应对复杂跨域场景的智慧所在。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门