Ohhnews

分类导航

$ cd ..
foojay原文

如何创建Spring Boot欺诈评分服务

#spring boot#欺诈检测#机器学习#deep netts#java

大多数希望将机器学习模型投入生产的 Java 团队,最终都会另起一个 Python 服务,然后通过 HTTP 调用它。这种方式确实能跑通,但作为一名 Java 开发者,你为此付出了额外的代价:多一套运行时、多一套部署流水线、每次预测都多一次网络跳转,还多出一条团队边界,最终让“重新训练模型”变成了别人工单里的事。

Deep Netts 消除了这种割裂:它是一个纯 Java 的深度学习库。模型用 Java 训练,序列化到文件,再像普通 Bean 一样加载回 Spring Boot 应用。预测变成进程内的方法调用,耗时以微秒计,不需要额外部署、加固或监控任何组件。

本教程从 Deep Netts 的信用卡欺诈检测示例出发,并把它带到了示例刻意没涉及的地方——一个可以部署到真实流量面前的 HTTP 服务。你将完成以下工作:

  1. 在公开的 Kaggle 交易数据集上训练一个前馈网络。
  2. 将模型连同它隐式依赖的缩放参数一并导出。
  3. 将模型包装进 Spring Boot 应用,并确保线程安全、有经过深思熟虑的判定阈值、健康检查和回归测试。

你可能需要花与网络架构同样多的时间来处理运维问题——CI 中的依赖解析、模型版本管理、许可限制——因为在实践中,这些才是决定这个服务能否撑过第一个月的东西。

预算大约需要两个小时。

第 1 步——让原始示例先跑起来

1.1 将 Deep Netts 安装到本地 Maven 仓库

Deep Netts Pro 不在 Maven Central 上,但本场景需要用到它。它是免费的(除非你的收入超过 100 万美元)。

deepnetts.com/download-latest 下载 Deep Netts Pro ZIP 包,解压后,在解压目录中运行导入脚本:

./importToLocalMaven.sh      # 在 Windows 上使用 importToLocalMaven.bat

确认它已安装成功:

ls ~/.m2/repository/com/deepnetts/
# 预期结果:deepnetts-core-pro/  deepnetts-license/

1.2 克隆并运行示例

git clone https://github.com/deepnetts/CreditCardFraudDetection.git
cd CreditCardFraudDetection
unzip creditcard.zip          # 完整的 284k 行数据集
mvn clean package
mvn exec:java

项目的 pom.xml 声明了以下依赖:

依赖作用
com.deepnetts:deepnetts-core-pro:3.2.0神经网络引擎
com.deepnetts:deepnetts-license:1.0许可 JAR——运行时同样需要
tech.tablesaw:tablesaw-core / tablesaw-jsplot用于数据探索的数据框和图表
javax.visrec:visrec-ri:1.0.3JSR-381 参考实现

在写任何 Spring 代码之前,先把这一步跑通。如果许可 JAR 无法解析,后续所有工作都会卡住,你还会浪费一个小时去排查是不是 Spring 的问题。

1.3 理解数据内容

仓库中附带两个文件:

  • creditcard.csv(在 zip 文件内)——约 284,807 笔交易,其中约 492 笔是欺诈。这大约是 0.17%
  • creditcard-balanced.csv——经过欠采样、类别比例大致均衡的版本。

列包括:TimeV1V28AmountClass(0 = 正常,1 = 欺诈)。也就是 30 个可用输入和 1 个二分类输出。

类别不平衡是整个问题的核心。一个对每笔交易都预测“不是欺诈”的模型,准确率高达 99.83%,但毫无价值。在这个问题上永远不要只报告准确率。 你需要关注精确率、召回率和混淆矩阵——这些将在第 3 步中介绍。

在均衡文件上训练是标准的起步做法,但要理解其中的取舍:欠采样会丢弃约 99% 的正常样本,最终模型的原始输出是针对一个并不存在的 50/50 世界进行校准的。你的判定阈值需要额外调整(见第 4 步)。


第 2 步——让数据划分变得可复现

示例在运行时通过随机洗牌划分数据。对于一个服务来说,你需要将缩放器与模型一起持久化,而缩放器的参数必须来自严格的训练数据行——不能是每次重新随机划分的结果。因此,只划分一次,输出到文件,并保留这些文件。

创建 src/main/java/com/example/fraud/training/SplitData.java

package com.example.fraud.training;

import java.io.*;
import java.nio.file.*;
import java.util.*;

/** 分层 70/30 划分为 train.csv 和 test.csv,在给定随机种子下可复现。 */
public class SplitData {

    private static final long SEED = 42L;
    private static final double TRAIN_FRACTION = 0.7;

    public static void main(String[] args) throws IOException {
        Path source = Path.of(args.length > 0 ? args[0] : "creditcard-balanced.csv");
        List<String> lines = Files.readAllLines(source);
        String header = lines.get(0);

        List<String> fraud = new ArrayList<>();
        List<String> legit = new ArrayList<>();
        for (String line : lines.subList(1, lines.size())) {
            if (line.isBlank()) continue;
            (line.trim().endsWith(",1") ? fraud : legit).add(line);
        }

        Random rnd = new Random(SEED);
        Collections.shuffle(fraud, rnd);
        Collections.shuffle(legit, rnd);

        List<String> train = new ArrayList<>(), test = new ArrayList<>();
        partition(fraud, train, test);
        partition(legit, train, test);
        Collections.shuffle(train, rnd);   // 避免按类别顺序排列的批次
        Collections.shuffle(test, rnd);

        write(Path.of("data/train.csv"), header, train);
        write(Path.of("data/test.csv"), header, test);
        System.out.printf("train=%d test=%d%n", train.size(), test.size());
    }

    private static void partition(List<String> rows, List<String> train, List<String> test) {
        int cut = (int) Math.round(rows.size() * TRAIN_FRACTION);
        train.addAll(rows.subList(0, cut));
        test.addAll(rows.subList(cut, rows.size()));
    }

    private static void write(Path path, String header, List<String> rows) throws IOException {
        Files.createDirectories(path.getParent());
        try (BufferedWriter w = Files.newBufferedWriter(path)) {
            w.write(header); w.newLine();
            for (String r : rows) { w.write(r); w.newLine(); }
        }
    }
}

运行一次。提交随机种子,而不是 CSV 文件。

endsWith(",1") 这个判断假设 Class 是最后一列,且没有尾随空格。在使用前请对照你的表头确认这一点。


第 3 步——训练并导出可部署的产物

训练会生成两个文件。所有人都记得导出模型,但被遗忘的那个文件才是大多数生产事故的根源。

创建 TrainFraudModel.java

package com.example.fraud.training;

import deepnetts.data.DataSet;
import deepnetts.data.DataSets;
import deepnetts.data.norm.MaxNormalizer;
import deepnetts.eval.Evaluators;
import deepnetts.eval.ClassifierEvaluationResult;
import deepnetts.net.FeedForwardNetwork;
import deepnetts.net.layers.activation.ActivationType;
import deepnetts.net.loss.LossType;
import deepnetts.net.train.BackpropagationTrainer;
import deepnetts.util.DeepNetts;
import deepnetts.util.FileIO;

import java.nio.file.*;
import java.util.*;

public class TrainFraudModel {

    private static final int NUM_INPUTS  = 30;   // Time, V1..V28, Amount
    private static final int NUM_OUTPUTS = 1;    // Class

    public static void main(String[] args) throws Exception {
        DataSet trainSet = DataSets.readCsv("data/train.csv", NUM_INPUTS, NUM_OUTPUTS, true, ",");
        DataSet testSet  = DataSets.readCsv("data/test.csv",  NUM_INPUTS, NUM_OUTPUTS, true, ",");

        // 仅使用训练集的统计数据进行缩放,然后对测试集应用同样的缩放。
        MaxNormalizer normalizer = new MaxNormalizer(trainSet);
        normalizer.normalize(trainSet);
        normalizer.normalize(testSet);

        // 独立计算同样的列最大值,以便在服务阶段持久化使用。
        float[] columnMax = columnMaxFrom("data/train.csv", NUM_INPUTS);
        writeScaler(Path.of("target/model/scaler.json"), columnMax);

        FeedForwardNetwork net = FeedForwardNetwork.builder()
                .addInputLayer(NUM_INPUTS)
                .addFullyConnectedLayer(32, ActivationType.RELU)
                .addFullyConnectedLayer(16, ActivationType.RELU)
                .addOutputLayer(NUM_OUTPUTS, ActivationType.SIGMOID)
                .lossFunction(LossType.CROSS_ENTROPY)
                .randomSeed(123)
                .build();

        BackpropagationTrainer trainer = net.getTrainer();
        trainer.setMaxError(0.03f)
               .setMaxEpochs(3000)
               .setLearningRate(0.01f);
        trainer.train(trainSet);

        ClassifierEvaluationResult result = Evaluators.evaluateClassifier(net, testSet);
        System.out.println(result);   // 在这里查看精确率/召回率/F1,忽略准确率

        Files.createDirectories(Path.of("target/model"));
        FileIO.writeToFile(net, "target/model/fraud-model.dnet");

        DeepNetts.shutdown();
    }

    /** 每个输入列的最大绝对值,与 MaxNormalizer 的缩放方式保持一致。 */
    private static float[] columnMaxFrom(String csv, int numInputs) throws Exception {
        float[] max = new float[numInputs];
        List<String> lines = Files.readAllLines(Path.of(csv));
        for (String line : lines.subList(1, lines.size())) {
            if (line.isBlank()) continue;
            String[] parts = line.split(",");
            for (int i = 0; i < numInputs; i++) {
                max[i] = Math.max(max[i], Math.abs(Float.parseFloat(parts[i].trim())));
            }
        }
        for (int i = 0; i < numInputs; i++) if (max[i] == 0f) max[i] = 1f;  // 防御处理
        return max;
    }

    private static void writeScaler(Path path, float[] max) throws Exception {
        Files.createDirectories(path.getParent());
        StringJoiner j = new StringJoiner(",", "{\"columnMax\":[", "]}");
        for (float m : max) j.add(Float.toString(m));
        Files.writeString(path, j.toString());
    }
}

为什么要重复计算最大值的逻辑? 服务在推理阶段需要缩放常量,而如何从 MaxNormalizer 中取出这些常量在不同版本中并不一致。从同一个文件中独立计算是一种不受版本影响的方案,而且只需要 10 行代码。如果你的 Deep Netts 版本能直接暴露这些值,可以加一个测试来断言两种方式的结果一致。

这是 Java ML 服务中最常见的生产环境 Bug: 模型部署了,但缩放器没有部署,未归一化的原始值被直接输入到基于 [-1, 1] 范围训练的网络中,所有输出都会饱和。系统不会抛出任何异常。你的欺诈率只会悄无声息地变成 0% 或 100%。

请始终把 fraud-model.dnetscaler.json 作为一对产物一起发布,并同步进行版本管理。


第 4 步——选择阈值(不要用 0.5)

在编写任何 Spring 代码之前,先决定什么样的分数意味着“欺诈”。将以下代码追加到训练过程的末尾:

System.out.println("threshold\tTP\tFP\tFN\tprecision\trecall");
for (float t = 0.05f; t < 1.0f; t += 0.05f) {
    int tp = 0, fp = 0, fn = 0;
    for (var item : testSet) {
        net.setInput(item.getInput());
        boolean predicted = net.getOutput()[0] >= t;
        boolean actual = item.getTargetOutput().get(0) >= 0.5f;
        if (predicted && actual) tp++;
        else if (predicted) fp++;
        else if (actual) fn++;
    }
    System.out.printf("%.2f\t%d\t%d\t%d\t%.3f\t%.3f%n",
            t, tp, fp, fn,
            tp + fp == 0 ? 0 : (double) tp / (tp + fp),
            tp + fn == 0 ? 0 : (double) tp / (tp + fn));
}

现在是基于成本做决策,而不是选一个看起来很整的数。漏报一次欺诈意味着一次拒付再加上欺诈损失;误报一次则意味着顾客的卡被拒绝、顾客不满、还要接一次客服电话。如果一次漏报欺诈的代价是误报拒绝的 40 倍,你应该选择较低的阈值,并接受随之而来的噪音。这是一个业务决策——把上面的表交给相关负责人,并把最终答案放在配置里,而不是硬编码在代码中。

还记得第 1.3 步提到的校准问题吗:在均衡数据集上训练的模型,其输出看起来像是概率,但实际上并不是。阈值表是经验性的、诚实的;原始分数不是概率。不要把它作为概率展示给用户。


第 5 步——创建 Spring Boot 项目

curl https://start.spring.io/starter.zip \
  -d dependencies=web,actuator,validation \
  -d javaVersion=17 -d type=maven-project \
  -d groupId=com.example -d artifactId=fraud-service \
  -o fraud-service.zip && unzip fraud-service.zip -d fraud-service

pom.xml 中添加:

<dependency>
    <groupId>com.deepnetts</groupId>
    <artifactId>deepnetts-core-pro</artifactId>
    <version>3.2.0</version>
</dependency>
<dependency>
    <groupId>com.deepnetts</groupId>
    <artifactId>deepnetts-license</artifactId>
    <version>1.0</version>
</dependency>
<dependency>
    <groupId>org.apache.commons</groupId>
    <artifactId>commons-pool2</artifactId>
    <version>2.12.0</version>
</dependency>

CI 会在这里失败。 本地 .m2 安装只能在你自己的电脑上生效,换一个环境就不行了。请用 mvn deploy:deploy-file 将这两个 Deep Netts JAR 包发布到你内部的 Nexus/Artifactory 仓库,并让 CI 指向该仓库。现在就做,不要等到构建服务器第一次出问题那天再处理。

src/main/resources/application.yml

fraud:
  model-path: classpath:model/fraud-model.dnet
  scaler-path: classpath:model/scaler.json
  threshold: 0.35          # 来自第 4 步,而不是教科书
  pool-size: 8
management:
  endpoints.web.exposure.include: health,metrics,prometheus
  endpoint.health.show-details: always

fraud-model.dnetscaler.json 复制到 src/main/resources/model/ 目录下。

---## 步骤 6 —— 加载模型和缩放器

FraudProperties.java:

package com.example.fraud;

import org.springframework.boot.context.properties.ConfigurationProperties;
import org.springframework.core.io.Resource;

@ConfigurationProperties(prefix = "fraud")
public record FraudProperties(Resource modelPath, Resource scalerPath,
                              float threshold, int poolSize) {}

ModelConfig.java:

package com.example.fraud;

import com.fasterxml.jackson.databind.ObjectMapper;
import deepnetts.net.FeedForwardNetwork;
import deepnetts.util.DeepNetts;
import deepnetts.util.FileIO;
import jakarta.annotation.PreDestroy;
import org.apache.commons.pool2.BasePooledObjectFactory;
import org.apache.commons.pool2.PooledObject;
import org.apache.commons.pool2.impl.*;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.*;

import java.io.*;
import java.nio.file.*;

@Configuration
@EnableConfigurationProperties(FraudProperties.class)
public class ModelConfig {

    /** Copy the model out of the jar once; pooled instances are deserialized from this file. */
    @Bean
    File modelFile(FraudProperties props) throws IOException {
        Path tmp = Files.createTempFile("fraud-model", ".dnet");
        tmp.toFile().deleteOnExit();
        try (InputStream in = props.modelPath().getInputStream()) {
            Files.copy(in, tmp, StandardCopyOption.REPLACE_EXISTING);
        }
        return tmp.toFile();
    }

    @Bean
    Scaler scaler(FraudProperties props, ObjectMapper mapper) throws IOException {
        try (InputStream in = props.scalerPath().getInputStream()) {
            return mapper.readValue(in, Scaler.class);
        }
    }

    @Bean(destroyMethod = "close")
    GenericObjectPool<FeedForwardNetwork> networkPool(File modelFile, FraudProperties props) {
        var config = new GenericObjectPoolConfig<FeedForwardNetwork>();
        config.setMaxTotal(props.poolSize());
        config.setMinIdle(props.poolSize());          // pre-warm: deserialization is slow
        config.setMaxWait(java.time.Duration.ofMillis(200));
        config.setBlockWhenExhausted(true);

        var pool = new GenericObjectPool<>(new BasePooledObjectFactory<FeedForwardNetwork>() {
            @Override public FeedForwardNetwork create() throws Exception {
                return FileIO.createFromFile(modelFile, FeedForwardNetwork.class);
            }
            @Override public PooledObject<FeedForwardNetwork> wrap(FeedForwardNetwork net) {
                return new DefaultPooledObject<>(net);
            }
        }, config);

        try { pool.preparePool(); } catch (Exception e) {
            throw new IllegalStateException("Could not initialise model pool", e);
        }
        return pool;
    }

    @PreDestroy
    void shutdownEngine() {
        DeepNetts.shutdown();   // Deep Netts runs its own thread pool
    }
}

这里有两件重要的事情:

对象池之所以存在,是因为 Deep Netts 网络是有状态的。 先调用 setInput(),再调用 getOutput(),这是针对实例字段的两步操作。如果两个并发 Tomcat 线程共享同一个网络实例,就会互相交错执行,并把对方的结果返回给错误的调用方。它不会抛异常,不会记录日志,只会出现在客户的投诉里。解决方案是使用一个由独立反序列化实例组成的对象池;如果你的吞吐量不大,用 synchronized 块是更简单的修复方式——先测量,再选择。

DeepNetts.shutdown() 不是可选项。 不调用它,上下文关闭会挂起,Pod 会一直等到完整终止宽限期结束才会停止。


步骤 7 —— 缩放器和评分服务

Scaler.java:

package com.example.fraud;

public record Scaler(float[] columnMax) {

    public float[] apply(float[] raw) {
        if (raw.length != columnMax.length) {
            throw new IllegalArgumentException(
                "Expected " + columnMax.length + " features, got " + raw.length);
        }
        float[] scaled = new float[raw.length];
        for (int i = 0; i < raw.length; i++) scaled[i] = raw[i] / columnMax[i];
        return scaled;
    }
}

FraudScoringService.java:

package com.example.fraud;

import deepnetts.net.FeedForwardNetwork;
import deepnetts.tensor.Tensor;
import io.micrometer.core.annotation.Timed;
import org.apache.commons.pool2.impl.GenericObjectPool;
import org.springframework.stereotype.Service;

@Service
public class FraudScoringService {

    private final GenericObjectPool<FeedForwardNetwork> pool;
    private final Scaler scaler;
    private final float threshold;

    public FraudScoringService(GenericObjectPool<FeedForwardNetwork> pool,
                               Scaler scaler, FraudProperties props) {
        this.pool = pool;
        this.scaler = scaler;
        this.threshold = props.threshold();
    }

    @Timed(value = "fraud.score", percentiles = {0.5, 0.95, 0.99})
    public ScoreResult score(float[] rawFeatures) {
        float[] input = scaler.apply(rawFeatures);
        FeedForwardNetwork net = null;
        try {
            net = pool.borrowObject();
            net.setInput(new Tensor(input));
            float score = net.getOutput()[0];
            return new ScoreResult(score, score >= threshold);
        } catch (Exception e) {
            throw new ScoringUnavailableException(e);
        } finally {
            if (net != null) pool.returnObject(net);
        }
    }

    public record ScoreResult(float score, boolean flagged) {}
}

注意 finally 块即使在失败时也会归还实例。如果漏掉这一步,错误条件下对象池会被耗尽,然后每个请求都会阻塞 200 毫秒并超时——这是一场缓慢而令人困惑的服务中断。

ScoringUnavailableException.java:

package com.example.fraud;

import org.springframework.http.HttpStatus;
import org.springframework.web.bind.annotation.ResponseStatus;

@ResponseStatus(HttpStatus.SERVICE_UNAVAILABLE)
public class ScoringUnavailableException extends RuntimeException {
    public ScoringUnavailableException(Throwable cause) { super("Scoring unavailable", cause); }
}

步骤 8 —— REST 层

把 30 个按位置排列的 float 放进 JSON 数组是一个很糟糕的对外契约——只要有一个字段被悄悄重排,模型就会把 Amount 当成 V7 来读。给这些字段起名字吧。

TransactionRequest.java:

package com.example.fraud;

import jakarta.validation.constraints.*;

public record TransactionRequest(
        @NotNull String transactionId,
        @NotNull Float time,
        @NotNull @Size(min = 28, max = 28) float[] v,   // V1..V28, order matters
        @NotNull @PositiveOrZero Float amount) {

    /** Must match the training CSV column order exactly: Time, V1..V28, Amount. */
    public float[] toFeatureVector() {
        float[] features = new float[30];
        features[0] = time;
        System.arraycopy(v, 0, features, 1, 28);
        features[29] = amount;
        return features;
    }
}

FraudController.java:

package com.example.fraud;

import jakarta.validation.Valid;
import org.slf4j.*;
import org.springframework.web.bind.annotation.*;

@RestController
@RequestMapping("/api/v1/transactions")
public class FraudController {

    private static final Logger log = LoggerFactory.getLogger(FraudController.class);
    private final FraudScoringService scoringService;

    public FraudController(FraudScoringService scoringService) {
        this.scoringService = scoringService;
    }

    @PostMapping("/score")
    public ScoreResponse score(@Valid @RequestBody TransactionRequest request) {
        var result = scoringService.score(request.toFeatureVector());
        log.info("scored transaction={} score={} flagged={}",
                 request.transactionId(), result.score(), result.flagged());
        return new ScoreResponse(request.transactionId(), result.score(),
                                 result.flagged(), "fraud-model-v1");
    }

    public record ScoreResponse(String transactionId, float score,
                                boolean flagged, String modelVersion) {}
}

modelVersion 字段不是装饰品。六个月后有人问某笔交易为什么被拒绝时,你需要知道是哪一版模型做出的决定。每次决策都要记录评分和版本——在大多数司法辖区,自动化金融决策需要审计跟踪,事后重建是不可能的。

试一试:

curl -X POST localhost:8080/api/v1/transactions/score \
  -H 'Content-Type: application/json' \
  -d '{"transactionId":"t-1","time":406,
       "v":[-2.31,1.95,-1.61,3.99,-0.52,-1.43,-2.54,1.39,-2.77,-2.77,
            3.20,-2.90,-0.60,-4.29,0.39,-1.14,-2.83,-0.02,0.42,0.13,
            0.52,-0.03,-0.47,0.32,0.04,0.18,0.26,-0.14],
       "amount":0.0}'

步骤 9 —— 健康检查与测试

只报告“Bean 存在”的健康检查毫无意义。用一个已知向量跑一遍模型:

package com.example.fraud;

import org.springframework.boot.actuate.health.*;
import org.springframework.stereotype.Component;

@Component
public class ModelHealthIndicator implements HealthIndicator {

    private static final float[] PROBE = new float[30];   // replace with a real known-good row
    private final FraudScoringService service;

    public ModelHealthIndicator(FraudScoringService service) { this.service = service; }

    @Override
    public Health health() {
        try {
            var result = service.score(PROBE);
            return Float.isNaN(result.score())
                    ? Health.down().withDetail("reason", "model returned NaN").build()
                    : Health.up().withDetail("probeScore", result.score()).build();
        } catch (Exception e) {
            return Health.down(e).build();
        }
    }
}

然后是黄金向量测试——这是 ML 服务中最有价值的测试,因为它能发现其他测试发现不了的模型/缩放器不匹配问题:

@SpringBootTest
class ScoringRegressionTest {

    @Autowired FraudScoringService service;

    @Test
    void knownFraudulentTransactionScoresHigh() {
        float[] features = { /* a row from test.csv with Class=1 */ };
        assertThat(service.score(features).score()).isGreaterThan(0.7f);
    }

    @Test
    void knownLegitimateTransactionScoresLow() {
        float[] features = { /* a row from test.csv with Class=0 */ };
        assertThat(service.score(features).score()).isLessThan(0.3f);
    }
}

test.csv 中每类各取 10 条,硬编码到测试里。当有人换了模型却没有同步替换对应的缩放器时,让测试明确地失败。

还有一点值得补充:一个并发测试,用 50 个线程同时提交同一个向量,并断言所有响应都一致。如果你跳过了对象池,这个测试就能发现问题。


步骤 10 —— 打包与部署

模型位置。.dnet 打包进 jar 是最简单的做法,能让你获得不可变、原子化部署的构建产物。这也意味着每次重新训练都要完整部署一次应用。另一种做法——从存储卷挂载,或在启动时从 S3 拉取——将两者解耦,但要求你严格对模型/缩放器配对进行版本管理,并在启动时处理损坏的产物。先采用打包方式;当重新训练的频率真正带来困扰时,再把它挪出来。

Dockerfile —— Deep Netts 的 jar 必须存在于镜像中,这意味着你的构建阶段需要能访问内部制品库:

FROM maven:3.9-eclipse-temurin-17 AS build
WORKDIR /app
COPY settings.xml /root/.m2/settings.xml    # points at your internal Nexus
COPY pom.xml .
RUN mvn -B dependency:go-offline
COPY src ./src
RUN mvn -B clean package -DskipTests

FROM eclipse-temurin:17-jre
COPY --from=build /app/target/fraud-service-*.jar /app/app.jar
ENV JAVA_OPTS="-Xmx1g -XX:MaxRAMPercentage=75"
ENTRYPOINT ["sh","-c","java $JAVA_OPTS -jar /app/app.jar"]

Deep Netts 是纯 Java 实现,因此所有内容都在堆上——不会有堆外内存的意外,但设置 -Xmx 时要按 池大小 × 模型大小 再加上正常的应用开销来计算。大型网络的 8 个副本就是 8 倍内存。

就绪探针与存活探针。 启动时反序列化 8 个模型实例需要真实时间。将就绪探针指向 /actuator/health/readiness,并给它一个充足的 initialDelaySeconds;否则 Kubernetes 会在预热途中杀死 Pod,你永远无法达到稳定状态。


步骤 11 —— 上线前的许可证问题

免费层级对这种服务来说限制确实相当严格。它只允许部署在不超过一个生产环境中,要求通过该产品产生的年收入低于 10 万美元,且公司总年收入低于 100 万美元,并明确禁止用于运营或支持任何托管 AI 平台、托管服务或 SaaS 产品

小型公司中单个生产环境内的内部欺诈评分服务可能符合要求。任何面向客户、多区域,或作为服务出售的用途都不符合。在基于免费层级规划路线图之前,请阅读 EULA 并与 Deep Netts 沟通。


步骤 12 —— 生产环境还需要什么

服务现在能运行了,但下面这些才是它与演示项目区分开来的事项:

  • 重训练管道。 欺诈模式每个月都会变化。没有定期重训练的模型是一种不断贬值的资产。在 CI 中自动化第 2–4 步,并依据阈值表来把关发布,而不是只看单一指标。
  • 漂移监控。 输出评分的直方图。当分布发生偏移时,说明世界已经先于你的指标发生了变化。这是你最早的预警,成本只是一个 Micrometer 计数器。
  • 先以影子模式运行。 将评分与现有规则引擎并行部署,记录两种决策,不改变任何行为。运行几周,进行对比,然后再依据模型行动。绝不要让一个全新的模型第一天就去拒绝交易。
  • 一个应急开关。 一个配置开关,可以绕过模型并回退到规则引擎,无需部署即可切换。
  • 特征向量中不要包含 PII(个人身份信息)。 Kaggle 数据已经为你做了匿名化。你自己的特征不会——请慎重决定哪些信息进入模型,哪些信息写入日志。

参考:项目结构

fraud-service/
├── pom.xml
├── data/
│   ├── train.csv                    # generated, gitignored
│   └── test.csv
└── src/
    ├── main/
    │   ├── java/com/example/fraud/
    │   │   ├── FraudApplication.java
    │   │   ├── FraudProperties.java
    │   │   ├── ModelConfig.java
    │   │   ├── Scaler.java
    │   │   ├── FraudScoringService.java
    │   │   ├── ScoringUnavailableException.java
    │   │   ├── ModelHealthIndicator.java
    │   │   ├── TransactionRequest.java
    │   │   ├── FraudController.java
    │   │   └── training/
    │   │       ├── SplitData.java
    │   │       └── TrainFraudModel.java
    │   └── resources/
    │       ├── application.yml
    │       └── model/
    │           ├── fraud-model.dnet
    │           └── scaler.json
    └── test/java/com/example/fraud/
        └── ScoringRegressionTest.java

训练相关类先放在同一个模块里是可以的。一旦训练依赖(tablesaw、绘图库)开始让服务镜像变得臃肿,就把它们拆分到独立的 fraud-training 模块中。

《如何创建一个 Spring Boot 欺诈评分服务》一文最初发布在 foojay 上。