最近陆续有朋友问我同一个问题:你是做Java的,怎么突然研究起PyTorch了?其实这事儿对很多后端团队来说挺现实——模型训练阶段大家都在Python生态里玩得飞起,可一到部署,面对Java为主的微服务体系,就有点尴尬。要么单独起一个Python服务,要么硬着头皮把推理逻辑用Java重写一遍。两个方案都痛:多一个服务就多一个运维节点,重写模型等于自找麻烦。PyTorch On Java 的出现,就是为了填上这个坑。这篇文章是 PyTorch On Java 系列第一章第二节,重点讲讲张量操作。我会从一个Java工程师的视角,把张量怎么创建、怎么操作、怎么和底层内存交互讲清楚,顺便把这一路踩过的坑都抖出来。
1. 为什么Java项目里要跑PyTorch:先看清这套技术栈的真实价值
1.1 Java后端接深度学习的两种常见姿势,以及它们的痛点
先说个最常见的场景。你在一家做电商或金融系统的公司,后端全线是Spring Boot,服务拆了几十个,都是Java写的老实人。某天业务方提需求:要把用户评论做情感分析,或者对商品图片做分类。算法团队早就把模型训好了,代码基于PyTorch,效果也验证了。现在问题来了:怎么把这个模型塞进你的Java服务里?
传统做法无非两条路。第一条,单独部署一个Python推理服务,Java这边走HTTP或者RPC调用。好处是算法团队能直接维护,坏处也明显——你的架构图上多了一个用不同语言写的“外挂”,部署、监控、权限、依赖管理全部要重来一遍,而且每次模型更新,两边还得对接口。第二条路,把模型导出成某种中间格式,比如ONNX,然后Java端用别的推理引擎加载。这条路的坑在于,训练时的算子和推理引擎的算子未必完全对齐,碰到一些特殊操作就会报错,排查起来非常消耗耐心。
两条路我都走过,说实话没有一条是省心的。所以当我第一次看到官方提供了完整的Java绑定,可以直接在Java进程里加载PyTorch模型、执行张量计算的时候,第一反应是:这才是Java后端该有的接入姿势。
1.2 PyTorch On Java到底包了什么
很多人以为PyTorch Java支持还是个半成品,其实官方从1.8开始就一直在推进这块。现在的Java绑定主要分两层:
- 底层的LibTorch:这是PyTorch的C++核心库,Java通过JNI方式调用它。所有张量操作、模型推理、自动求导,底层都是同一套C++实现,不是模拟实现,也不是“近似实现”。
- 上层的Java API:官方提供了
org.pytorch:pytorch-java这样的Maven依赖,把C++操作封装成Java对象和方法,让Java开发者可以用面向对象的方式执行张量运算、加载模型、做推理。
这套架构最大的价值在于:训练时用的PyTorch模型,推理时还是同一个引擎,算子对齐问题天然不存在。从Python端torch.save出来的模型文件,Java端直接用Module.load加载,语义完全一致。对做AI Infra的团队来说,这意味着你可以把“算法训练”和“服务部署”的边界划得更干净。
1.3 这门课的定位和适合人群
我说句实在话,这系列课程不面向纯新手。你要是连Java的基本语法、Maven依赖、线程模型都还没弄明白,建议先补基础。但如果你是这些情况,这套内容正好补上你的知识缺口:
- Java后端工程师,团队里开始引入AI能力,你不想永远靠调别人的HTTP接口过活,希望自己能直接操作张量、加载模型、做推理;
- 算法工程师,模型训练经验丰富,但对Java服务端一窍不通,想知道算法产出物怎么真正落到业务系统里;
- 做AI Infra的研发,需要对模型部署、推理加速、内存管理有深入理解,而不只是跑通demo。
这节的核心是张量操作。张量是PyTorch一切计算的基础,你后面写的每一段推理代码,本质上都是张量的创建、变换、运算和读取。把这块吃透了,再往后的模型加载和推理,你会觉得顺理成章。
2. 环境搭建是这个样子的:版本对齐、依赖引入和第一行张量代码
2.1 版本对齐是第一道坎:Java、PyTorch、LibTorch要匹配
先说结论:版本不匹配是初学者最常见的报错来源,没有之一。你在Maven里引一个pytorch-java依赖,它背后对应着某个版本的LibTorch,如果你机器上装了另一个版本的Python PyTorch,很容易在加载模型时出现IllegalStateException或UnsatisfiedLinkError,提示某些符号找不到。
我个人的建议是:训练环境的PyTorch版本、Java依赖声明的版本、LibTorch native库的版本,三者保持一致。比如你训练模型时用的是PyTorch 2.1.0,那么Java端就把pytorch-java和对应的pytorch-java-native都锁在2.1.0。
以Maven为例,核心依赖加这些:
<dependency> <groupId>org.pytorch</groupId> <artifactId>pytorch-java</artifactId> <version>2.1.0</version> </dependency> <dependency> <groupId>org.pytorch</groupId> <artifactId>pytorch-java-native</artifactId> <version>2.1.0</version> <classifier>linux-x86_64</classifier> </dependency>注意这个<classifier>,它指定了操作系统的类型。如果是CPU环境,用linux-x86_64就行;如果是GPU环境,可能得选带gpu标识的版本,具体看官方仓库发布的构件。Mac用户对应的是macosx-x86_64,Apple Silicon需要看是否单独出了arm版本。
2.2 Maven依赖引入:区分CPU版和GPU版的细节
这里想在依赖层面多说一句。PyTorch Java的native库体积不小,默认拉下来可能是几百MB。如果你只是本地跑跑测试,不建议一上来就上GPU版,因为GPU版还需要额外匹配CUDA版本,而CUDA的版本又得看显卡驱动。CPU版装好了,把代码逻辑跑通,再考虑上GPU,这样排错范围会小很多。
如果你用的构建工具是Gradle,对应写法也不复杂:
implementation 'org.pytorch:pytorch-java:2.1.0' implementation 'org.pytorch:pytorch-java-native:2.1.0:linux-x86_64'依赖拉取完成后,可以写个最简代码验证环境:
import org.pytorch.Tensor; import org.pytorch.torch.Tensor as TorchTensor; // 这是示意,实际import路径看API版本 // 简化验证 Tensor t = Tensor.fromBlob(new float[] {1f, 2f, 3f}, new long[] {3}); System.out.println(t.toString());这段代码不一定能直接复现,因为官方Java API在不同版本里的包名和类名有过调整,但核心目的就是验证JNI加载、native库和基础张量操作是否正常。如果你连一个3元素的一维张量都打印不出来,那说明环境层有问题,先回头检查版本和classifier。
2.3 环境踩坑:JNI库加载失败和native库路径问题
我见过最典型的报错长这样:
java.lang.UnsatisfiedLinkError: no jni_pytorch in java.library.path这个错误翻译成人话就是:Java启动的时候,没找到PyTorch的JNI动态链接库。虽然Maven依赖里声明了native库,但如果agent或者IDE的启动配置里没有把native库目录加进去,就有可能出现这个情况。
解决办法有两个方向。一是在启动参数里显式指定:
-Djava.library.path=/你的本地路径/libtorch/lib二是确认依赖完整后,让Maven帮你把native库解压到本地,再观察java.library.path是否包含了对应目录。如果你用的是spring-boot-maven-plugin,还要注意打包时会不会把native库忽略掉。
注意:如果你的服务最终要容器化部署,Docker镜像里也得装好对应的C++运行库(比如
libgomp),否则镜像里一切正常,一跑JVM就报JNI错误。而且容器镜像的CPU指令集要和编译LibTorch时的指令集兼容,低端CPU跑高版本编译的lib,偶尔会碰到非法指令错误。
3. 张量核心操作实战:从创建、索引到广播机制的完整认知
3.1 Tensor创建的五种常用方式,以及它们的语义
Java里写张量,最直观的对象是Tensor,它对应PyTorch Python端的torch.Tensor。创建方式我总结下来主要有五种,各有各的适用场景。
第一种,fromBlob,从Java原生数组直接创建。这是最常用的方式,因为你从数据库、网络请求、文件里读到的数据,大概率就是float[]、double[]或者int[],把它们原样包装成张量,最省事:
float[] data = new float[] {1, 2, 3, 4}; Tensor t = Tensor.fromBlob(data, new long[] {2, 2});这里第二个参数是shape,我上面写的{2, 2}就表示2行2列。要特别注意,fromBlob在多数实现里是引用,不是复制。也就是说,改动Java数组,张量内部的数据也会变。这个特性用好了能省内存,用不好就是数据错乱的坑。
第二种,zeros和ones,创建全0或全1的张量,常用于初始化或mask操作。
第三种,rand和randn,创建均匀分布或标准正态分布的随机张量。做测试或者模拟数据时很好使。
第四种,arange,创建一个数值范围张量,类似JavaIntStream.range的语义。
第五种,从已有的Tensor用empty或new Tensor扩展,这个更底层一些,实际项目里相对少用。
3.2 张量索引与切片:和Numpy趋同,但Java API的表现形式不同
和Python的Numpy相比,Java API在索引切片上的设计稍微“官方”一点:不直接用中括号语法,而是封装了Index和Slice工具类。比如想取某一行,Python代码一行搞定:
t[:, 1]Java这边就得这样写:
Tensor sliced = t.select(1, 1);或者用Slice表示范围,配合get方法:
Tensor sub = t.indexSelect(1, Tensor.fromBlob(new long[] {0, 2}, new long[]{2}));第一眼看上去确实比Python繁琐,但这其实更符合Java语言一贯的显式风格。用习惯了以后,写代码时脑子里要清晰地存着一句话:索引操作不创建新数据,它只是创建了一个视图(view)。这跟Numpy的切片一样,可以做到内存零拷贝。
这个“视图”概念很重要。你操作视图,原张量也会变,这和Java基础类型数组的直觉是反的。很多从Python刚转过来的同事,在视图上做修改,结果发现源数据被动改了,一脸懵。
3.3 张量数学运算:逐元素操作、矩阵乘法、归约
数学运算是张量操作里占比最大的一块。官方Java API把运算函数设计成了Tensor的实例方法,比如:
Tensor a = Tensor.fromBlob(new float[] {1, 2, 3}, new long[] {3}); Tensor b = Tensor.fromBlob(new float[] {4, 5, 6}, new long[] {3}); Tensor sum = a.add(b); // 逐元素相加 Tensor mul = a.mul(b); // 逐元素相乘 Tensor matmul = a.reshape(new long[]{1, 3}).mm(b.reshape(new long[]{3, 1})); // 矩阵乘这些操作底层都对应LibTorch里的ATen算子,性能和Python训练时是同一套实现,不会因为换了个语言就变慢。真正影响性能的,是你有没有在Java层反复转换数据、有没有频繁创建不必要的中间张量。
归约操作也很常用,比如求sum、mean、max,对应Java API里有sum()、mean()、max()等方法。如果你想指定在哪条维度上归约,需要传入dim参数,具体参数顺序建议查一下对应版本的API文档,因为不同版本有调整。
3.4 广播机制:Java里的广播和Python同样强大,但要小心隐式行为
广播(broadcasting)是个老朋友了。PyTorch的广播规则简单来说就是:从尾部维度开始对齐,维度相等或者其中一个为1,就能对齐;不满足就报错。Java API完全继承了这套规则,所以下面这个操作是合法的:
Tensor base = Tensor.ones(new long[] {2, 3}); Tensor bias = Tensor.ones(new long[] {3}); Tensor result = base.add(bias); // 结果还是 2x3这里bias是一维的3元素,base是2行3列,尾部维度都是3,合法广播。
但我要提醒的是:Java API里广播操作容易让人误以为做了内存扩展。其实广播是“逻辑上的扩展”,底层内存并没有真正复制。如果你处理的是超大张量,不要害怕广播,它比手动repeat快得多。
3.5 类型转换与shape变换
深度学习模型里经常要求输入是特定dtype、特定shape。Java API对dtype的支持比较直白,主要有FLOAT、DOUBLE、INT64、BOOL等。转换方法类似:
Tensor floatTensor = longTensor.toType(org.pytorch.DType.FLOAT);shape变换则是reshape和view两个方法容易搞混。简单说,reshape更灵活,如果底层数据不连续,它可能帮你复制数据;view只适合能共享内存的场景,底层不连续时会报错。刚上手建议优先用reshape。
4. 从张量运算到模型推理:模型加载、输入预处理和输出解析的完整链路
4.1 模型加载:Module.load和它背后的状态管理
张量操作练得差不多了,下一步必然是怎么把它用到实际模型上。PyTorch Java加载模型很简单,核心就一个类Module:
Module model = Module.load("/models/resnet18.pt");这个方法内部会初始化LibTorch的JIT解释器,然后加载TorchScript格式的模型。注意这里有个关键前提:模型必须是TorchScript格式,不能用Python的torch.save(model.state_dict())直接保存的那种。简单来说,你得在Python环境里先用torch.jit.script或者torch.jit.trace把模型导出成.pt或.torchscript文件,Java端才能加载。
导出文件这个环节,很多Java同学不熟悉。我用Python端示意一下trace的写法:
import torch import torchvision.models as models model = models.resnet18(pretrained=True) model.eval() example = torch.rand(1, 3, 224, 224) traced = torch.jit.trace(model, example) traced.save("resnet18.pt")这个导出过程就生成了Java端需要的文件。
4.2 输入预处理:从图片到张量的三种路径对比
模型加载完成后,推理前紧接着就是预处理。比如图像分类模型,输入是[1, 3, 224, 224]的浮点张量,数值范围通常是0到1,还要按通道做个归一化。Java端怎么从一张图片变成这种张量?我试过几种方案,对比如下:
| 方案 | 操作方式 | 优点 | 缺点 |
|---|---|---|---|
| 使用Java图像IO + 手工像素遍历 | BufferedImage.getRGB()逐像素读取,手动填充多维数组 | 依赖少,控制力强 | 代码量大,性能一般 |
| 使用OpenCV Java绑定 | Imgproc做缩放和通道转换,再转成Mat,再读取数据 | 图像处理能力强,性能好 | 依赖较重,需要单独引入OpenCV |
| 直接用Tensor直接加载PNG等 | 部分PyTorch版本支持从文件直接解码 | 代码最简 | 支持格式有限,不好定制预处理 |
我自己的项目里用的第二套方案。先说思路:用OpenCV读图、缩放、转RGB、归一化,然后把Mat的数据用fromBlob包装成张量。整个过程核心代码大概这么写:
// 以OpenCV为例,示意预处理流程 Mat src = Imgcodecs.imread(imagePath); Mat resized = new Mat(); Imgproc.resize(src, resized, new Size(224, 224)); Imgproc.cvtColor(resized, resized, Imgproc.COLOR_BGR2RGB); resized.convertTo(resized, CvType.CV_32FC3, 1.0 / 255.0); float[] pixels = new float[3 * 224 * 224]; resized.get(0, 0, pixels); // 注意:此时pixels是HWC排列,还需要转成CHW,再做RGB通道的均值方差归一化 Tensor input = Tensor.fromBlob(pixels, new long[] {1, 224, 224, 3}); Tensor inputChw = input.permute(new long[] {0, 3, 1, 2}); Tensor normalized = inputChw.sub(0.485).div(0.229); // 简化写法,实际每个通道单独归一化这段只是骨架,真正项目里你要把mean和std换成ImageNet的标准值,并且对R、G、B三个通道分开处理。
顺带提一句,很多Java教程里会直接跳过预处理这一步,拿一个整形数组直接塞给张量,然后推理结果乱七八糟。这不怪模型,是你的输入分布和训练时不一致。预处理和训练时的数据预处理必须保持一致。
4.3 推理执行和输出解析:从Tensor到业务结果
模型推理走一个forward方法就行:
IValue output = model.forward(IValue.from(inputTensor)); Tensor outputTensor = output.toTensor();IValue是PyTorch Java里包装输入输出的通用对象,相当于Python端的IValue概念。拿到输出张量后,分类任务通常要读取置信度和类别索引:
long[] shape = outputTensor.shape(); // shape是 [1, num_classes],所以调用 Tensor probabilities = outputTensor.softmax(1); long maxIdx = probabilities.argmax(1).item().toLong(); float maxVal = probabilities.get(0, maxIdx).item().toFloat();这里留意一下:argmax(1)表示在第1个维度(类别维度)上取最大值索引,返回的还是一个Tensor,想拿到Java基础类型,要用.item()方法解包。
item()是张量操作里容易被忽略但是特别实用的方法,它把单元素张量转换成Java标量。如果你拿到一个单元素张量后忘了调item(),后面做if (maxIdx == 1)这类比较时会很痛苦,因为Tensor不是long。
4.4 一个完整的线性模型推理Demo
光说不练假把式。这里给一个极简可跑的完整例子,模型就用Python端torch生成的线性层来模拟:
public class DemoInference { public static void main(String[] args) { Module model = Module.load("linear.pt"); try (Tensor input = Tensor.fromBlob(new float[] {1.0f, 2.0f, 3.0f}, new long[] {1, 3})) { IValue out = model.forward(IValue.from(input)); Tensor outTensor = out.toTensor(); System.out.println("Model output shape: " + outTensor.shape().length); for (int i = 0; i < outTensor.shape()[1]; i++) { System.out.println(outTensor.get(0, i).item().toFloat()); } } } }这段代码里用到了Java的 try-with-resources 语法,Tensor实现了Closeable,这点和Python不一样,后面我会展开讲为什么必须关。
5. JNI内存模型与性能优化:Java端必须盯紧的底层细节
5.1 张量生命周期:为什么Java版的Tensor要手动close
如果你从Python转过来,对张量内存的第一直觉是“用完就扔,回收交给GC”。在Java API里,这个直觉会害了你。Tensor对象虽然看起来是个Java对象,但它持有的原生内存是由LibTorch管理的,不在JVM堆内。JVM的GC根本感知不到这块内存,除非Tensor对象被GC回收时调用了close或finalize,否则这块内存会被一直占用。
大型模型部署场景里,我见过最夸张的情况是:每次推理都创建一个输入Tensor,用完不关,跑了一个小时,JVM堆内存才几百兆,机器整体内存却快满了。这就是直接内存泄漏。
所以Java PyTorch的实战铁律是:凡是自己创建的Tensor,用完就关;凡是从模型输出拿到的Tensor,用完也关。写法上有两种方式,一种手动close,更稳妥的是用 try-with-resources:
try (Tensor input = createInputTensor()) { IValue out = model.forward(IValue.from(input)); try (Tensor output = out.toTensor()) { // 处理输出 } }如果是在循环里做批量推理,尤其要小心不要在循环里new一堆Tensor然后不管。有人觉得“每次才几MB,无所谓”,但压测时几百QPS一上,内存曲线会教你做人。
注意:
Module和IValue也有对应的释放机制。一般来说Module在整个进程生命周期里只加载一次,不要反复load,每次load都相当于重新初始化一次LibTorch的模型解释器。
5.2 fromBlob是引用还是拷贝:理解了它,你就理解了性能关键
前面提到过fromBlob的引用语义,这里再多说几句。在多数情况下,fromBlob不会立即复制Java数组的数据,而是让Tensor直接指向这个数组的内存区域。这样做的好处是零拷贝,坏处是:你在Java数组后续写入数据,等于直接修改张量。
这既是优势也是坑。例如从OpenCV拿到的像素字节数组,如果直接用fromBlob包装成Tensor,省掉了一次数组复制,在大尺寸图片上能省不少时间。但如果你是为了异步推理,把数组留着后续复用,那一定要搞清楚数组里的数据被Tensor引用着,改了就是脏数据。
如果你明确需要独立的内存副本,用Tensor.clone()或者先通过Tensor.fromBlob再复制到新张量。这块逻辑在优化时非常关键,但刚入门时不用过度设计,先搞清楚默认行为,遇到bug时才有排查方向。
5.3 推理性能优化:预热、批处理与资源复用
再往深走一点,性能优化有几个方向是通用的。
第一是预热。LibTorch首次推理可能包含一些初始化开销,比如算子的库加载、内存分配。生产环境上的常见做法是:应用启动后,先拿一个假数据跑一次推理,把该初始化的都初始化完,再对外提供流量。
第二是批处理。如果你一次要处理100张图片,与其循环100次调用forward,不如拼成一个[100, 3, 224, 224]的大张量,一次推理。GPU环境下批处理提升明显,CPU环境也能减少函数调用和调度开销。Java端的做法就是创建一个更大的Tensor,把多张图片的数据沿着batch维度拼接起来。
第三是线程模型。LibTorch Java绑定能否多线程并发推理,取决于你的模型和做法。通常可以用多个线程各自持有独立的Module实例,或者共享一个Module但控制并发上限。踩过坑的结论是:不要无脑把同一个Tensor丢给多个线程一起修改,原生Tensor不是线程安全的。
5.4 怎么定位原生内存占用:JVM之外的那块内存
排查Java进程内存问题时,光看jmap不够,因为原生内存不在堆上。我的经验是:
- 用
jcmd <pid> VM.native_memory查看JVM自身的内存分类统计,能看到部分JNI分配(如果开启了NMT); - 用系统级命令(比如容器里的
cat /proc/<pid>/status或top)确认RSS增长趋势; - 结合代码审计:看Tensor有没有都走try-with-resources,看
fromBlob是否造成意外的长期引用; - 在压测环境里做前后对比,每次改动控制一个变量。
这一步本身就是AI Infra团队的核心工作之一:调优不是靠猜,而是靠统计和验证。把这些内存和性能基本功练好,后面做大规模部署时才不会翻车。
6. 张量调试三板斧:打印、形状检查和本地复现
6.1 打印张量内容和形状:Java下的Numpy体验替代方案
调试张量操作,第一需求永远是“看看到底长什么样”。Java API里打印Tensor,不同版本的默认toString()输出详细度不一样,有的版本默认只打印形状不打印内容。这时候可以用一个笨办法但非常有效:先把张量转回Java数组再打印。
private static void printTensor(String tag, Tensor t) { System.out.println(tag + " shape: " + Arrays.toString(t.shape())); // 仅针对小张量,大张量不要这样干 FloatBuffer buffer = FloatBuffer.allocate((int) t.numel()); t.copyTo(buffer); float[] arr = new float[buffer.remaining()]; buffer.get(arr); System.out.println(Arrays.toString(arr)); }numel()表示张量总元素个数,copyTo把原生数据读到Java的FloatBuffer里。小张量调试利器,但大张量别打全量,一个百MB的Tensor你打出来,终端直接卡死。
6.2 形状不匹配是头号bug来源:三个可复用的自查步骤
张量相关的bug,一大半都是shape问题。我的自查顺序是:
第一步,确认预期shape。比如模型要求的输入是[N, C, H, W],那你的预处理就得按这个顺序和维度填。
第二步,确认实际shape。用上面那个打印方法,把每一步的张量shape打出来,看到底是哪一步开始偏的。
第三步,确认dtype。很多算子对类型敏感,Java里的float[]转换成Tensor后是FLOAT,但如果模型参数是DOUBLE,两者做运算有时会隐式提升、有时报错,具体情况版本不同。
这三步走完,大概80%的“张量操作报错”都能定位。
6.3 跨语言复现问题:Java看不明白,回Python验证一遍
遇到比较诡异的问题,我经常用一招:同样的张量操作,回Python端跑一遍,对比数值。PyTorch Java和Python共用底层C++引擎,算子数值一致性是有保证的,如果两边结果不一致,先怀疑Java端的数据有没有被错误修改,比如fromBlob引用了同一个底层数组,但多次操作产生了相互影响。
另外,模型输入预处理不一致导致的推理结果偏差,也得靠跨语言对比来揪出来。比如Python端用的是(x - mean) / std,Java端写成了x / std - mean,这两种写法在数值上完全不是一回事。
7. 我的几个实践心得,以及下一步可以聊什么
回到最开始的问题:Java项目里到底要不要上PyTorch On Java?我现在的态度是,只要你的团队以Java为核心,模型部署又要长期迭代,这套技术栈就值得投入。它不需要你额外维护Python服务,还和训练生态天然同源,本质上是把AI能力“内化”到现有技术体系里。
但别把学习曲线想得太短。Java版本的API成熟度、文档丰富度、社区案例都远不如Python生态,遇到问题能搜到的资料少很多,很多时候得自己读官方的C++源码、看API仓库里的测试用例来逆推用法。这恰恰是它筛选人的地方——能吃下这块硬骨头的人,在AI Infra这种岗位上会很值钱。
我的实际体会:刚开始接触Java张量时,最大的阻碍不是张量本身,而是思维切换。Python张量是“默认零成本操作”,Java张量必须时刻想着生命周期、引用还是副本、底层内存还住着谁。一旦你把这些底层机制理顺了,Java张量操作基本不会再出幺蛾子,而且你会比始终停留在Python高级封装的人更能理解深度学习框架的本质。
最后分享一个小技巧:写Java张量相关代码时,给自己定个规矩,所有临时Tensor都写在try-with-resources里,所有shape转换都加断言。
try (Tensor input = Tensor.fromBlob(rawData, new long[] {1, 3, 224, 224})) { assert input.shape()[0] == 1; assert input.shape()[1] == 3; // 继续后面的逻辑 }这一步,能让你的代码在变成正式服务后,少掉一半的内存追踪噩梦。下一步可以考虑聊一聊模型的TorchScript导出细节、GPU推理配置,或者完整的Java推理服务封装,看大家更关心哪一块。