Deep Java Library (DJL) 使用

将 PyTorch 模型转换为 TorchScript

DJL 加载 PyTorch 模型,要求模型必须是 TorchScript 格式(.pt 或 .pth 文件)。这是因为 TorchScript 是一种序列化的、可独立于 Python 运行的模型表示。

1、用分类模型来举例:

python的转换代码:

用python训练好的模型,直接转换,然后保存。

import torch
import torchvision.models as models

# 1. 加载一个预训练模型 (以 ResNet18 为例)
model = models.resnet18(pretrained=True)
model.eval()

# 2. 创建一个示例输入,用于追踪模型
dummy_input = torch.randn(1, 3, 224, 224)

# 3. 使用 torch.jit.trace 转换为 TorchScript
traced_model = torch.jit.trace(model, dummy_input)

# 4. 保存为 .pt 文件
traced_model.save("traced_resnet18.pt")

pom.xml (Maven) 文件中,添加 DJL 的核心 API 和 PyTorch 后端引擎依

<dependencies>
    <!-- DJL 核心 API -->
    <dependency>
        <groupId>ai.djl</groupId>
        <artifactId>api</artifactId>
        <version>0.28.0</version>
    </dependency>
    <!-- PyTorch 后端引擎 -->
    <dependency>
        <groupId>ai.djl.pytorch</groupId>
        <artifactId>pytorch-engine</artifactId>
        <version>0.28.0</version>
        <scope>runtime</scope>
    </dependency>
    <!-- 简单日志实现,用于查看 DJL 运行日志 -->
    <dependency>
        <groupId>org.slf4j</groupId>
        <artifactId>slf4j-simple</artifactId>
        <version>1.7.36</version>
    </dependency>
</dependencies>

Java 推理代码实现

import ai.djl.Model;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
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.modality.cv.translator.ImageClassificationTranslator;
import ai.djl.translate.TranslateException;
import ai.djl.application.ImageClassification;

import java.io.IOException;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.List;

public class DJLPyTorchDemo {

    public static void main(String[] args) throws IOException, TranslateException {
        // 1. 加载模型
        //    指定模型路径和模型名称(不包含 .pt 后缀)
        Path modelDir = Paths.get("build/pytorch_models/resnet18");
        String modelName = "resnet18"; // 对应 traced_resnet18.pt 文件

        try (Model model = Model.newInstance(modelName)) {
            // 加载模型文件
            model.load(modelDir, modelName);
            System.out.println("模型加载成功!");

            // 2. 准备图像预处理流程 (Translator)
            //    这里的预处理需要和 PyTorch 训练时保持一致
            ImageClassificationTranslator translator = ImageClassificationTranslator.builder()
                    .addTransform(new Resize(256))
                    .addTransform(new CenterCrop(224))
                    .addTransform(new ToTensor())
                    .addTransform(new Normalize(
                            new float[]{0.485f, 0.456f, 0.406f},  // mean
                            new float[]{0.229f, 0.224f, 0.225f}   // std
                    ))
                    // 可选:加载类别标签文件 (synset.txt)
                    // .optSynset(Paths.get("build/pytorch_models/resnet18/synset.txt"))
                    .build();

            // 3. 创建 Predictor 对象进行推理
            try (Predictor<Image, ImageClassification> predictor = model.newPredictor(translator)) {
                // 4. 准备输入图像 (此处为示例,你需要替换为实际图片路径)
                Image image = ImageFactory.getInstance().fromUrl(
                        "https://djl-ai.s3.amazonaws.com/resources/images/kitten.jpg"
                );

                // 5. 执行推理
                ImageClassification result = predictor.predict(image);
                List<ImageClassification.Class> topClasses = result.topK(3);

                // 6. 输出结果
                System.out.println("推理结果:");
                for (ImageClassification.Class clazz : topClasses) {
                    System.out.printf("类别: %s, 概率: %.5f%n", clazz.getClassName(), clazz.getProbability());
                }
            }
        }
    }
}

2、线性回归举例

python转换:

import torch
import torch.nn as nn

# 1. 定义并训练一个简单的线性模型
model = nn.Linear(1, 1)
# ... (此处省略训练过程) ...

# 2. 切换到评估模式
model.eval()

# 3. 创建一个示例输入,用于追踪模型
#    形状必须与模型训练时的输入一致,这里是 [batch_size, features]
example_input = torch.tensor([[5.0]], dtype=torch.float32)

# 4. 使用 torch.jit.trace 导出
traced_model = torch.jit.trace(model, example_input)

# 5. 保存为 .pt 文件 (也可以是 .zip)
traced_model.save("my_regression_model.pt")

java推理实现:

<dependencies>
    <dependency>
        <groupId>ai.djl</groupId>
        <artifactId>api</artifactId>
        <version>0.28.0</version>
    </dependency>
    <dependency>
        <groupId>ai.djl.pytorch</groupId>
        <artifactId>pytorch-engine</artifactId>
        <version>0.28.0</version>
        <scope>runtime</scope>
    </dependency>
    <!-- 日志实现 -->
    <dependency>
        <groupId>org.slf4j</groupId>
        <artifactId>slf4j-simple</artifactId>
        <version>1.7.36</version>
    </dependency>
</dependencies>
import ai.djl.Model;
import ai.djl.inference.Predictor;
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;

// 1. 定义输入为 float[],输出为 float 的转换器
public class RegressionTranslator implements Translator<float[], Float> {

    @Override
    public NDList processInput(TranslatorContext ctx, float[] input) {
        // 将 Java 的 float[] 转换为 DJL 的 NDArray (张量)
        // 形状需要是 [1, input.length],1 代表 batch size
        NDArray array = ctx.getNDManager().create(input).reshape(1, input.length);
        return new NDList(array);
    }

    @Override
    public Float processOutput(TranslatorContext ctx, NDList list) {
        // 从模型输出的 NDArray 中提取第一个结果
        // 假设模型输出形状为 [1, 1]
        NDArray result = list.singletonOrThrow();
        return result.getFloat(); // 返回 float 值
    }
}

3、神经网络模型

Python端导出:

import torch
import torch.nn as nn

# 1. 定义一个自定义神经网络(例如:10维输入 -> 5维输出)
class MyNet(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(10, 64)
        self.fc2 = nn.Linear(64, 32)
        self.fc3 = nn.Linear(32, 5)

    def forward(self, x):
        x = torch.relu(self.fc1(x))
        x = torch.relu(self.fc2(x))
        return self.fc3(x)  # 回归任务无激活函数

model = MyNet()
model.eval()

# 2. 构造示例输入(batch_size=1, features=10)
dummy_input = torch.randn(1, 10)

# 3. 追踪并保存
traced_model = torch.jit.trace(model, dummy_input)
traced_model.save("my_custom_net.pt")  # 关键输出文件

Java端加载:

<dependency>
    <groupId>ai.djl</groupId>
    <artifactId>api</artifactId>
    <version>0.28.0</version>
</dependency>
<dependency>
    <groupId>ai.djl.pytorch</groupId>
    <artifactId>pytorch-engine</artifactId>
    <version>0.28.0</version>
    <scope>runtime</scope>
</dependency>
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;

// 输入:float[](特征数组),输出:float[](预测结果数组)
public class ArrayToArrayTranslator implements Translator<float[], float[]> {

    @Override
    public NDList processInput(TranslatorContext ctx, float[] input) {
        // 将 Java 数组转为 DJL 张量,形状为 [1, 10](1个样本,10个特征)
        NDArray array = ctx.getNDManager().create(input).reshape(1, input.length);
        return new NDList(array);
    }

    @Override
    public float[] processOutput(TranslatorContext ctx, NDList list) {
        // 获取模型输出的第一个张量(形状为 [1, 5])
        NDArray result = list.singletonOrThrow();
        
        // 将张量数据复制到 Java 的 float 数组中
        float[] output = result.toFloatArray();
        return output; // 返回包含5个预测值的数组
    }
}
import ai.djl.Model;
import ai.djl.inference.Predictor;
import ai.djl.translate.TranslateException;
import java.io.IOException;
import java.nio.file.Paths;
import java.util.Arrays;

public class LoadCustomNeuralNetwork {

    public static void main(String[] args) throws IOException, TranslateException {
        // 1. 加载模型(指向 my_custom_net.pt 所在的文件夹)
        String modelDir = "path/to/your/model/folder/";
        String modelName = "my_custom_net"; 

        try (Model model = Model.newInstance(modelName)) {
            model.load(Paths.get(modelDir), modelName);
            System.out.println("自定义神经网络加载成功!");

            // 2. 绑定自定义转换器
            ArrayToArrayTranslator translator = new ArrayToArrayTranslator();
            
            // 3. 创建 Predictor
            try (Predictor<float[], float[]> predictor = model.newPredictor(translator)) {
                
                // 4. 模拟输入数据:10个特征值
                float[] inputFeatures = new float[]{1.0f, 2.1f, 3.2f, 4.3f, 5.4f, 
                                                    6.5f, 7.6f, 8.7f, 9.8f, 10.9f};

                // 5. 执行推理
                float[] predictions = predictor.predict(inputFeatures);

                // 6. 打印输出(5个预测值)
                System.out.println("预测结果: " + Arrays.toString(predictions));
            }
        }
    }
}

软件和硬件的配置:

centerOS系统,java11以上,maven3.6.3,需要安装cuda,需要根据显卡的型号,cuDNN 是 NVIDIA 的深度学习加速库,能进一步提升 GPU 推理性能。安装同样需要从 NVIDIA 官网下载对应 CUDA 版本的 cuDNN 库。

提供几档配置参考:

场景一:入门学习 / 轻量级推理 (预算敏感)

适合新手学习、运行小模型(如TinyLlama)、进行简单的图像分类(如CIFAR-10)或课程演示。

  • GPUNVIDIA RTX 3060 (12GB显存) 或二手 GTX 1080 (8GB显存)。12GB显存是入门体验的“甜点”配置。

  • CPU≥ 8核,如 Intel i5 或 AMD Ryzen 5 系列。

  • 内存32GB 起步。

  • 存储240GB - 512GB SSD

场景二:进阶开发 / 均衡实用 (主流选择)

适合大多数开发者和研究人员,能流畅运行13B-20B参数的模型(如 Llama 2-13B, ChatGLM4),进行中等规模GAN训练或 Stable Diffusion 图像生成。

  • GPUNVIDIA RTX 4090 (24GB显存) 是此阶段的“黄金组合”。24GB显存能避免多数模型的显存溢出错误。

  • CPU≥ 16核,如 Intel i7/i9 或 AMD Ryzen 7/9 系列。

  • 内存64GB 或更高。

  • 存储512GB - 1TB NVMe SSD

场景三:大型模型 / 多卡并行 (专业/企业级)

适用于需要训练或部署70B以上大模型(如 Llama 2-70B)、进行大规模分布式训练的场景。

  • GPU多卡并行,如 2×RTX 4090 (24GB)8×A100 (40GB/80GB)或 8×H100 (80GB)

  • CPU服务器级CPU,如 双路 Intel Xeon 或 AMD EPYC,核心数 ≥32核

  • 内存≥256GB,通常需要 DDR4/DDR5 ECC 内存以保证稳定性。

  • 存储大容量NVMe SSD(如 2TB - 4TB)。

补充说明:对于70B参数的大模型,使用半精度(FP16/BF16)加载时,模型本身大小约为 70B * 2 bytes = 140GB。因此,需要多张显卡的显存总和(如8×A100的320GB)才能装下。

1、在深度学习中,“B”代表“Billion”(十亿)
所以,70B 指的是模型有 700亿 个参数(即权重和偏置的总数量)。

70B=70×1,000,000,000=70,000,000,000(700亿)

2、在计算机中,数据的存储单位是字节(Byte)。

1 Byte = 8 bit(比特)。

深度学习模型在计算时,为了平衡精度和显存占用,通常不采用高精度的 FP32(32位浮点数,占4字节),而是采用半精度的 FP16 或 BF16(16位浮点数,占2字节)

因此,每个参数在显存中占据 2 个字节 的空间

3. 第三步:计算总字节数

将参数总数乘以每个参数占用的字节数,得到模型纯权重所占用的总字节数

70,000,000,000(个参数)×2(字节/个)=140,000,000,000(字节)

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值