
上周帮团队排查一个线上推荐服务的性能问题Python 侧跑模型Java 侧通过 Redis 拿推理结果一次请求要跨两个进程模型分发、监控、版本回滚都要维护两套系统。我当时就在想如果模型能直接加载进 JVM让 Java 进程自己完成从数据预处理到模型输出的全部工作这一堆事至少能省一半。后来我真把 PyTorch 的 Java 引擎接进了生产环境发现第一个卡住我的不是模型结构也不是网络调参而是最基础的张量基本操作——数据类型、shape、切片、reshape、广播。这些在 Python 里一行语法糖就能解决的事换到 Java 侧全都要重新适应。这个系列是面向 Java 工程师的 PyTorch 实战课程按硕士研一强度设计不要求你有 Python 经验。今天讲的是第一章第三讲张量基本操作我尽量把背后的数据逻辑也讲清楚而不只是给几个 API 让你抄完就完事。1. 从一次模型部署扯开Java工程师为什么需要掌握张量1.1 被 Python 和 Java 割裂的深度学习工作流过去几年大部分团队的深度学习工作流是“训练和推理分离”的算法工程师在 Python 里训模型后端工程师在 Java 里把它包成一个 HTTP 服务。听起来没毛病但一旦请求量上去问题就来了——模型服务要保高可用Python 侧的 GIL、内存回收、依赖冲突、环境迁移每一样都是运营噩梦。更难受的是数据前置处理比如归一化、维度变换、batch 组装如果放在 Python 侧做Java 服务就得把原始数据传过去一个来回就是几十毫秒的延迟如果放在 Java 侧做你就必须能够在 Java 代码里玩转张量。这就是我对当前行业阶段的理解训练框架已经稳定模型文件格式已经标准化真正开始成为重头戏的是 AI 基础设施如何跟业务系统融合。过去我们讲 AI Infra 1.0 是搭 GPU 集群2.0 是训练框架和调度标准化到 3.0 这个阶段核心已经变成了“让 AI 能力成为常规后端组件”Java 这种生产主力语言必须能直接调度模型、操作张量而不是永远隔一层 Python 进程。所以 PyTorch On Java 不是一个玩具项目它解决的是从训练到部署之间被语言割裂的那一段距离。1.2 这门课适合谁你会学到什么如果你是 Java 后端工程师、中间件开发者或者正在做推荐、搜索、风控这类算法接入工作的平台工程师那么 PyTorch 的 Java 能力值得你花时间掌握。课程的前置要求不高懂 Java 基本语法、理解 JVM 内存模型、会 Maven 依赖管理Linux 基本命令知道一些就够了。完全不会 Python 也没关系你只需要知道“模型是在 Python 里训练出来的”我们的重点是把训练好的模型和它的数据流接进 Java 应用。这一讲的产出目标是四个第一能说清楚张量到底是什么而不是把它当成一个绕口的名词第二能在 Java 代码里创建、查看、裁剪一个张量第三理解 reshape、permute 这些操作到底动没动数据第四用这些操作跑通一条最朴素的模型推理链路。前两个是“会用”后两个是“用明白”后面课程讲自动微分和 torch.nn 的时候地基就从这里来。2. 环境准备把 PyTorch 引擎请进 JVM2.1 官方 Java API 与 DJL两条路的取舍现在想在 Java 里用 PyTorch主流有两条路我建议你先把它们分清不要混着抄代码。第一条是 PyTorch 官方的 Java APIMaven 坐标是org.pytorch:pytorch_java_all。它是通过 JNI 直接封装 libtorch 的加载 TorchScript 模型非常直接适合做纯推理。但它的硬伤也很明显高层张量运算方法很少文档稀疏你想做x.mul(2).add(y)这种组合操作相当别扭更像是一个“能把模型跑起来”的最小绑定。第二条是 AWS 开源的 DJLDeep Java Library核心思想是给 Java 工程师一个更友好的前端底层引擎可以换成 PyTorch、TensorFlow 等。DJL 提供的NDArray封装了大部分深度学习常见的张量运算API 更符合工程直觉同时又保留了对底层 PyTorch 引擎的直接访问。我这一讲的教学思路就是拿 DJL 作为载体因为这是普通 Java 工程师最难卡壳、最容易上手的路线。严格来说你学到的是“PyTorch 引擎 Java 前端”模型格式还是 TorchScript底层 JNI 库还是 libtorch并没有跑偏。下表可以帮你快速选型维度官方 Java APIDJL PyTorch Engine定位薄封装适合纯推理全功能前端适合完整工程张量操作偏少丰富接近 Python 体验模型加载TorchScript 模型TorchScript 模型团队熟悉度需要自己封装API 贴近 Java 习惯适合场景快速验证生产系统、复杂预处理2.2 Maven 依赖与首次加载以我常用的版本组合为例先用官方 API 体验一下最原始的模型加载是什么样子dependency groupIdorg.pytorch/groupId artifactIdpytorch_java_all/artifactId version1.13.1/version /dependency然后写一段最朴素的代码Module model Module.load(resnet18.pt); Tensor input Tensor.fromBlob(new float[1 * 3 * 224 * 224], new long[]{1, 3, 224, 224}); IValue result model.forward(IValue.from(input)); Tensor output result.toTensor(); System.out.println(output.toString());注意一个关键点Tensor.fromBlob的第一个参数是一维 Java 数组第二个参数是 shape。PyTorch Java API 里Module.load加载的是 TorchScript 格式不是直接塞一个 Python 训练出的.pt文件就行。如果你手里只有 Python 训练好的权重需要先在 Python 侧做一次torch.jit.script或torch.jit.trace导出。如果你用 DJL依赖会长这样dependency groupIdai.djl/groupId artifactIdapi/artifactId version0.30.0/version /dependency dependency groupIdai.djl.pytorch/groupId artifactIdpytorch-engine/artifactId version0.30.0/version /dependency dependency groupIdai.djl.pytorch/groupId artifactIdpytorch-native-auto/artifactId version1.13.1/version /dependency这里我不建议你原样照抄版本号。依赖版本之间必须匹配尤其 native 包要和 engine 版本保持一致否则 JNI 加载时报错会让你怀疑人生。我在实际项目里见过太多次UnsatisfiedLinkError最后定位下来全是 native 库和 Java 层版本错位。2.3 原生库版本的三个经典坑第一个坑是 JDK 版本。新版 DJL 推荐 Java 11 以上虽然 Java 8 也能跑但某些内存访问和模块化行为在 Java 8 下表现不一样。如果你在生产机器上还在用 JDK 8建议先把升级计划排上日程否则后面做 JNI 调试会很痛苦。第二个坑是 MKL 和 OpenBLAS 冲突。PyTorch 原生库依赖底层 BLAS 优化库如果你的机器上还有别的原生数值库加载时可能出现符号冲突。我踩过最诡异的一次是模型第一次加载成功第二次并发访问时直接 JVM 崩溃。排查到最后发现是 GPU 版本和 CPU 版本的 native 包同时出现在 classpath 里。解决方式很简单确认你的 pom 里只保留一个 native classifier不要cpu和cu113混着引。第三个坑是本地库路径。生产环境一般用 Docker 镜像native 库会解压到一个临时目录如果有只读文件系统权限限制native 库可能解压失败。稳妥的做法是设置DJL_CACHE_DIR环境变量让 DJL 把 native 库解压到确定的位置镜像构建时提前把该目录放行。3. 张量的数据哲学shape、dtype、layout 与第一个 Tensor3.1 张量不是“多维数组”那层皮很多 Java 工程师第一次接触张量会把它当成int[][]的直接升级版。这个直觉对一半但你如果真按照“数组的数组”去理解后面 permute、transpose、broadcast 会把你绕晕。我更愿意用快递柜来打比方。一组长方体的快递柜有 4 排 3 列你取件时会说“第 2 排第 3 格”这就是 shape。但快递柜的服务员在后台记录包裹位置时用的是“从入口数第几个格子”可能是连续编号的这就是内存里真正的存储顺序。张量和快递柜的相似点就在这里逻辑上是多维的物理上却是一段一维连续空间。shape 告诉你有几个维度、每个维度多大stride 告诉你沿着某个维度移动一格实际要跨过多少个元素。这两个信息组合起来计算机才知道“第 2 排第 3 格”到底对应物理上的第几个元素。很多操作比如reshape、permute本质上只是在改 shape 和 stride 这套“元数据”数据本身没有复制。理解这一点之后你就不会再担心“一次 reshape 是不是把整个数组又复制了一份”。除了 shape还有两个属性必须养成看的好习惯dtype 和 device。dtype 是元素类型深度学习默认用 float32因为训练时梯度和精度需要平衡推理时很多人用 float16 或 int8本质是用精度换吞吐。device 则决定了数据在哪个设备上计算CPU 还是 GPU。跨设备搬运数据是一个重操作在代码里要尽量少做。3.2 在 Java 里创建你的第一批张量用 DJL 创建张量非常 Java 化try (NDManager manager NDManager.newBaseManager()) { NDArray a manager.create( new float[]{1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f}) .reshape(2, 3); System.out.println(a.toDebugString()); System.out.println(a.getShape()); System.out.println(a.getDataType()); }这里有个非常重要的组件NDManager。你可以把它理解成所有张量的“主管”负责创建、回收张量背后的原生内存。后面我会专门讲内存管理这里你只需要记住第一条铁律张量原生内存不归 JVM GC 管NDManager必须在 try-with-resources 里使用否则迟早吃内存泄漏的亏。除了从数组创建你还会高频使用这些工厂方法NDArray zeros manager.zeros(new Shape(2, 3)); NDArray ones manager.ones(new Shape(2, 3)); NDArray random manager.randomNormal(new Shape(2, 3));zeros和ones很好理解一个全零一个全一。randomNormal生成的是服从标准正态分布的随机数模型权重初始化和数据增强里经常用到。这里我想提醒你不要小看manager.randomNormal它底层走的是 PyTorch 的随机数生成器不是Math.random()那种默认实现。随机种子、生成算法、多线程并发下的随机行为在深度学习里都是敏感点。4. 张量基本操作索引、切片、reshape 与内存视图4.1 切片和索引不是所有下标访问都值得写循环Java 里操作数组最常见的姿势就是 for 循环。到了张量世界你要尽快改掉这个习惯因为循环会把性能和可读性一起拖垮。NDArray的索引操作长这样NDArray x manager.create(new float[]{1, 2, 3, 4, 5, 6}).reshape(2, 3); // 取第一行 NDArray row x.get(new NDIndex(0)); // 取第二列的所有行 NDArray col x.get(new NDIndex(:, 1)); // 修改某个位置的值 x.set(new NDIndex(1, 2), 100f);NDIndex里写:表示这一维全部取类似 Python 的冒号。如果要取第 0 到第 1 列写作new NDIndex(:, 0:2)。这套表达能力覆盖了“连续块”“跳跃”“布尔选择”等场景。别再写嵌套 for 循环去抠元素了一方面慢另一方面写出来的代码没法跟算法工程师沟通——他们习惯用切片思维描述数据流。有一个细节值得注意切片操作有些返回的是原始数据的视图有些则会拷贝。视图的好处是几乎零开销坏处是你改了切片原始数据也可能变。在 DJL 里当你需要一份独立数据时要显式调用duplicate()。我做数据处理时凡是要进推理模型的数组几乎都会用duplicate()复制一份避免后续的中间操作不小心污染了原始输入。4.2 reshape、permute、squeeze元数据在变数据不一定动在图像领域最常见的数据格式变化是 HWC 和 CHW。HWC 表示一个图像按高、宽、通道排布比如 224x224 的 RGB 三通道图就是(224, 224, 3)CHW 则是通道在前(3, 224, 224)。很多训练好的模型输入要求 CHW而图片解码库输出的通常是 HWC。用张量操作来转换NDArray img manager.create(new float[224 * 224 * 3]).reshape(224, 224, 3); NDArray chw img.permute(2, 0, 1); // 维度顺序从 (H, W, C) 变成 (C, H, W) NDArray batch chw.expandDims(0); // 加一个 batch 维度 - (1, C, H, W)permute是维度重排它不复制数据只是改了每个维度的 stride 信息。这里就体现出前面讲“快递柜”的价值了逻辑上我们把维度换了个顺序物理上的数据还是原来那段连续内存只是计算机读取时跳步方式不同了。reshape的逻辑更贴近“重新划分格子”总元素数量不变的前提下把一维数据重新组织成任意合理形状。需要注意如果原始数据在内存里不是连续的比如你刚permute完又要reshape底层就可能发生一次隐式复制否则无法保证数据逻辑顺序。expandDims则是给张量插入一个大小为 1 的维度你不需要为了增加一个 batch 轴而新建一个大数组。4.3 组合操作中的常见误解我见过不少新手写出类似“先 permute 再 reshape”的代码然后发现结果完全不对。原因是 permute 之后数据的逻辑顺序和存储顺序已经不一致此时 reshape 会先尝试把数据“拉平”这个拉平顺序和你脑子里想象的“按新维度顺序拉平”并不是一回事。遇到这种情况Python 社区的习惯是先调用contiguous()让数据在内存里按当前逻辑顺序重排再 reshape。DJL 里没有直接叫contiguous()的方法但相似的思路是如果你不确定数据连续性就不要连续链式调用 reshape 和 permute。要么每次调用后打印toDebugString()看结果要么直接用duplicate()拿一份连续数据。还有一个容易忽略的地方是argMax。分类模型输出的通常是一个概率分布向量你要取最大值的下标作为预测类别NDArray logits predictor.predict(batch); long[] maxIdxArray logits.argMax(1, false).toLongArray();注意第二个参数keepDim我习惯传false这样返回的 shape 会去掉被归约的那个维度。细节上差一个维度下游代码就很容易出现数组越界或者维度不匹配。5. 广播机制与矩阵乘法把向量化思维带到 Java5.1 广播到底是什么一个对齐位置的规则广播英文叫 broadcast是张量运算里最反直觉、但也最省心的机制。它让你可以对两个形状不完全一样的张量做运算只要维度能“对齐”就行。规则只有两条从最后一个维度往前比对如果两个维度相等或者其中一个为 1就能对齐否则直接报错。看这个例子NDArray a manager.create(new float[]{1, 2, 3}).reshape(3, 1); NDArray b manager.create(new float[]{4, 5, 6, 7}).reshape(1, 4); NDArray c a.add(b); System.out.println(c.getShape()); // [3, 4]? 不对是 [3, 4]? 我看一下a的形状是(3, 1)b的形状是(1, 4)。从右往左对齐第一个维1 和 4 比1 可以扩展到 4第二个维3 和 1 比1 可以扩展到 3。最后结果就是(3, 4)。这个过程不会真的把 3 和 4 复制成 12 个元素再相加底层是零拷贝的“虚拟扩展”所以性能没有想象中那么差。广播的好处是让代码极其简洁。如果让你用两个 for 循环算这 12 个元素你不仅要写循环还要处理边界很容易出错。而在深度学习里归一化操作、偏置相加、mask 填充全部依赖广播。比如模型输出是(N, C)你要给每个样本加上一个均值向量(C,)直接output.add(meanVector)即可。5.2 矩阵乘法与 batch 计算MatMul 不只是“双重循环”的代替品Java 工程师做数据计算天然会用 for 循环实现矩阵乘法那个三重循环我写过太多次了。在张量世界里矩阵乘法有专门的算子matMul它底层调用的是经过高度优化的 BLAS 库速度远超手写循环。写法非常直接NDArray input manager.randomNormal(new Shape(1024, 128)); NDArray weight manager.randomNormal(new Shape(128, 64)); NDArray output input.matMul(weight); // shape (1024, 64)这里你看到的就是一个线性层Linear Layer的核心计算输入特征矩阵乘以权重矩阵。Transformer 里的 attention score、全连接网络、embedding 查询归根到底都是矩阵乘法的不同变形。当数据带到 batch 维度后三维张量的矩阵乘法用batchMatMulNDArray q manager.randomNormal(new Shape(4, 8, 16)); NDArray k manager.randomNormal(new Shape(4, 16, 8)); NDArray scores q.batchMatMul(k); // shape (4, 8, 8)这里4是 batch 大小每一批内部做一次(8,16)x(16,8)的矩阵乘。没有了手动循环每个 batch 的乘法和后面做 softmax 的步骤可以天然衔接。我之前一个跑 transformer 的项目attention 分数上游用 Python 算、下游 Java 算时就是被batchMatMul这个 API 救回来的。6. 内存管理与设备切换JVM 之外的资源同样要关6.1 NDManager 的层级结构与自动回收前面反复提到NDManager现在必须把它讲透。如果你只把张量当作 Java 对象写代码时可能很爽但生产环境一定挂给内存溢出。PyTorch 引擎分配的张量内存大部分在 JVM 堆外这部分内存 JVM 的垃圾回收器是管不到的。只要你在 Java 侧丢掉了NDArray引用那个对象本身可以被 GC 回收但它背后的原生内存可能没有被及时释放。DJL 解决这个问题的方式是NDManager的层级结构你创建的每一个数组都挂在一个NDManager下面。父 manager 关闭时它会递归关闭所有子 manager并释放所有挂在这个子图上的张量内存。所以正确的写法是try (NDManager parent NDManager.newBaseManager()) { NDArray a parent.create(new float[]{1, 2, 3}); try (NDManager child parent.newSubManager()) { NDArray b child.create(new float[]{4, 5, 6}); // 用 a 和 b 做计算 } // 到这里 b 已经被自动释放a 还可以继续用 }我自己的习惯是凡是循环里创建张量必须确保它们在子 manager 中创建并随循环关闭凡是推理请求进来我会把每次请求的数据放在独立的子 manager 里请求结束就关闭避免请求量一大把 JVM 堆外内存吃掉。6.2 Device 切换与 CPU/GPU 版本坑张量在哪个设备上计算取决于它的Device。默认情况下创建出的张量都在 CPU 上。如果服务器有 GPU你想把数据搬到 GPUNDArray cpuArray manager.create(new float[]{1, 2, 3}); NDArray gpuArray cpuArray.toDevice(Device.gpu(), true);如果当前环境没有可用的 NVIDIA GPU调用Device.gpu()很可能直接抛异常所以代码里要提前判断。一个实用技巧是启动时检测一次设备可用性然后根据结果决定计算策略不要在每次请求时都做设备探测。这里还要强调依赖坑你引入的 native 包如果是 CPU 版那么即使有 GPU 也用不上如果是 CUDA 版机器上又必须存在匹配版本的 NVIDIA 驱动和 CUDA 运行时。否则加载 JNI 库时经常报找不到符号或库文件。团队协作时最好把 native 依赖写进一个专门的 BOM 或 profiles 里区分 CPU 环境与 GPU 环境避免哪天有人在没有 GPU 的机器上强行跑 CUDA 版本。我见过最隐蔽的一次问题是 Docker 镜像里没有安装对应驱动依赖模型一加载就libcuda.so not found排查了整整一个下午。7. 综合小练习用 Tensor 手写一个推理前处理管线7.1 场景图像分类前的数据准备理论讲再多不落地都是虚的。最后我用一个图像分类的推理前处理例子把前面讲的操作全部串起来。假设你在 Python 侧训练了一个 ResNet-18并且已经用 TorchScript 导出成resnet18.pt。模型期望的输入格式是(1, 3, 224, 224)像素值要归一化到[-1, 1]附近。这段逻辑在 Python 里通常写成x x / 255.0 x (x - mean) / std x x.transpose(2, 0, 1) x x.unsqueeze(0)在 Java 侧我们用 DJL 逐行复刻这个流程。先加载模型CriteriaNDArray, NDArray criteria Criteria.builder() .optEngine(PyTorch) .optModelPath(Paths.get(resnet18.pt)) .optTypes(NDArray.class, NDArray.class) .build(); ZooModelNDArray, NDArray model criteria.loadModel(); PredictorNDArray, NDArray predictor model.newPredictor();Criteria是 DJL 里声明“我要加载什么模型、用什么引擎、输入输出长什么样”的配置对象。这里指定NDArray作为输入输出类型意思是原始数据和推理结果都用张量进行交换。7.2 串联Java 里完成归一化、维度变换和批量推理伪代码级别的核心流程try (NDManager manager NDManager.newBaseManager()) { // 假设已经通过图片解码得到 HWC 顺序的像素float 值在 [0, 255] NDArray pixels manager.create(hwcFloatArray).reshape(224, 224, 3); // 归一化 NDArray normalized pixels.div(255.0f) .sub(meanVector) .div(stdVector); // 维度调整HWC - CHW - 增加 batch 维 NDArray chw normalized.permute(2, 0, 1); NDArray batch chw.expandDims(0); // 推理 NDArray logits predictor.predict(batch); long[] classIds logits.argMax(1, false).toLongArray(); System.out.println(预测类别下标: classIds[0]); }这段代码基本复刻了 Python 侧的全部处理逻辑而且每一步都看得清楚。pixels.div(255.0f)把像素从 0 到 255 缩放到 0 到 1.sub(meanVector)减去均值.div(stdVector)除以标准差permute(2, 0, 1)把 HWC 转换成 CHWexpandDims(0)加上 batch 维。最后argMax取出最大概率对应的类别。写到这里我想多说一句真正放到生产里你不会把解码逻辑写成manager.create一长串而是会抽成一个ImagePreprocessor类专门负责从图片字节到模型输入的转换。不要小看这个“小练习”当年我在团队里做的第一个 AI Infra 服务核心链路其实就是这么朴素读图、缩放、归一化、转格式、推理、取最大值。分布式训练、多卡调度那些听起来高级的东西都是后话能不能把最基础的数据战线理顺决定了一个 AI 系统在业务手里好不好用。如果你刚接触这套技术栈我的建议很直接先别急着学自动微分和模型训练把 shape、dtype、reshape、broadcast 这四个概念练到不需要查文档的程度。我这边的经验是很多项目里真正的耗时点不在模型精度而在数据在各个环节之间的形状对不对得上。张量基本操作就是这个系列的地基我当初在这里磨了两天换来后面读模型源码、写推理服务时完全不卡壳。下一讲我们会进入梯度自动计算到时候这批张量就要参与真正的反向传播了那时候你再回头看这一讲就会明白为什么每一行操作都值得认真对待。