ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

Java图片相似度计算实战:DJL+ResNet50特征提取与检索

Java图片相似度计算实战:DJL+ResNet50特征提取与检索 简介一套基于Java深度学习库DJL和ImageNet预训练ResNet50模型的图片相似度计算工程面向需要实现图像检索、人脸识别、重复图片排查或版权保护的Java开发者。工程支持提取图片512维特征并采用余弦相似度、欧氏距离等度量方式完成1:1比对直接输出相似度分数附带可直接运行的工程源码、ResNet50模型文件、测试图片及相关配置下载后即可体验。压缩包共292个文件约756.59MB主要包含80个jpg与12个jpeg测试图片、34个xml工程配置、6个java源码、6个class编译产物、3个md说明文档以及模型文件和多类型资源目录结构清晰完整。已有434人学习下载。借助DJL统一封装开发者无需深入底层即可在Java生态中快速集成深度学习能力便捷构建图像搜索、相似度分析等应用适合作为课程设计、毕业设计或生产环境二次开发的基础范本实测效果可参考配套博客。1. 用 Java 做图片相似度计算为什么我把赌注压给 DJL图片相似度计算在 Java 后端里并不冷门商品图去重、相册聚类、违规图片比对甚至提交表单时的重复图片校验都是这一需求的具体场景。很多人第一反应是调 Python 服务但多一个 Python 进程就多一套部署、监控和权限边界而用 Deep Java LibraryDJL可以直接在 JVM 内完成“图片解码 → 特征提取 → 相似度计算”全流程工程结构简单得多。DJL 是 AWS 开源的 Java 深度学习库底层引擎支持 PyTorch、TensorFlow、MXNet接口风格贴近 JPMS 和 Java 习惯没有把 C 的细节全部暴露给调用方。本文按一条完整落地路径来写先准备一个适合做相似度计算的特征提取模型再用 DJL 加载并输出特征向量最后实现余弦相似度与批量检索。下面每一步都附带可复现的代码和参数说明方便你直接接到自己的项目里。2. 图片相似度的前置一步用 TorchScript 导出 ResNet50 特征模型2.1 为什么相似度计算不能直接比较像素矩阵计算两张图片相似度最朴素的做法是逐像素相减但这对尺寸、光照、物体位置极其敏感。一个物体平移几个像素像素差就会剧烈变化但视觉上仍然是同一张图。更稳定的做法是先通过卷积神经网络把图片编码成一个固定长度的向量让向量之间的距离近似表示图片之间的语义距离。这个过程通常被称为“特征提取”或“embedding 提取”。在深度学习中分类模型的最后一层输出是各个类别的概率而倒数第二层往往保存着高层次的抽象特征。比如 ResNet50 在 ImageNet 预训练后图像的全局池化特征向量是 2048 维这 2048 个浮点数可以用于一万张规模下的近似去重和以图搜图。需要提前说明DJL 自身不带可视化标注好的特征专用模型社区里常见做法是加载 PyTorch 的预训练模型去掉末尾的全连接分类层再导出为 TorchScript供 DJL 在 Java 端加载。2.2 选 ResNet50 作为特征提取基座的三个理由ResNet50 不是最强的特征提取模型但在 Java 后端做推理服务时它有三个不可替代的优势第一识别精度与推理速度的平衡点很好CPU 上跑一张 224x224 的图大约几十毫秒到一两百毫秒瓶颈很容易被多线程消化第二输出维度只有 2048和动辄 5120 维的 ViT 模型相比后续做余弦相似度和建立索引的内存压力低得多第三预训练权重随处可得导出 TorchScript 的过程非常成熟遇到问题的资料也比新模型多。EfficientNet 和 ConvNeXt 也可以替换但替换时要注意输入尺寸和归一化参数必须跟着预训练配置走。下面以 ResNet50 为例说明如何导出一个只输出特征向量的 TorchScript 模型。2.3 用 PyTorch 导出特征向量的最小可运行代码在安装好 PyTorch 的环境里执行下面的 Python 脚本import torch import torchvision.models as models # 使用 ImageNet 预训练权重若网络受限可先离线下载权重文件 model models.resnet50(weightsmodels.ResNet50_Weights.DEFAULT) # 去掉最后一层全连接分类层输出直接变成 2048 维特征向量 model.fc torch.nn.Identity() # 切换到 eval 模式让 BatchNorm 使用训练阶段缓存的均值和方差 model.eval() # 构造一个与预期输入一致的假数据input shape 是 (batch, channel, height, width) dummy_input torch.randn(1, 3, 224, 224) # 用 torch.jit.trace 生成 TorchScript 图 traced_model torch.jit.trace(model, dummy_input) # 保存为 .pt 文件后续由 DJL 直接加载 traced_model.save(resnet50_feature.pt)这段代码的重点在三处model.fc torch.nn.Identity()使模型不再输出 1000 类概率model.eval()控制 BatchNorm 和 Dropout 的行为避免推理时统计量漂移torch.jit.trace用假数据走一次前向把动态流图固定成静态图。如果你希望同时压缩体积可以在保存前调用traced_model torch.jit.optimize_for_inference(traced_model)但这会把 BatchNorm 折叠进卷积层修改后仍可直接使用。导出的resnet50_feature.pt输入为[1, 3, 224, 224]输出为[1, 2048]这个形状后面在 Java 端定义 Translator 时一定要保持一致。如果你想让模型直接接受更大尺寸可以改dummy_input的宽高但 ResNet 的预训练规模通常意味着 224 或 288 效果更好超过原训练尺寸不会提升特征质量反而增加延迟。3. 在 Java 中用 DJL 加载模型并完成图片向量化3.1 Maven 依赖与 Java 环境准备DJL 的核心依赖只有一个api具体引擎实现需要单独引入。这里以 PyTorch 引擎为例在pom.xml中加入properties !-- 用占位符避免版本过期实际构建时替换为 Maven Central 上的稳定版 -- djl.version${djl.version}/djl.version /properties dependencies dependency groupIdai.djl/groupId artifactIdapi/artifactId version${djl.version}/version /dependency dependency groupIdai.djl.pytorch/groupId artifactIdpytorch-engine/artifactId version${djl.version}/version /dependency dependency groupIdai.djl.pytorch/groupId artifactIdpytorch-native-auto/artifactId version${djl.version}/version scoperuntime/scope /dependency /dependenciesJava 环境建议使用 JDK 11 以上命令行里先确认java -version与环境变量配置没有问题。pytorch-native-auto会根据当前操作系统自动拉取对应的原生库所以不需要手动安装 PyTorch如果你的生产服务器是离线环境可以换成pytorch-native-cpu并手动上传对应的 so/dll 文件。3.2 自定义 Translator把图片变成模型需要的张量模型输入不能直接是BufferedImageDJL 通过Translator负责把输入图片解码、缩放、归一化再转换为NDList。这里是完整实现import ai.djl.modality.cv.Image; import ai.djl.modality.cv.transform.CenterCrop; import ai.djl.modality.cv.transform.Normalize; import ai.djl.modality.cv.transform.Resize; import ai.djl.modality.cv.transform.ToTensor; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDList; import ai.djl.ndarray.NDManager; import ai.djl.translate.Translator; import ai.djl.translate.TranslatorContext; public class FeatureTranslator implements TranslatorImage, float[] { private static final float IMAGE_SIZE 224f; private static final Resize RESIZE new Resize(IMAGE_SIZE, IMAGE_SIZE); private static final CenterCrop CROP new CenterCrop(IMAGE_SIZE, IMAGE_SIZE); private static final ToTensor TO_TENSOR new ToTensor(); private static final Normalize NORMALIZE new Normalize( new float[]{0.485f, 0.456f, 0.406f}, new float[]{0.229f, 0.224f, 0.225f}); Override public NDList processInput(TranslatorContext ctx, Image input) { NDArray array input.toNDArray(ctx.getNDManager(), Image.Flag.COLOR); array RESIZE.transform(array); array CROP.transform(array); array TO_TENSOR.transform(array); array NORMALIZE.transform(array); // 增加 batch 维模型要求输入形状为 (1, 3, 224, 224) array array.expandDims(0); return new NDList(array); } Override public float[] processOutput(TranslatorContext ctx, NDList list) { // 模型输出是 (1, 2048)squeeze 后变成一维数组 try (NDArray array list.singletonOrThrow().squeeze()) { return array.toFloatArray(); } } }需要注意Resize与CenterCrop的顺序两个都设置为 224x224真实执行时先 Resize 再 CenterCrop 基本等价于直接缩放但保留 Crop 可以让图片从中心截取避免角落信息干扰。ToTensor会把像素值从0~255缩放到0.0~1.0因此Normalize用的是 PyTorch 标准 ImageNet 均值方差。如果模型是从其他框架导出的这组参数必须改成对应的预训练配置否则特征向量会严重漂移。3.3 封装特征提取服务有了 Translator加载模型并对外提供服务只需要很少的代码import ai.djl.Application; import ai.djl.inference.Predictor; import ai.djl.modality.cv.Image; import ai.djl.modality.cv.ImageFactory; import ai.djl.repository.zoo.Criteria; import ai.djl.repository.zoo.ZooModel; import java.nio.file.Path; public class FeatureExtractor implements AutoCloseable { private final ZooModelImage, float[] model; public FeatureExtractor() { CriteriaImage, float[] criteria Criteria.builder() .setTypes(Image.class, float[].class) .optModelPath(Path.of(models/resnet50_feature.pt)) .optTranslator(new FeatureTranslator()) .optEngine(PyTorch) .build(); try { model criteria.loadModel(); } catch (Exception e) { throw new IllegalStateException(加载 DJL 模型失败, e); } } public float[] extract(Image image) throws Exception { try (PredictorImage, float[] predictor model.newPredictor()) { return predictor.predict(image); } } Override public void close() { if (model ! null) { model.close(); } } public static void main(String[] args) throws Exception { Image image ImageFactory.getInstance().fromFile(Path.of(demo.jpg)); float[] vector new FeatureExtractor().extract(image); System.out.println(特征维度: vector.length); } }main方法验证整个链路ImageFactory负责读取图片文件并解码为 RGB 数据newPredictor()会为当前调用分配一个推理上下文用完关闭。这里有一个高频坑Predictor不是线程安全的但创建成本不高所以最稳妥的用法是每次推理都新建或者用 ThreadLocal 缓存。实际项目中我一般会在 Spring 的 Service 里持有ZooModel每次请求单独创建Predictor避免多线程共享导致的引擎崩溃。4. 余弦相似度计算与索引优化的 Java 实现4.1 余弦相似度、欧氏距离和向量夹角的关系拿到两张图的 2048 维特征向量后相似度计算就变成纯数学问题。最常用的是余弦相似度它只关心向量的方向不关心模长而欧氏距离关心绝对距离。举例来说同一张图分别缩放成 100x100 和 500x500 后特征向量方向接近但模长可能差异很大余弦相似度会给出更符合直觉的结果。计算公式是cos(a, b) (a·b) / (|a| × |b|)。用 Java 实现时建议先对特征向量做 L2 归一化这样后续只需要算点积能省掉每次计算夹角的开根号开销。下面代码同时保留了原始计算和归一化后的点积方便观察差异public class VectorUtils { public static double cosineSimilarity(float[] a, float[] b) { double dot 0; double normA 0; double normB 0; for (int i 0; i a.length; i) { dot a[i] * b[i]; normA a[i] * a[i]; normB b[i] * b[i]; } return dot / (Math.sqrt(normA) * Math.sqrt(normB) 1e-8); } public static void normalize(float[] vector) { double norm 0; for (float v : vector) { norm v * v; } norm Math.sqrt(norm); for (int i 0; i vector.length; i) { vector[i] (float) (vector[i] / norm); } } public static double dotProduct(float[] a, float[] b) { double dot 0; for (int i 0; i a.length; i) { dot a[i] * b[i]; } return dot; } }归一化后余弦相似度与点积等价因此可在批量比对时先对库中所有向量归一化再循环遍历点积。1e-8是防止零向量产生 NaN 的安全兜底几乎没有计算成本。阈值方面在 ImageNet 特征上做一些版权图去重场景时0.92 以上基本可认定为同一张图如果只是相近内容的聚类可以放到 0.80 以下这个没有通用标准必须用你自己的业务图集标定。4.2 不同距离度量方法的适用场景下表整理了四种常见度量在图片相似度任务中的表现后续选型时可以直接参考度量方法计算复杂度对光照变化对几何变化典型场景余弦相似度低较稳定较稳定以图搜图、重复图片判定欧氏距离低不稳定不稳定人脸特征打印比对曼哈顿距离低不稳定中等紧凑编码后的粗筛汉明距离极低不稳定中等二进制哈希特征汉明距离一般配合二值化向量使用比如把每维特征量化为 0/1 后做快速粗筛。这样做能大幅降低内存和计算量但精度损失也明显适合十万级以上的候选集召回阶段。下一节的哈希索引思路就是基于这个思想。4.3 一万张图片的批量检索极简局部敏感哈希如果库图只有几百张直接两两点积没有问题。但一万张图需要一亿次点积每次点积遍历 2048 维整体耗时在单线程下会到达不可接受的程度。常见做法是先建立倒排桶只有落在同一桶的向量才做精确对比。这里给出一个基于随机投影的简易实现import java.util.*; public class SimpleHashIndex { private final int bits; private final float[][] randomVectors; private final MapInteger, Listfloat[] buckets new HashMap(); public SimpleHashIndex(int bits, int dim) { this.bits bits; this.randomVectors new float[bits][dim]; Random random new Random(42); for (int i 0; i bits; i) { for (int j 0; j dim; j) { randomVectors[i][j] random.nextGaussian(); } } } public int hash(float[] vector) { int code 0; for (int b 0; b bits; b) { float dot 0; for (int i 0; i vector.length; i) { dot vector[i] * randomVectors[b][i]; } code (code 1) | (dot 0 ? 1 : 0); } return code; } public void add(float[] vector) { buckets.computeIfAbsent(hash(vector), k - new ArrayList()).add(vector); } public Listfloat[] search(float[] query, int topK) { Listfloat[] candidates buckets.getOrDefault(hash(query), List.of()); candidates.sort(Comparator.comparingDouble(v - -VectorUtils.cosineSimilarity(v, query))); return candidates.subList(0, Math.min(topK, candidates.size())); } }bits控制哈希桶的数量bits20时理论上能产生一百万个桶随机投影向量用固定种子 42 初始化保证每次构建出来的索引一致。这个实现的召回率依赖向量在空间中的分布实际项目中可以对比不同bits下最终返回的 topK 是否稳定。不要追求完全精确的最近邻重点是先用哈希把十万张图过滤到几百个候选再对候选做精排最终效果与全量扫描一致响应时间却能缩短两个数量级。5. DJL 图片相似度计算上线前验证流程与高频坑5.1 验证相似度排序是否正确的三步法先准备三张图A 是原始图B 是对 A 做轻微压缩的副本C 是与 A 完全不相关的另一张图。期望输出是 sim(A,B) 远高于 sim(A,C)。用以下脚本确认public class Validation { public static void main(String[] args) throws Exception { FeatureExtractor extractor new FeatureExtractor(); float[] a extractor.extract(ImageFactory.getInstance().fromFile(Path.of(a.jpg))); float[] b extractor.extract(ImageFactory.getInstance().fromFile(Path.of(b.jpg))); float[] c extractor.extract(ImageFactory.getInstance().fromFile(Path.of(c.jpg))); System.out.printf(A-B 相似度: %.4f%n, VectorUtils.cosineSimilarity(a, b)); System.out.printf(A-C 相似度: %.4f%n, VectorUtils.cosineSimilarity(a, c)); } }如果 A-B 的相似度低于 0.9先检查两张图是否真的经过了常见的缩放或压缩如果压缩比例过高ResNet50 的输出也会下降。在进入批量索引之前至少要有一组“相似为正例、不相似为负例”的标注集用召回率和误报率决定阈值。5.2 三个影响结果的高频坑第一个坑是NDArray内存泄漏。DJL 的每个NDArray都受NDManager管理如果你在循环里调extract但忘记关闭Predictor最后会触发 Native 内存溢出。处理方式很简单所有涉及NDArray.toFloatArray()的输出都放在 try-with-resources 里转成 Java 数组后立刻释放。第二个坑是图像通道顺序。Image.Flag.COLOR会把图片解码为 RGB 三通道但某些客户上传的是 RGBA 或 CMYK 的 JPEG直接交给ImageFactory可能出现颜色偏移。保险做法在上传阶段统一转成 RGB或者调image.getWrappedImage()做一次校验。第三个坑是模型没有处于 eval 状态。从 PyTorch 导出时如果漏写model.eval()TorchScript 中 BatchNorm 会在推理时继续更新统计量导致同样两张图在不同时间提取的特征不一致。检查方法是在 Java 端对同一张图片连续提取两次特征计算两次输出的余弦相似度理论上应该接近 1.0。5.3 用 ONNX Runtime 替代 PyTorch 的迁移边界如果生产服务器希望减少 PyTorch 原生库的体积可以把导出步骤从 TorchScript 换成 ONNX并让 DJL 搭配onnxruntime-engine加载.onnx文件。需要注意的是ONNX 导出的特征向量可能与 PyTorch 存在1e-6量级差异阈值需要重新标定。上线时先用pytorch-engine跑历史图集生成一批基线向量再与 ONNX 输出做逐维比对确认最大误差不超过1e-4再切换引擎。把这条验证逻辑写进 CI图片相似度计算的升级过程就不会被底层引擎的微小差异悄悄改变。本文还有配套的精品资源点击获取
返回列表