ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

ai-engineering-hub 实战:基于 PyTorch 从零实现孪生网络(Siamese Network),用 MNIST 判断两张手写数字图像是否属于同一数字

2026/9/10 11:45:17 拓冰建站 浏览量
ai-engineering-hub 实战:基于 PyTorch 从零实现孪生网络(Siamese Network),用 MNIST 判断两张手写数字图像是否属于同一数字 ai-engineering-hub 实战基于 PyTorch 从零实现孪生网络Siamese Network用 MNIST 判断两张手写数字图像是否属于同一数字【免费下载链接】ai-engineering-hubIn-depth tutorials on LLMs, RAGs and real-world AI agent applications.项目地址: https://gitcode.com/GitHub_Trending/ai/ai-engineering-hub本篇技术指南以 siamese-network 项目为主体讲解如何在 MNIST 手写数字数据集上从零实现并训练一个孪生网络Siamese Network用于判断任意两张输入图像是否为同一数字。读者学完本文后将掌握图像对数据集构造、权重共享孪生网络架构、对比损失Contrastive Loss的原理与实现以及用欧氏距离导出相似度分数的完整实战流程并可直接在 Siamese-Network.ipynb 中按序复现全部代码。一、核心问题孪生网络要解决什么传统分类模型解决的是这张图是哪个数字而孪生网络解决的是这两张图是不是同一个数字。后者属于度量学习Metric Learning范畴模型不再直接输出类别标签而是学习一个嵌入Embedding映射把输入图像映射到特征空间中的向量使得同类样本在空间中距离近、异类样本距离远。本项目在 MNIST 上验证这一思路训练完成后给模型任意两张手写数字图像它输出一个 01 之间的相似度分数接近 1 表示同一数字接近 0 表示不同数字。这种能力在现实世界有广泛应用方向例如人脸验证、签名核验、指纹比对以及只需要每类少量样本即可工作的一次学习one-shot / few-shot learning场景。二、环境与依赖导入项目的全部实现位于 Siamese-Network.ipynb依赖 PyTorch 生态与常规科学计算库import torch import torch.nn as nn import torch.nn.functional as F import numpy as np import random import torchvision.transforms as transforms import matplotlib.pyplot as plt from torch.utils.data import DataLoader, Dataset from torchvision.datasets import MNIST from torch import optim依赖用途torch / torch.nn / torch.nn.functional张量计算、网络层定义与激活/距离函数如F.pairwise_distance、F.pad等torchvisiondatasets、transforms加载 MNIST、图像张量化numpy / random数值处理、图像对随机采样matplotlib.pyplot训练后可视化图像对与相似度运行前提需要安装torch、torchvision、numpy、matplotlibNotebook 中网络被显式放在 CUDA 上net SiameseNetwork().cuda()因此默认假设存在可用 GPU纯 CPU 环境需将.cuda()改为.to(device)device cuda if torch.cuda.is_available() else cpu。三、数据集构造把 MNIST 改造成图像对孪生网络的训练样本不再是单张图像而是图像对 配对标签。本项目自定义了SiameseDataset在每次取样本时随机生成一个同类对或异类对class SiameseDataset(Dataset): def __init__(self, data, transformNone): self.data data self.transform transform def __getitem__(self, index): imgA, labelA self.data[index] same_class_flag random.randint(0, 1) if same_class_flag: labelB -1 while labelB ! labelA: imgB, labelB random.choice(self.data) else: labelB labelA while labelB labelA: imgB, labelB random.choice(self.data) if self.transform: imgA self.transform(imgA) imgB self.transform(imgB) return imgA, imgB, torch.tensor([(labelA ! labelB)], dtypetorch.float32) def __len__(self): return len(self.data)关键设计点逐条拆解配对策略same_class_flag random.randint(0, 1)以 1/2 概率决定构造同类对还是异类对。同类对通过循环不断从全量数据中随机抽取直到labelB labelA异类对则抽到labelB ! labelA为止。注意两个分支的初始值-1与labelA只是保证循环至少执行一次取样的技巧。标签语义返回的配对标签是(labelA ! labelB)的布尔值转float32——同类对为 0.0异类对为 1.0。这个约定与后续对比损失的公式直接对应理解错误会导致损失函数行为颠倒。数据变换transform通过transforms.Compose([transforms.ToTensor()])传入将 PIL 图像转为[1, 28, 28]的张量。本项目未做旋转、缩放等数据增强这一点在后续改进小节会讨论。随后用官方 MNIST 接口一次性加载训练集与测试集并包装成孪生数据集mnist_train MNIST(root./data, trainTrue, downloadTrue) mnist_test MNIST(root./data, trainFalse, downloadTrue) transform transforms.Compose([transforms.ToTensor()]) siamese_train SiameseDataset(mnist_train, transform) siamese_test SiameseDataset(mnist_test, transform)downloadTrue会在首次运行时自动把 MNIST 下载到./data目录trainTrue/False分别对应 6 万张训练样本与 1 万张测试样本。值得注意__getitem__中的循环抽样在每次迭代都做全量随机选择属于朴素但直观的配对方式当数据规模更大时可改用预先构建固定配对表 分桶抽样来提升效率并保证同类/异类对的比例稳定。四、网络架构权重共享的卷积孪生网络孪生网络的核心特征是两个分支共享同一套参数weight sharing两张图像分别经过同一个特征提取网络得到各自的嵌入向量再在嵌入空间计算距离。共享权重意味着模型学到的是通用的相似度衡量能力而不是两套独立特征。本项目用 CNN 作为共享主干class SiameseNetwork(nn.Module): def __init__(self): super(SiameseNetwork, self).__init__() self.cnn nn.Sequential( nn.Conv2d(1, 64, kernel_size5, stride1, padding2), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, stride2), nn.Conv2d(64, 128, kernel_size5, stride1, padding2), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, stride2), nn.Conv2d(128, 256, kernel_size3, stride1, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, stride2) ) self.fc nn.Sequential( nn.Linear(256 * 3 * 3, 1024), nn.ReLU(inplaceTrue), nn.Linear(1024, 256), nn.ReLU(inplaceTrue), nn.Linear(256, 2) ) def forward_once(self, x): output self.cnn(x) output output.view(output.size()[0], -1) output self.fc(output) return output def forward(self, inputA, inputB): outputA self.forward_once(inputA) outputB self.forward_once(inputB) return outputA, outputB架构要点与逐层尺寸推演输入 MNIST 单通道 28×28层输出尺寸说明Conv2d(1, 64, k5, p2) ReLU64×28×28padding2 保持空间尺寸不变MaxPool2d(2)64×14×14下采样减半Conv2d(64, 128, k5, p2) ReLU128×14×14通道翻倍尺寸不变MaxPool2d(2)128×7×7下采样Conv2d(128, 256, k3, p1) ReLU256×7×7第三层卷积改用 3×3 卷积核MaxPool2d(2)256×3×37×7 经步长 2 池化后取整为 3×3展平 view(-1)256×3×3 2304送入全连接层Linear(2304, 1024) ReLU1024第一层全连接Linear(1024, 256) ReLU256第二层全连接Linear(256, 2)2最终嵌入为 2 维向量可以看到本项目最终把每张图像压缩为仅 2 维的嵌入向量——这正是便于后续F.pairwise_distance直接计算欧氏距离并可视化的设计2 维空间中的距离直观可解释。forward_once负责单张图像的嵌入计算forward则对两张输入分别调用forward_once体现了同一网络、两份输入、共享权重的孪生结构。五、对比损失训练孪生网络的驱动信号孪生网络的训练目标是让同类嵌入靠近、异类嵌入远离这由**对比损失Contrastive Loss**实现class ContrastiveLoss(torch.nn.Module): def __init__(self, margin2.0): super(ContrastiveLoss, self).__init__() self.margin margin def forward(self, outputA, outputB, label): euclidean_distance F.pairwise_distance(output1, output2, keepdim True) same_class_loss (1-label) * (euclidean_distance**2) diff_class_loss (label) * (torch.clamp(self.margin - euclidean_distance, min0.0)**2) return torch.mean(same_class_loss diff_class_loss)损失函数背后的数学形式为$$L \frac{1}{N}\sum_{i}\left[(1-y_i)\cdot d_i^2 y_i\cdot \max(0, m - d_i)^2\right]$$其中 $d_i |f(x_A^{(i)}) - f(x_B^{(i)})|_2$ 是两张图像嵌入向量的欧氏距离$y_i$ 是配对标签同类为 0、异类为 1$m$ 是 margin 超参数默认margin2.0。两项的含义同类项 $(1-y)d^2$当两张图属于同一数字时$y0$惩罚它们嵌入距离的平方迫使距离趋于 0异类项 $y\cdot\max(0, m-d)^2$当两张图不同数字时$y1$只有距离小于 margin 才产生惩罚鼓励距离至少拉开到 $m$ 以上距离超过 margin 后损失为 0不再继续推远——这正是 margin 的作用为异类距离划定一个合理的下界避免模型过度放大所有距离。F.pairwise_distance默认使用 $p2$ 的欧氏距离torch.clamp(..., min0.0)实现 $\max(0, \cdot)$ 的截断。从仓库源码看的一个坑Notebook 中ContrastiveLoss.forward的形参名为outputA/outputB而函数体内F.pairwise_distance引用的是output1/output2二者不一致。从源码内容看直接逐字执行会触发NameError建议在运行前将函数体统一改为outputA/outputB这与训练 cell 中criterion(outputA, outputB, label)的调用保持一致。这也是复现 Notebook 时最需要留意的一处细节。六、训练流程与超参配置数据加载、网络、损失与优化器配置如下train_dataloader DataLoader(siamese_train, shuffleTrue, num_workers8, batch_size64) net SiameseNetwork().cuda() criterion ContrastiveLoss() optimizer optim.Adam(net.parameters(), lr 0.001)训练循环对每个 batch 完成前向 → 计算对比损失 → 反向传播 → 参数更新的标准流程并按 epoch 累计打印损失for epoch in range(5): total_loss 0 for imgA, imgB, label in train_dataloader: imgA, imgB, label imgA.cuda(), imgB.cuda(), label.cuda() optimizer.zero_grad() outputA, outputB net(imgA, imgB) loss_contrastive criterion(outputA, outputB, label) loss_contrastive.backward() total_loss loss_contrastive.item() optimizer.step() print(fEpoch {epoch}; Loss {total_loss})本项目训练超参汇总超参数取值说明优化器Adam自适应学习率优化器学习率 lr0.001Adam 默认学习率batch_size64训练 DataLoader 的批大小num_workers8数据加载并行进程数需不超过 CPU 核数epochs5训练轮数margin2.0对比损失的异类距离下界设备CUDAnet、imgA/imgB/label均显式.cuda()Notebook 中记录的 5 个 epoch 累计损失输出为Epoch累计 Loss0294.241102.85263.37344.62432.78可以看到损失从约 294 快速下降至约 33前几个 epoch 收敛明显说明孪生网络与对比损失在 MNIST 上训练稳定、收敛迅速——这与同类距离拉近、异类距离拉开的学习目标一致。七、推理评估欧氏距离与相似度分数训练完成后用测试集构造单样本 batchbatch_size1进行验证。评估思路对测试集中的每一对图像用训练好的网络分别计算嵌入再求两者欧氏距离并进一步转换为直观的相似度分数test_dataloader DataLoader(siamese_test, shuffleTrue, num_workers8, batch_size1) def show_image_pair(imgA, imgB, label, similarity_score, i): fig, ax plt.subplots(1, 2, figsize(4, 4)) ax[0].imshow(imgA.squeeze(), cmapgray) ax[0].set_title(Image 1) ax[1].imshow(imgB.squeeze(), cmapgray) ax[1].set_title(Image 2) plt.savefig(fimage_{i}.jpeg, bbox_inchestight, dpi 300) print(similarity_score) plt.show() def visualize_siamese_pairs(data_loader, total_images4): for idx, batch in enumerate(data_loader): if idx total_images: return imgA, imgB, label batch outputA, outputB net(imgA.cuda(), imgB.cuda()) euclidean_distance F.pairwise_distance(outputA, outputB) similarity_score torch.exp(-euclidean_distance) imgA imgA[0].numpy() imgB imgB[0].numpy() label label[0].item() show_image_pair(imgA, imgB, label, round(similarity_score.item(), 4), idx) visualize_siamese_pairs(test_dataloader, total_images4)相似度分数的设计值得单独说明similarity exp(-d)将欧氏距离 $d \in [0, \infty)$ 单调映射到相似度 $(0, 1]$——距离为 0 时相似度为 1完全一致距离越大相似度越趋近于 0。相比直接输出距离exp(-d)的形式更符合人类直觉也便于设定阈值做二分类例如相似度 0.5 判定为同一数字。Notebook 中 4 组测试对的输出分数为测试对相似度分数第 1 对0.9199第 2 对0.0439第 3 对0.9008第 4 对0.0399从分数分布可以清晰推断模型的行为模式同一数字的配对得到接近 1 的高相似度如 0.92、0.90不同数字的配对得到接近 0 的低相似度如 0.04二者之间存在巨大的分数间隔说明经过 5 个 epoch 的训练模型已学会将相同数字映射到嵌入空间中的邻近区域、将不同数字推到远处。show_image_pair还会把每对图像以 300 dpi 保存为image_0.jpegimage_3.jpeg方便后续人工核验。八、复现指引在仓库根目录下按以下步骤即可复现本项目打开 siamese-network/Siamese-Network.ipynbSiameseNetwork的实现入口确保环境已安装torch、torchvision、numpy、matplotlib按 cell 顺序从上到下依次运行导入依赖 → 定义SiameseDataset→ 加载 MNIST → 定义SiameseNetwork→ 定义ContrastiveLoss→ 配置 DataLoader/优化器 → 训练 5 个 epoch → 可视化评估首次运行会自动下载 MNIST 到./data需要网络连接训练约 5 个 epoch 后观察测试对相似度分数的分布约 0.9 与约 0.04 两个簇若为纯 CPU 环境需将训练与评估 cell 中的.cuda()替换为.to(device)device cuda if torch.cuda.is_available() else cpu若运行ContrastiveLoss报NameError按第五节说明统一F.pairwise_distance的变量名。整体链路可概括为MNIST 图像 → 随机配对SiameseDataset→ 共享权重 CNN 嵌入SiameseNetwork.forward_once→ 欧氏距离F.pairwise_distance→ 对比损失ContrastiveLoss→ 反向传播 → 推理时以exp(-d)输出相似度。九、从 MNIST 走向真实世界延伸应用与改进方向孪生网络的价值在于它把分类问题转化为比较问题因此特别适合类别很多、每类样本很少甚至类别在训练时未出现的场景。基于本项目可以自然延伸的方向包括人脸/签名/指纹验证注册时保存特征嵌入验证时对比现场采集样本与注册样本的嵌入距离超过阈值即拒绝——这是孪生网络最经典的生产应用一次学习one-shot / few-shot learning训练集只含每类极少量样本时孪生网络仍能通过相似 vs 不相似的二元信号学会通用判别能力文本/多模态相似度把共享 CNN 换成共享的 Transformer 编码器孪生 BERT 结构即可将同一套双塔 对比损失范式迁移到语义相似度、检索排序、去重等 LLM/RAG 工程场景中——这与 ai-engineering-hub 仓库中大量 RAG 与 Agent 实战教程的检索目标一脉相承。结合本仓库实现可尝试的改进点包括改进方向说明margin 调参margin2.0是默认值margin 过大可能过度挤压特征空间过小则异类区分不充分可做网格搜索数据增强当前transform仅ToTensor()加入随机旋转、平移、噪声可显著提升对书写变体的泛化嵌入维度当前输出 2 维嵌入便于可视化真实应用通常提高到 64256 维以容纳更丰富的判别信息损失函数升级可对比三元组损失Triplet Loss、Circle Loss 等观察收敛速度与嵌入质量的差异阈值设定生产落地时需在验证集上确定相似度阈值平衡误接受率与误拒绝率通过本文的完整拆解读者不仅掌握了孪生网络在 MNIST 上的可复现实现更理解了嵌入 距离 对比损失这一度量学习范式——它既是人脸验证类应用的基石也是现代向量检索与 RAG 系统中把相似性变成距离度量的核心思想来源。【免费下载链接】ai-engineering-hubIn-depth tutorials on LLMs, RAGs and real-world AI agent applications.项目地址: https://gitcode.com/GitHub_Trending/ai/ai-engineering-hub创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考