ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

pytorch-metric-learning 官方 Colab 示例全解析:从 TripletMarginLoss 到分布式训练与推理部署

pytorch-metric-learning 官方 Colab 示例全解析:从 TripletMarginLoss 到分布式训练与推理部署 人工智能机器学习深度学习计算机视觉【免费下载链接】pytorch-metric-learningThe easiest way to use deep metric learning in your application. Modular, flexible, and extensible. Written in PyTorch.项目地址https://gitcode.com/gh_mirrors/py/pytorch-metric-learning点击查看免费下载本文以examples/README.md即 examples 目录 下的官方示例指南为主线系统梳理 pytorch-metric-learning 提供的 11 个可直接运行的 Colab Notebook其中 4 个聚焦单一组件loss、miner、cross-batch memory、分布式7 个展示带日志记录与模型保存的完整训练/测试工作流并延伸到训练完成后的推理部署。读完本文你将掌握在 GPU Colab 环境中快速启动 metric learning 实验的完整路径理解losses、miners、samplers、trainers、testers、AccuracyCalculator、logging_presets与inference模块之间的真实调用关系并能够把这些示例迁移到自己的数据集上。运行前提Colab 环境准备所有 Notebook 均设计为在Google Colab上运行。原文档明确指出两点前置操作将运行时类型设置为GPU通过菜单Runtime→Change runtime type选择 GPU点击 Colab 页面顶部的Open in playground即可交互式运行 Notebook无需先保存到自己的 Drive。安装依赖是每个 Notebook 的第一步。简单示例如 TripletMarginLossMNIST.ipynb使用两条命令!pip install pytorch-metric-learning !pip install faiss-gpu # 分布式示例中按需替换为 faiss-cpu而展示完整工作流的 Notebook如 MetricLossOnly.ipynb统一采用!pip install -q pytorch-metric-learning[with-hooks] !pip install umap-learn其中[with-hooks]是setup.py中声明的 extras 依赖见 setup.py用于启用 logging_presets 的完整日志与模型保存能力umap-learn用于在测试阶段对嵌入空间做降维可视化。若缺少 record-keeper / tensorboardget_record_keeper会退化为空容器并打印警告此时将没有日志和模型保存对应 logging_presets.py 的降级逻辑。一、简单示例只使用某个损失或挖掘器原文档将前 4 个 Notebook 归为 Simple examples定位是在你自己代码里单独使用某个 loss 或 miner的参考不涉及完整训练框架。1. MNIST TripletMarginLoss AccuracyCalculatorTripletMarginLossMNIST.ipynb 是最直观的入门示例完整走通手工训练循环 测试评估from pytorch_metric_learning import distances, losses, miners, reducers, testers from pytorch_metric_learning.utils.accuracy_calculator import AccuracyCalculator distance distances.CosineSimilarity() reducer reducers.ThresholdReducer(low0) loss_func losses.TripletMarginLoss(margin0.2, distancedistance, reducerreducer) mining_func miners.TripletMarginMiner(margin0.2, distancedistance, type_of_tripletssemihard) accuracy_calculator AccuracyCalculator(include(precision_at_1,), k1)训练循环展示了一个重要机制miner 先于 loss 运行miner 输出的indices_tuple再传给 loss——embeddings model(data) indices_tuple mining_func(embeddings, labels) loss loss_func(embeddings, labels, indices_tuple) loss.backward() optimizer.step()从源码看TripletMarginLoss会把indices_tuple通过lmu.convert_to_triplets转成三元组anchor/positive/negative 索引然后计算margin(ap_dists, an_dists)的违规量取relu(violation)作为损失默认 reducer 为AvgNonZeroReducer见 triplet_margin_loss.py。TripletMarginMiner则依据距离差ap_dist - an_dist在CosineSimilarity这类 inverted 距离下取反与 margin 的关系筛选三元组type_of_triplets可选all、hard、semihard或easy见 triplet_margin_miner.py。示例还打印mining_func.num_triplets用于监控挖掘出的三元组数量。评估阶段用testers.BaseTester().get_all_embeddings()抽取全量嵌入再交给AccuracyCalculator计算precision_at_1。BaseTester默认会对嵌入做 L2 归一化normalize_embeddingsTrue见 base_tester.py。2. MNIST SubCenterArcFaceLoss 离群点查看SubCenterArcFaceMNIST.ipynb 使用losses.SubCenterArcFaceLoss训练亮点是训练结束后调用outliers, _ loss_func.get_outliers(train_embeddings, train_labels) print(fThere are {len(outliers)} outliers)get_outliers返回与主导 center 偏离超过阈值的样本索引Notebook 随后用imshow_many(dataset, outliers)批量展示这些离群样本帮助直观理解 sub-center 结构下哪些样本难以归位。3. MoCo on CIFAR10CrossBatchMemory 与自监督MoCoCIFAR10.ipynb 演示用CrossBatchMemory做 MoCo 风格的自监督对比学习。Notebook 说明选择 CrossBatchMemory 的三个理由队列与模型解耦模块化、可以用任意 tuple loss 而非局限于 InfoNCE/NTXent、可选配任意 pair/triplet miner 从动量编码器队列中挖难样本。核心配置loss_fn losses.CrossBatchMemory( losslosses.NTXentLoss(temperature0.1), embedding_size128, memory_size4096, )其余部分大量复用官方 MoCo 代码TwoCropsTransform、动量更新、kNN 监控等通过logging_presets.get_record_keeper记录 loss 与 kNN 精度。该 Notebook 自身记录了两个参考数字训练 200 epoch 后本 Notebook 精度 82.9%官方 MoCo CIFAR10 notebook 为 82.6%此为 Notebook 内记载的历史运行结果仅作参考对照。4. DistributedDataParallel 多进程训练DistributedTripletMarginLossMNIST.ipynb 展示如何用pytorch_metric_learning.utils.distributed在多进程DDP环境下训练。代码结构借鉴 PyTorch 官方分布式教程通过DataPartitioner把数据集切分给各进程再用torch.multiprocessing.Process启动多个 worker。Metric learning 特有的部分是把 miner 挖出的三元组索引做跨进程的 all-gather 对齐让每个进程都能看到全局 batch 的困难样本。该示例特意使用faiss-cpu因为纯 CPU 分布式环境足以跑通 MNIST。二、完整训练/测试工作流日志与模型保存原文档强调下面 7 个 Notebook 用于展示完整的训练/测试工作流若只想在自己的代码里单独使用某个 loss 或 miner请看上面的简单示例。它们普遍遵循同一个 4 步流程初始化模型trunk embedder、优化器和图像变换创建类不相交class-disjoint的 train/validation 划分初始化 loss、miner、sampler、trainer 和 tester训练模型、记录精度、绘制嵌入空间。通用脚手架类不相交的 CIFAR100 划分除 scRNAseq 与 Inference 外工作流 Notebook 都复用了同一段数据集代码——下载 CIFAR100 后按类别阈值前 50 类训练、后 50 类验证拼出类不相交的ClassDisjointCIFAR100数据集并用assert set(train_dataset.targets).isdisjoint(set(val_dataset.targets))强制校验。类不相交的意义在于验证集类别在训练中从未见过测的是嵌入空间的真实泛化能力而不是对训练类别的记忆。标准组件组合TripletMarginLoss MultiSimilarityMiner MPerClassSampler工作流示例中的多数MetricLossOnly、TrainWithClassifier、scRNAseq、DeepAdversarialMetricLearning采用相同的主力组合loss losses.TripletMarginLoss(margin0.1) miner miners.MultiSimilarityMiner(epsilon0.1) sampler samplers.MPerClassSampler( train_dataset.targets, m4, length_before_new_iterlen(train_dataset) ) batch_size 32 num_epochs 4 models {trunk: trunk, embedder: embedder} optimizers {trunk_optimizer: trunk_optimizer, embedder_optimizer: embedder_optimizer} loss_funcs {metric_loss: loss} mining_funcs {tuple_miner: miner}这里体现出 pytorch-metric-learning 的字典约定models必须有trunk必需键与embedder可省略缺省时自动补nn.Identity()loss_funcs用metric_loss键miner 放进mining_funcs的tuple_miner键。这些键名约束由KeyCheckerDict在 trainer 初始化时校验见 base_trainer.py。MPerClassSampler保证每个 batch 中每个类别恰好出现 m 个样本是 tuple 挖掘有效性的前提。trunk 的经典改法是去掉分类头trunk torchvision.models.resnet18(pretrainedTrue) trunk_output_size trunk.fc.in_features trunk.fc nn.Identity() # 用恒等函数替换 softmax 层 trunk torch.nn.DataParallel(trunk.to(device)) embedder torch.nn.DataParallel(MLP([trunk_output_size, 64]).to(device))embedder 把 trunk 输出的高维特征压成 64 维嵌入DataParallel适配多卡。优化器常设为 trunk 小学习率1e-5 embedder 大学习率1e-4图像变换使用带RandomResizedCrop与水平翻转的 ImageNet 归一化。日志与模型保存get_record_keeper get_hook_container所有工作流 Notebook 统一用 logging_presets 搭起记录 早停 画图三件套record_keeper, _, _ logging_presets.get_record_keeper(example_logs, example_tensorboard) hooks logging_presets.get_hook_container(record_keeper) dataset_dict {val: val_dataset} model_folder example_saved_models tester testers.GlobalEmbeddingSpaceTester( end_of_testing_hookhooks.end_of_testing_hook, visualizerumap.UMAP(), visualizer_hookvisualizer_hook, dataloader_num_workers2, accuracy_calculatorAccuracyCalculator(kmax_bin_count), ) end_of_epoch_hook hooks.end_of_epoch_hook( tester, dataset_dict, model_folder, test_interval1, patience1 )其中visualizer_hook用 nipy_spectral 色环按类别着色把 UMAP 降维后的嵌入画成散点图。end_of_epoch_hook每个 epoch 运行 tester 评测 val 集按primary_metric默认mean_average_precision_at_r判断是否刷新最佳精度并在model_folder下保存最新模型与best 模型两份 checkpoint见 logging_presets.pypatience1表示验证精度连续 1 个 epoch 未提升就提前停止patience_remaining逻辑见 logging_presets.py。随后可启动 TensorBoard%load_ext tensorboard %tensorboard --logdir example_tensorboard1. MetricLossOnly纯度量损失训练MetricLossOnly.ipynb 是最简工作流只用一个 metric loss。对应 trainer 源码极其精简——MetricLossOnly.calculate_loss依次完成compute_embeddings→maybe_mine_embeddings→maybe_get_metric_loss三步见 metric_loss_only.py。它也是理解BaseTrainer训练主循环train→forward_and_backward→calculate_loss→backward→step_optimizers见 base_trainer.py的最佳入口。2. scRNAseq_MetricEmbedding单细胞转录组度量嵌入scRNAseq_MetricEmbedding.ipynbUCSF Keiser 实验室作者2020-04 发布是跨界示例用 canonical single-cell RNAseq 细胞类型训练嵌入。它使用scanpy加载 Paul 2015 髓系祖细胞数据选择 10 个代表性 cluster 做训练/验证把其余 cluster 作为 holdout 集投影进学到的度量空间。模型是 fast.ai 风格的带 BN/Dropout 的EmbeddingNetemb_szs[1000,500,250,100]输出 25 维。评估上tester 设置use_trunk_outputTrue直接用 trunk 输出评估并借助hooks.get_loss_history()/hooks.get_accuracy_history()绘制损失曲线与 Adjusted Mutual InfoAMI曲线可视化改用 tSNE最后用簇中心原型prototype与 holdout 样本的pairwise_distances生成层级聚类热图。该示例证明了 pytorch-metric-learning 不局限于图像也适用于表格型高维数据。3. TrainWithClassifier度量损失 分类损失联合训练TrainWithClassifier.ipynb 在 trunk/embedder 之外增加classifier网络并用分类损失联合监督classifier torch.nn.DataParallel(MLP([64, 50])).to(device) # 50 类 models {trunk: trunk, embedder: embedder, classifier: classifier} loss_funcs {metric_loss: loss, classifier_loss: torch.nn.CrossEntropyLoss()} loss_weights {metric_loss: 1, classifier_loss: 0.5}loss_weights是可选的缺省时所有 loss 权重为 1见 base_trainer.py。这里classifier_loss的权重取 0.5体现度量为主、分类为辅的调参思路。分类器输出维度需与训练类别数50一致。4. CascadedEmbeddings多子网络级联与分块挖掘CascadedEmbeddings.ipynb 展示如何把多个子网络3 个不同规模的 trunkshufflenet_v2_x0_5 / x1_0 / resnet18封装成单一 trunk 与 embedder并对每段嵌入分别应用不同 lossloss0 losses.ContrastiveLoss(pos_margin0, neg_margin0.5) loss1 losses.MultiSimilarityLoss(alpha0.1, beta40, base0.5) loss2 losses.CircleLoss() mining_funcs {tuple_miner_1: miners.MultiSimilarityMiner(epsilon0.1), tuple_miner_2: miners.HDCMiner(filter_percentage0.25)}ListOfModels把多个子网络输出沿最后一维拼接torch.cat(outputs, dim-1)embedder 输出维度为64*3。trainer 通过embedding_sizes[64, 64, 64]告诉CascadedEmbeddings把嵌入切成 3 段第 i 段对应metric_loss_i。Notebook 还提示CascadedEmbeddings允许但不要求加分类网络与分类损失因此初始化时会看到loss_funcs is missing classifier_loss_[0-9]$的警告属正常现象。5. DeepAdversarialMetricLearning生成器制造难负样本DeepAdversarialMetricLearning.ipynb 使用一个 generator 在训练中制造难负样本。generator 输入维度要求为3 * trunk_output_size、输出维度为trunk_output_sizeMLP([3 * trunk_output_size, trunk_output_size, trunk_output_size], final_reluTrue)。配置上除了metric_loss还定义了synth_loss与g_adv_loss并可选配g_hard_loss、g_reg_loss这两个由 trainer 内部定义loss_weights { metric_loss: 1, synth_loss: 0.1, g_adv_loss: 0.1, g_hard_loss: 0.1, g_reg_loss: 0.1, } trainer trainers.DeepAdversarialMetricLearning( models, optimizers, batch_size, loss_funcs, train_dataset, mining_funcsmining_funcs, loss_weightsloss_weights, samplersampler, end_of_iteration_hookhooks.end_of_iteration_hook, end_of_epoch_hookend_of_epoch_hook, metric_alone_epochs1, g_alone_epochs1, g_triplets_per_anchor100, )metric_alone_epochs/g_alone_epochs控制度量网络与生成器分别单独训练的预热轮数g_triplets_per_anchor控制生成难样本时每个 anchor 采样的三元组数量。同样地缺少classifier_loss的警告是允许的。6. TwoStreamMetricLoss双流数据集TwoStreamMetricLoss.ipynb 面向 anchor 与 positive/negative 来自不同来源的双流场景。它用CIFAR100TwoStreamDataset把数据按 80%/20% 切成 anchor 流与 posneg 流__getitem__返回(anchor, posneg, target)——anchor 从第一流取posneg 从第二流中随机选同类的样本loss losses.TripletMarginLoss(margin0.2) miner miners.TripletMarginMiner(margin0.2) tester testers.GlobalTwoStreamEmbeddingSpaceTester(...) trainer trainers.TwoStreamMetricLoss(models, optimizers, batch_size, loss_funcs, train_dataset, mining_funcsmining_funcs, samplersampler, ...)对应的 tester 换成GlobalTwoStreamEmbeddingSpaceTester其 visualizer_hook 会以不同点形区分两个流小点 第一流 anchor大方块 第二流 posneg直观检查两类嵌入是否对齐。7. Inference训练后的推理部署Inference.ipynb 独立于训练工作流专门演示训练完成后的检索与匹配。它加载 CIFAR10 上预训练的 resnet20去掉最后的线性层然后用inference模块封装from pytorch_metric_learning.distances import CosineSimilarity from pytorch_metric_learning.utils.inference import InferenceModel, MatchFinder match_finder MatchFinder(distanceCosineSimilarity(), threshold0.7) inference_model InferenceModel(model, match_findermatch_finder)随后依次演示 5 类操作inference_model.train_knn(dataset)用 faiss 建立索引FaissKNN为默认knn_func见 inference.pyget_nearest_neighbors(img, k10)查询并返回 10 个最近邻样本及其距离is_match(x, y)判断两张图是否同类相似度 threshold 判为匹配get_matches(x)批量计算 batch 内全对匹配矩阵get_matches(x, y, return_tuplesTrue)query 与 reference 两组样本间的匹配元组列表Notebook 把 threshold 调高到 0.95 以减少误配。MatchFinder的阈值方向由距离的is_inverted属性决定相似度类距离如CosineSimilarity是dist threshold判匹配距离类则反之见 inference.py。三、示例矩阵速查表Notebook核心组件场景/说明TripletMarginLossMNISTTripletMarginLossTripletMarginMinerAccuracyCalculator手写训练循环 精度评估的最简入门SubCenterArcFaceMNISTSubCenterArcFaceLoss训练 离群样本查看MoCoCIFAR10CrossBatchMemoryNTXentLoss自监督对比学习MoCo 风格DistributedTripletMarginLossMNISTutils.distributed DDP多进程分布式训练MetricLossOnlyMetricLossOnlytrainer纯度量损失的标准工作流scRNAseq_MetricEmbeddingMetricLossOnly scanpy单细胞转录组度量嵌入 holdout 投影TrainWithClassifierTrainWithClassifiertrainer度量损失 分类损失联合CascadedEmbeddingsCascadedEmbeddingstrainer多子网络级联、分段挖掘与分段损失DeepAdversarialMetricLearningDeepAdversarialMetricLearningtrainer生成器制造难负样本TwoStreamMetricLossTwoStreamMetricLossGlobalTwoStreamEmbeddingSpaceTester双流数据集anchor 与 posneg 异源InferenceInferenceModelMatchFinder faiss训练后的最近邻检索与配对判定四、迁移到自有数据的建议综合上述示例将整套流程迁移到自己的数据集通常只需改动三处数据层把ClassDisjointCIFAR100替换为自定义torch.utils.data.Dataset保证返回(样本, 标签)若属双流场景则按CIFAR100TwoStreamDataset的模式返回(anchor, posneg, target)模型层trunk换成任意特征提取主干去掉分类头、输出接nn.Identity()embedder的输入维度改为trunk.fc.in_features输出维度即目标嵌入维度组件层按任务特点选 loss分类数多且标注充足可加TrainWithClassifier自监督用CrossBatchMemory、minerMultiSimilarityMiner或TripletMarginMiner、sampler有标签用MPerClassSampler然后沿用字典打包 trainer tester hooks的标准骨架即可获得带日志、早停、模型保存和嵌入可视化的完整实验闭环。赞分享人工智能机器学习深度学习计算机视觉【免费下载链接】pytorch-metric-learningThe easiest way to use deep metric learning in your application. Modular, flexible, and extensible. Written in PyTorch.项目地址https://gitcode.com/gh_mirrors/py/pytorch-metric-learning点击查看免费下载相关推荐DeepSeek-V3-Base模型部署成本计算器云服务与本地部署经济性对比DeepSeek V3 Base模型部署成本计算器云服务与本地部署经济性对比 引言大模型部署的成本困境 你是否正面临这样的困境想要部署性能强大的DeepS人工智能机器学习深度学习计算机视觉Ludwig 官方示例全解析从单机训练到分布式调参的 Python API 实战指南Ludwig 官方示例全解析从单机训练到分布式调参的 Python API 实战指南 本篇指南以 Ludwig 仓库 examples 目录为线索系统梳理该人工智能深度学习机器学习大模型预训练微调LoRA多模态NLP计算机视觉模型推理服务PyTorch Metric Learning推理模型高效部署与实时预测解决方案PyTorch Metric Learning推理模型高效部署与实时预测解决方案 PyTorch Metric Learning是一个强大的深度学习度量学习库人工智能机器学习深度学习计算机视觉上一篇网盘直链下载助手完全指南如何让百度云盘、阿里云盘等8大平台下载速度显著提升下一篇ZenlessZoneZero-OneDragon绝区零全自动游戏助手终极指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进