ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

如何用 ai-engineering-hub 的 Siamese Network 判断两张 MNIST 图片是否为同一数字?

如何用 ai-engineering-hub 的 Siamese Network 判断两张 MNIST 图片是否为同一数字? 如何用 ai-engineering-hub 的 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-hubai-engineering-hub 仓库的siamese-network目录实现了一个在 MNIST 数据集上的 Siamese Networksiamese-network/README.md 对其目标的描述是 detect if two images are of the same digit即判断两张 MNIST 手写数字图片是否属于同一个数字。整个流程集中在 Siamese-Network.ipynb 这一个 notebook 中构造同/不同数字的配对数据集、定义共享权重的双输入网络、用对比损失训练 5 个 epoch最后在测试集图片对上输出相似度分数并保存成对图片供核对。跑通它的结果是终端打印 4 个相似度分数工作目录生成image_0.jpeg到image_3.jpeg四张成对图片。注意一个硬前提notebook 中模型与张量都无条件调用.cuda()因此必须在带 CUDA GPU 的机器上运行MNIST 数据依赖网络自动下载。准备条件运行环境与数据下载一个能运行.ipynb的 Jupyter 环境jupyter 或任何 ipynb 编辑器均可。依赖即 notebook 第一个代码单元导入的库torch、torchvision、numpy、matplotlib。安装方式按你本地的包管理工具完成即可notebook 未指定版本。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数据不需要手动准备数据集单元用MNIST(root./data, trainTrue, downloadTrue)和MNIST(root./data, trainFalse, downloadTrue)加载首次运行时会把 MNIST 自动下载到相对工作目录的./data下需要网络可用。训练与推理的 DataLoader 都使用了num_workers8请保证运行机器支持该进程数。构造 MNIST 同数字/不同数字的配对数据集Siamese 任务需要的是图片对而不是单张图片。notebook 用SiameseDataset包装标准 MNIST对每一张图imgA用random.randint(0, 1)随机决定配对类型——同一 flag 分支下从训练集中随机重选直到imgB的标签与imgA相同同数字对另一分支下重选直到标签不同不同数字对两类对随机混在一起。返回的第三个值是torch.tensor([(labelA ! labelB)], dtypetorch.float32)即1.0 表示两张图不是同一数字0.0 表示是同一数字这个标签方向后面对比损失会用到。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)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)这里只有一个转换transforms.ToTensor()没有归一化或随机增强照跑即可。定义共享权重的 Siamese 网络网络结构是双塔共享权重forward_once把单张单通道图片依次过三层卷积64、128、256 通道层间为 ReLU MaxPool2d(2)展平后经三层全连接256*3*3 → 1024 → 256 → 2最终得到2 维嵌入向量forward(inputA, inputB)只是让两张图分别走同一套forward_once返回outputA和outputB两个嵌入。是否同一数字最终体现为这两个嵌入之间的距离。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训练对比损失与 5 个 epoch损失是ContrastiveLossmargin2.0对两个嵌入取欧氏距离label为 0同数字时惩罚距离的平方label为 1不同数字时惩罚clamp(margin - distance, min0)的平方与上面SiameseDataset的标签方向一一对应。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)必须指出一个执行前需要处理的问题原 notebook 中F.pairwise_distance(output1, output2, ...)引用的output1、output2是未定义变量forward的形参名是outputA、outputB原样运行该单元格会抛NameError。要把output1、output2改成outputA、outputB后即可运行这是该单元格唯一需要相对原代码的改动其余部分保持原样。训练配置与循环如下batch_size64、num_workers8、Adam 优化器lr0.001、共 5 个 epoch每轮结束打印一行Epoch {epoch}; Loss {total_loss}——这是文档中给出的训练进度观察方式逐 epoch 各打印一条。train_dataloader DataLoader(siamese_train, shuffleTrue, num_workers8, batch_size64) net SiameseNetwork().cuda() criterion ContrastiveLoss() optimizer optim.Adam(net.parameters(), lr 0.001 )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})判断结果相似度分数与成对图片训练完成后notebook 用测试集抽样成对来做是否同一数字的判断展示visualize_siamese_pairs取前 4 个 batchtotal_images4测试集 DataLoader 的batch_size1即 4 对图片对每对计算两个嵌入的欧氏距离然后按文档给出的公式计算相似度分数similarity_score torch.exp(-euclidean_distance)打印保留 4 位的分数并把图片对保存为image_0.jpeg至image_3.jpegdpi300同时弹窗显示。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)运行后的核对方式就是这一单元格自身的输出终端会依次打印 4 个相似度分数工作目录生成 4 张 jpeg。打开这些图片肉眼比对两张手写数字是否相同再对照分数——按公式exp(-distance)嵌入距离为 0 时分数为 1距离越大分数越低。需要明确的是文档没有给出同一数字的固定分数阈值也没有提供传入任意两张指定图片的交互入口判断发生在从测试集抽样的图片对上核对以分数 成对图片的形式完成不要自行补充固定判定线。限制与边界全流程依赖 CUDA GPU.cuda()无条件调用无 GPU 环境跑不起来。ContrastiveLoss单元格的output1/output2变量名问题不修复则无法进入训练。DataLoader 使用num_workers8在受限环境中如遇到 worker 相关问题这是文档中给出的可调参数。该实现的目标范围就是 MNIST 上成对图片的同数字检测见 siamese-network/README.mdnotebook 未包含迁移到其他数据集或部署服务的步骤。【免费下载链接】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),仅供参考
返回列表