欢迎光临
我们一直在努力

【Java PyTorch深度学习】PyTorch On Java 进阶课程 Flink特征工程 与PyTorch实时特征工程与流式推荐系统[PyTorch Java 硕士研一课程]

在这里插入图片描述

PyTorch Java 高校计算机硕士研一课程

Apache Flink 最新版集成 JavaCPP-PyTorch 2.10-1.5.13 实战:实时特征工程与流式推荐系统(MIND 算法)

在电商推荐场景中,实时性与个性化是提升用户体验与转化的核心。传统方案常面临实时特征计算延迟高、流式数据训练衔接繁琐、多语言部署成本高三大痛点。本文将基于Apache Flink 最新稳定版 2.2.0 与 JavaCPP-PyTorch 2.10-1.5.13,打造一套全 Java 生态的实时推荐解决方案:以 Flink 完成低延迟实时特征工程,通过流式数据集(Streaming Dataset)串联数据至 PyTorch,基于 MIND(Multi-Interest Network with Dynamic Routing)算法(严格遵循官方论文实现)实现电商商品的个性化推荐,实现“实时特征→流式训练→精准推荐”端到端闭环。

本文全程采用 Markdown 语法,完整覆盖 Flink 实时特征工程、Flink 与 PyTorch Streaming Dataset 串联、JavaCPP-PyTorch 实现标准 MIND 模型(嵌入层、行为胶囊聚合、动态路由、标签感知注意力)、模型流式训练与推荐推理全流程,可直接复制作为 MD 源文件使用,适配生产级电商推荐场景。

一、技术选型与核心概念

1. 核心组件与版本匹配

组件版本核心价值
Apache Flink 2.2.0 最新稳定版 提供存算分离架构与 Flink ML 2.0,支持高吞吐、低延迟实时特征工程与状态化计算,可消费 Kafka 实时用户行为流,生成用户实时兴趣特征,是实时推荐数据处理的基石。
JavaCPP-PyTorch 2.10-1.5.13 基于 JavaCPP 封装的 PyTorch Java 绑定库,无需安装 Python、无需部署 PyTorch 环境,直接在 JVM 上调用 PyTorch C++ 核心 API,严格遵循 MIND 官方论文实现模型,支持 CPU/GPU 加速,完美适配 Java/Flink 生态。
MIND 算法 官方论文标准实现 推荐召回层核心算法,解决用户兴趣多样性问题,通过“嵌入层→行为胶囊聚合→动态路由→标签感知注意力”四步,将用户行为序列提取为多个兴趣向量,相比单一向量召回精度提升 20%+。
PyTorch Streaming Dataset 适配 PyTorch 2.1 核心 以流式方式读取 Flink 输出的实时特征数据,实现“边读边训”,打破离线训练与实时数据的壁垒,确保模型能及时学习用户最新兴趣。
Kafka 2.8.0+ 作为实时数据源,存储用户行为流(点击、加购、收藏、浏览),供 Flink 消费处理。

2. 整体架构(端到端闭环)

整个系统分为 3 层,全程基于 Java 生态,无跨语言调用,架构如下:

  • 实时数据接入层:用户行为数据(点击、加购等)实时写入 Kafka 主题,商品维度数据(商品 ID、类目、价格等)存储在 Redis/HBase,供 Flink 实时关联。

  • 实时特征工程层(Flink 2.2.0):消费 Kafka 实时用户行为流,通过 Keyed State、滑动窗口(如 1 小时窗口),结合商品维度数据,生成用户多维度实时特征(用户近期行为序列、商品类目偏好、行为时间特征等),通过自定义 Sink 输出至 PyTorch Streaming Dataset。

  • 流式训练与推荐层(JavaCPP-PyTorch 2.10-1.5.13):通过 Streaming Dataset 读取 Flink 实时特征流,基于 MIND 官方论文实现多兴趣提取模型,进行流式训练;训练完成后,生成用户多兴趣向量,结合商品向量库,通过向量检索实时生成个性化商品推荐列表,推送至前端。

  • 3. MIND 算法官方论文核心逻辑(必遵循)

    严格参考 MIND 官方论文《Multi-Interest Network with Dynamic Routing for Recommendation at Tmall》,核心逻辑分为 4 步,也是本文模型实现的核心依据:

    • 嵌入层(Embedding Layer):将用户 ID、商品 ID、商品类目等离散特征,通过嵌入矩阵转换为固定维度的稠密向量(本文设为 64 维),为后续特征聚合奠定基础。

    • 行为胶囊聚合(Behavior Capsule Aggregation):将用户历史行为序列(如最近 10 次点击的商品向量),通过胶囊网络(Capsule Network)的 Squash 激活函数,聚合为多个行为胶囊(Behavior Capsule),每个行为胶囊对应用户的一个局部兴趣。

    • 动态路由(Dynamic Routing):通过迭代路由算法(默认 3 轮迭代),计算行为胶囊与兴趣胶囊(Interest Capsule)的耦合系数,将行为胶囊动态分配至 K 个兴趣胶囊(本文 K=4,即提取用户 4 个核心兴趣),最终生成用户多兴趣向量。

    • 标签感知注意力(Label-Aware Attention):训练时,根据目标商品的嵌入向量,对用户的 K 个兴趣向量计算注意力权重,选择最相关的兴趣向量与目标商品向量计算损失,提升推荐精准度(解决单一兴趣向量无法覆盖用户多样需求的问题)。

    二、环境准备

  • 开发环境:JDK 17+(必须,适配 Flink 2.2.0 与 JavaCPP-PyTorch 最新依赖)、Maven 3.8+、IntelliJ IDEA、Flink 2.2.0 本地集群(含 Flink ML 2.0)。

  • 中间件环境:Kafka 2.8.0+(启动 ZooKeeper 或使用 Kafka 内置 ZooKeeper)、Redis 6.2+(存储商品维度表,可选,也可使用本地缓存用于测试)。

  • 依赖准备:配置 Flink 环境变量 FLINK_HOME,确保 bin/flink 可正常执行;通过 Maven 引入 Flink、JavaCPP-PyTorch 核心依赖,自动下载 PyTorch 底层运行时(约 600MB,首次构建耗时较长,属正常现象)。

  • GPU 环境(可选):若需 GPU 加速模型训练,安装 NVIDIA CUDA 11.8+、cuDNN 8.6+,JavaCPP-PyTorch 会自动识别 GPU 并启用加速,训练速度提升 3-10 倍。

  • 数据准备:

    • 用户行为数据:构造 CSV 测试数据(含 userId、itemId、behaviorType、timestamp),导入 Kafka 主题user-behavior-topic。

    • 商品维度数据:构造商品信息(itemId、itemCategory、itemPrice),导入 Redis 或本地缓存。

  • 三、项目搭建与核心依赖(pom.xml)

    1. 项目初始化

    使用 Maven 骨架创建 Flink 项目,执行以下命令:

    mvn archetype:generate \\
    -DarchetypeGroupId=org.apache.flink \\
    -DarchetypeArtifactId=flink-quickstart-java \\
    -DarchetypeVersion=2.2.0

    2. 核心 Maven 依赖(关键,适配所有组件版本)

    引入 Flink 实时处理、Kafka 连接器、JavaCPP-PyTorch 核心依赖,确保 MIND 模型能基于 JavaCPP-PyTorch 正常运行,pom.xml 配置如下:

    <?xml version="1.0" encoding="UTF-8"?>
    <project xmlns="http://maven.apache.org/POM/4.0.0" xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
    xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 https://maven.apache.org/xsd/maven-4.0.0.xsd">

    <modelVersion>4.0.0</modelVersion>
    <parent>
    <groupId>org.springframework.boot</groupId>
    <artifactId>spring-boot-starter-parent</artifactId>
    <version>3.2.5</version>
    <relativePath/>
    </parent>
    <groupId>com.example</groupId>
    <artifactId>flink-mind-recommend</artifactId>
    <version>0.0.1-SNAPSHOT</version>
    <name>flink-mind-recommend</name>
    <description>Flink 2.2.0 集成 JavaCPP-PyTorch 2.10-1.5.13 实现 MIND 算法实时推荐系统</description>

    <properties>
    <java.version>17</java.version>
    <flink.version>2.2.0</flink.version>
    <javacpp.version>1.5.13</javacpp.version>
    <pytorch.version>2.10-1.5.13</pytorch.version>
    <scala.binary.version>2.12</scala.binary.version>
    <kafka.version>2.8.0</kafka.version>
    </properties>

    <dependencies>
    <!– Flink 核心依赖 –>
    <dependency>
    <groupId>org.apache.flink</groupId>
    <artifactId>flink-streaming-java</artifactId>
    <version>${flink.version}</version>
    </dependency>
    <dependency>
    <groupId>org.apache.flink</groupId>
    <artifactId>flink-connector-kafka</artifactId>
    <version>${flink.version}</version>
    </dependency>
    <dependency>
    <groupId>org.apache.flink</groupId>
    <artifactId>flink-ml-api</artifactId>
    <version>${flink.version}</version>
    </dependency>
    <dependency>
    <groupId>org.apache.flink</groupId>
    <artifactId>flink-statebackend-rocksdb</artifactId>
    <version>${flink.version}</version>
    </dependency>

    <!– JavaCPP-PyTorch 核心依赖(MIND 模型实现核心) –>
    <dependency>
    <groupId>org.bytedeco</groupId>
    <artifactId>javacpp</artifactId>
    <version>${javacpp.version}</version>
    </dependency>
    <dependency>
    <groupId>org.bytedeco</groupId>
    <artifactId>pytorch-platform</artifactId>
    <version>${pytorch.version}</version>
    </dependency>

    <!– 辅助依赖 –>
    <dependency>
    <groupId>org.projectlombok</groupId>
    <artifactId>lombok</artifactId>
    <optional>true</optional>
    </dependency>
    <dependency>
    <groupId>com.google.code.gson</groupId>
    <artifactId>gson</artifactId>
    <version>2.10.1</version>
    </dependency>
    <dependency>
    <groupId>redis.clients</groupId>
    <artifactId>jedis</artifactId>
    <version>4.4.6</version>
    </dependency>
    <dependency>
    <groupId>org.apache.flink</groupId>
    <artifactId>flink-clients</artifactId>
    <version>${flink.version}</version>
    <scope>provided</scope>
    </dependency>
    <dependency>
    <groupId>junit</groupId>
    <artifactId>junit</artifactId>
    <version>4.13.2</version>
    <scope>test</scope>
    </dependency>
    </dependencies>

    <build>
    <plugins>
    <plugin>
    <groupId>org.apache.maven.plugins</groupId>
    <artifactId>maven-compiler-plugin</artifactId>
    <configuration>

    <target>${java.version}</target>
    <encoding>UTF-8</encoding>
    </configuration>
    </plugin>
    <!– 打包插件,适配 Flink 集群部署,避免依赖冲突 –>
    <plugin>
    <groupId>org.apache.maven.plugins</groupId>
    <artifactId>maven-shade-plugin</artifactId>
    <version>3.5.1</version>
    <executions>
    <execution>
    <phase>package</phase>
    <goals>
    <goal>shade</goal>
    </goals>
    <configuration>
    <filters>
    <filter>
    <artifact>*:*</artifact>
    <excludes>
    <exclude>META-INF/*.SF</exclude>
    <exclude>META-INF/*.DSA</exclude>
    <exclude>META-INF/*.RSA</exclude>
    </excludes>
    </filter>
    </filters>
    <transformers>
    <transformer implementation="org.apache.maven.plugins.shade.resource.ManifestResourceTransformer"><mainClass>com.example.FlinkFeatureEngineering</mainClass>
    </transformer>
    </transformers>
    </configuration>
    </execution>
    </executions>
    </plugin>
    </plugins>
    </build>
    </project>

    四、Flink 2.2.0 实时特征工程实现(适配 MIND 模型输入)

    Flink 作为实时特征工程核心,需生成 MIND 模型所需的输入数据:用户 ID、商品 ID 序列、商品类目序列、行为时间特征等,同时通过自定义 Sink 串联 PyTorch Streaming Dataset,确保数据实时流转。

    1. 数据模型定义(贴合 MIND 模型输入需求)

    定义用户行为实体、商品维度实体、MIND 模型输入特征实体,确保字段与模型输入一一对应:

    import lombok.AllArgsConstructor;
    import lombok.Data;
    import lombok.NoArgsConstructor;
    import java.util.List;

    /**
    * 用户行为实体(Kafka 消费的原始数据)
    * behaviorType:click(点击)、cart(加购)、fav(收藏)、view(浏览)
    */

    @Data
    @NoArgsConstructor
    @AllArgsConstructor
    public class UserBehavior {
    private String userId; // 用户ID(离散特征)
    private String itemId; // 商品ID(离散特征)
    private String behaviorType; // 行为类型
    private long timestamp; // 行为时间戳(用于窗口计算)
    }

    /**
    * 商品维度实体(Redis/本地缓存存储)
    */

    @Data
    @AllArgsConstructor
    @NoArgsConstructor
    public class ItemInfo {
    private String itemId; // 商品ID
    private String itemCategory; // 商品类目(离散特征)
    private float itemPrice; // 商品价格(连续特征)
    }

    /**
    * MIND 模型输入特征实体(Flink 特征工程输出,适配 PyTorch 输入)
    * 严格对应 MIND 模型需求:用户ID、商品行为序列、类目序列、目标商品ID(用于训练)
    */

    @Data
    @NoArgsConstructor
    @AllArgsConstructor
    public class MINDInputFeature {
    private String userId; // 用户ID(嵌入层输入)
    private List<String> itemIdSequence; // 用户近期商品行为序列(如最近10次点击)
    private List<String> categorySequence; // 商品类目序列(与商品序列一一对应)
    private String targetItemId; // 目标商品ID(训练标签,用于计算损失)
    private float targetItemPrice; // 目标商品价格(辅助特征)
    private long timestamp; // 特征生成时间戳
    }

    /**
    * 商品向量实体(MIND 模型嵌入层输出,用于后续推荐检索)
    */

    @Data
    @AllArgsConstructor
    @NoArgsConstructor
    public class ItemEmbedding {
    private String itemId;
    private float[] embeddingVec; // 商品嵌入向量(64维)
    }
    }

    2. Flink 实时特征工程核心逻辑(低延迟、状态化)

    基于 Flink 2.2.0 的 Keyed State、滑动窗口(1 小时窗口,步长 5 分钟),消费 Kafka 用户行为流,关联商品维度数据,生成 MIND 模型所需的输入特征,核心逻辑如下:

    import org.apache.flink.api.common.eventtime.WatermarkStrategy;
    import org.apache.flink.api.common.functions.MapFunction;
    import org.apache.flink.api.common.serialization.SimpleStringSchema;
    import org.apache.flink.api.common.state.ListState;
    import org.apache.flink.api.common.state.ListStateDescriptor;
    import org.apache.flink.api.common.typeinfo.TypeHint;
    import org.apache.flink.api.common.typeinfo.TypeInformation;
    import org.apache.flink.configuration.Configuration;
    import org.apache.flink.connector.kafka.source.KafkaSource;
    import org.apache.flink.connector.kafka.source.enumerator.initializer.OffsetsInitializer;
    import org.apache.flink.streaming.api.datastream.DataStream;
    import org.apache.flink.streaming.api.environment.StreamExecutionEnvironment;
    import org.apache.flink.streaming.api.functions.KeyedProcessFunction;
    import org.apache.flink.util.Collector;
    import com.google.gson.Gson;
    import redis.clients.jedis.Jedis;
    import java.time.Duration;
    import java.util.ArrayList;
    import java.util.List;
    import java.util.stream.Collectors;

    /**
    * Flink 2.2.0 实时特征工程(适配 MIND 模型输入)
    * 核心:生成用户行为序列、类目序列,关联目标商品,输出适配 PyTorch Streaming Dataset 的特征流
    */

    public class FlinkFeatureEngineering {
    private static final Gson GSON = new Gson();
    private static final int MAX_BEHAVIOR_SIZE = 10; // 用户行为序列最大长度(MIND 模型默认)
    private static final long WINDOW_SIZE = 3600000; // 滑动窗口大小:1小时(毫秒)
    private static final long WINDOW_STEP = 300000; // 窗口步长:5分钟(毫秒)

    public static void main(String[] args) throws Exception {
    // 1. 初始化 Flink 执行环境(适配 Flink 2.2.0 新特性)
    StreamExecutionEnvironment env = StreamExecutionEnvironment.getExecutionEnvironment();
    env.setParallelism(1);
    env.enableCheckpointing(5000); // 开启检查点,保证状态一致性(生产级必备)
    // 配置 RocksDB 状态后端(存储用户行为序列状态,支持海量数据)
    env.setStateBackend(new org.apache.flink.contrib.streaming.state.RocksDBStateBackend());

    // 2. 构建 Kafka 源,消费用户行为数据
    KafkaSource<String> kafkaSource = KafkaSource.<String>builder()
    .setBootstrapServers("localhost:9092")
    .setTopics("user-behavior-topic")
    .setValueOnlyDeserializer(new SimpleStringSchema())
    .setStartingOffsets(OffsetsInitializer.latest()) // 从最新偏移量开始消费
    .build();

    // 3. 读取 Kafka 流,解析为 UserBehavior 实体
    DataStream<UserBehavior> behaviorStream = env.fromSource(
    kafkaSource,
    WatermarkStrategy.forMonotonousTimestamps(), // 单调时间戳水印策略
    "Kafka User Behavior Source"
    )
    .map(json -> GSON.fromJson(json, UserBehavior.class)) // JSON 解析
    .filter(behavior -> "click".equals(behavior.getBehaviorType())); // 仅保留点击行为(核心行为)

    // 4. 实时特征计算:生成用户行为序列、类目序列,关联目标商品
    DataStream<MINDInputFeature> mindFeatureStream = behaviorStream
    .keyBy(UserBehavior::getUserId) // 按用户ID分组,保证同一用户的行为有序处理
    .process(new UserBehaviorSequenceProcessFunction());

    // 5. 打印特征流(测试用),并通过自定义 Sink 串联 PyTorch Streaming Dataset
    mindFeatureStream.print("MIND Model Input Feature");
    mindFeatureStream.addSink(new PyTorchStreamingSink()); // 关键:串联 PyTorch 流式数据

    // 6. 提交 Flink 任务
    env.execute("Flink 2.2.0 Real-Time Feature Engineering for MIND");
    }

    /**
    * 核心 ProcessFunction:维护用户行为序列(状态化),生成 MIND 模型输入特征
    */

    static class UserBehaviorSequenceProcessFunction extends KeyedProcessFunction<String, UserBehavior, MINDInputFeature> {
    // 存储用户最近 MAX_BEHAVIOR_SIZE 次点击行为(状态化存储,避免数据丢失)
    private ListState<UserBehavior> userBehaviorState;
    // Redis 客户端(关联商品维度数据)
    private Jedis jedis;

    @Override
    public void open(Configuration parameters) {
    // 初始化状态:存储用户行为序列
    userBehaviorState = getRuntimeContext().getState(
    new ListStateDescriptor<>(
    "userClickBehaviorState",
    TypeInformation.of(new TypeHint<UserBehavior>() {})
    )
    );
    // 初始化 Redis 客户端(生产级需使用连接池)
    jedis = new Jedis("localhost", 6379);
    }

    @Override
    public void processElement(UserBehavior behavior, Context ctx, Collector<MINDInputFeature> out) throws Exception {
    // 1. 从状态中获取用户历史行为
    List<UserBehavior> historyBehaviors = new ArrayList<>();
    userBehaviorState.get().forEach(historyBehaviors::add);

    // 2. 添加当前行为,保持序列长度不超过 MAX_BEHAVIOR_SIZE
    historyBehaviors.add(behavior);
    if (historyBehaviors.size() > MAX_BEHAVIOR_SIZE) {
    historyBehaviors.remove(0); // 移除最早的行为
    }

    // 3. 更新状态
    userBehaviorState.update(historyBehaviors);

    // 4. 关联商品维度数据(从 Redis 获取)
    String itemId = behavior.getItemId();
    String itemInfoJson = jedis.get("item:" + itemId);
    ItemInfo itemInfo = GSON.fromJson(itemInfoJson, ItemInfo.class);
    if (itemInfo == null) {
    return; // 商品信息不存在,跳过当前行为
    }

    // 5. 生成用户商品行为序列、类目序列(取历史所有行为)
    List<String> itemIdSequence = historyBehaviors.stream()
    .map(UserBehavior::getItemId)
    .collect(Collectors.toList());
    List<String> categorySequence = historyBehaviors.stream()
    .map(ub -> {
    String infoJson = jedis.get("item:" + ub.getItemId());
    return GSON.fromJson(infoJson, ItemInfo.class).getItemCategory();
    })
    .collect(Collectors.toList());

    // 6. 构造 MIND 模型输入特征(当前行为作为目标商品,历史序列作为输入)
    MINDInputFeature inputFeature = new MINDInputFeature(
    behavior.getUserId(),
    itemIdSequence,
    categorySequence,
    itemId, // 目标商品ID(训练标签)
    itemInfo.getItemPrice(),
    behavior.getTimestamp()
    );

    // 7. 输出特征(触发窗口逻辑,每5分钟输出一次特征)
    ctx.timerService().registerEventTimeTimer(behavior.getTimestamp() (behavior.getTimestamp() % WINDOW_STEP) + WINDOW_SIZE);
    out.collect(inputFeature);
    }

    // 定时器:清理过期窗口的状态数据
    @Override
    public void onTimer(long timestamp, OnTimerContext ctx, Collector<MINDInputFeature> out) throws Exception {
    List<UserBehavior> historyBehaviors = new ArrayList<>();
    userBehaviorState.get().forEach(historyBehaviors::add);
    // 移除过期行为(超过1小时的行为)
    historyBehaviors.removeIf(behavior -> behavior.getTimestamp() < timestamp WINDOW_SIZE);
    userBehaviorState.update(historyBehaviors);
    }

    @Override
    public void close() {
    // 关闭 Redis 客户端,释放资源
    if (jedis != null) {
    jedis.close();
    }
    }
    }
    }

    3. 自定义 Sink:串联 Flink 与 PyTorch Streaming Dataset

    自定义 Flink Sink,将实时特征流写入内存队列(测试用)或 Kafka(生产用),供 PyTorch Streaming Dataset 实时读取,实现 Flink 与 PyTorch 的无缝串联,核心是保证数据流式传输,适配“边读边训”:

    import org.apache.flink.streaming.api.functions.sink.SinkFunction;
    import com.example.MINDInputFeature;
    import com.google.gson.Gson;
    import java.util.Queue;
    import java.util.concurrent.ConcurrentLinkedQueue;

    /**
    * 自定义 Sink:串联 Flink 特征流与 PyTorch Streaming Dataset
    * 核心:将 Flink 输出的 MINDInputFeature 转换为 JSON 字符串,写入线程安全队列,供 PyTorch 读取
    * 生产级可替换为 Kafka Sink,提升可靠性
    */

    public class PyTorchStreamingSink implements SinkFunction<MINDInputFeature> {
    private static final Gson GSON = new Gson();
    // 线程安全队列:作为 Flink 与 PyTorch 之间的流式数据桥梁
    public static final Queue<String> MIND_FEATURE_QUEUE = new ConcurrentLinkedQueue<>();

    @Override
    public void invoke(MINDInputFeature feature, Context context) {
    // 将 MIND 模型输入特征转换为 JSON 字符串,写入队列
    String featureJson = GSON.toJson(feature);
    MIND_FEATURE_QUEUE.offer(featureJson);
    // 控制队列大小,避免内存溢出(生产级可结合 Kafka 实现持久化)
    if (MIND_FEATURE_QUEUE.size() > 10000) {
    MIND_FEATURE_QUEUE.poll();
    }
    }
    }

    五、JavaCPP-PyTorch 2.10-1.5.13 实现 MIND 模型(严格遵循官方论文)

    这是本文核心,基于 JavaCPP-PyTorch 2.10-1.5.13 严格遵循 MIND 官方论文,实现嵌入层、行为胶囊聚合、动态路由、标签感知注意力全流程,同时适配 PyTorch Streaming Dataset 读取 Flink 实时特征流,完成模型流式训练与推荐推理。

    1. 核心工具类:数据转换与 Streaming Dataset 实现

    实现 PyTorch Streaming Dataset,读取 Flink Sink 输出的实时特征流,同时提供数据转换方法(将 JSON 特征转换为 PyTorch 张量),适配 MIND 模型输入:

    import org.bytedeco.pytorch.*;
    import org.bytedeco.pytorch.global.torch;
    import com.example.MINDInputFeature;
    import com.google.gson.Gson;
    import java.util.ArrayList;
    import java.util.List;
    import java.util.Queue;

    /**
    * PyTorch Streaming Dataset(适配 JavaCPP-PyTorch 2.10-1.5.13)
    * 核心:实时读取 Flink 输出的特征流,转换为 PyTorch 张量,供 MIND 模型流式训练
    */

    public class MINDStreamingDataset extends Dataset {
    private static final Gson GSON = new Gson();
    private final Queue<String> featureQueue; // Flink Sink 输出的特征队列
    private final int embeddingDim; // 嵌入向量维度(64维,与 MIND 论文一致)
    private final int maxSequenceLen; // 行为序列最大长度(10)

    // 嵌入层字典(用户ID、商品ID、类目ID → 索引映射,用于嵌入矩阵查找)
    private final Vocab userIdVocab;
    private final Vocab itemIdVocab;
    private final Vocab categoryVocab;

    public MINDStreamingDataset(Queue<String> featureQueue, int embeddingDim, int maxSequenceLen,
    Vocab userIdVocab, Vocab itemIdVocab, Vocab categoryVocab) {
    this.featureQueue = featureQueue;
    this.embeddingDim = embeddingDim;
    this.maxSequenceLen = maxSequenceLen;
    this.userIdVocab = userIdVocab;
    this.itemIdVocab = itemIdVocab;
    this.categoryVocab = categoryVocab;
    }

    /**
    * 重写 size 方法:返回当前队列中的特征数量(流式数据,动态变化)
    */

    @Override
    public long size() {
    return featureQueue.size();
    }

    /**
    * 重写 get 方法:获取指定索引的特征,转换为 PyTorch 张量
    */

    @Override
    public IValue get(long index) {
    // 从队列中获取特征 JSON 字符串
    String featureJson = featureQueue.poll();
    if (featureJson == null) {
    throw new RuntimeException("No more streaming features from Flink");
    }
    MINDInputFeature feature = GSON.fromJson(featureJson, MINDInputFeature.class);

    // 1. 转换用户ID为索引张量(嵌入层输入)
    long userIdIdx = userIdVocab.getIndex(feature.getUserId());
    Tensor userIdTensor = torch.tensor(new long[]{userIdIdx}, torch.torch_int64());

    // 2. 转换商品行为序列为索引张量(补齐/截断至 maxSequenceLen)
    List<String> itemIdSeq = feature.getItemIdSequence();
    long[] itemIdIdxSeq = new long[maxSequenceLen];
    for (int i = 0; i < maxSequenceLen; i++) {
    if (i < itemIdSeq.size()) {
    itemIdIdxSeq[i] = itemIdVocab.getIndex(itemIdSeq.get(i));
    } else {
    itemIdIdxSeq[i] = 0; // 补齐用0(未知商品)
    }
    }
    Tensor itemSeqTensor = torch.tensor(itemIdIdxSeq, torch.torch_int64()).reshape(new long[]{1, maxSequenceLen});

    // 3. 转换类目序列为索引张量
    List<String> categorySeq = feature.getCategorySequence();
    long[] categoryIdxSeq = new long[maxSequenceLen];
    for (int i = 0; i < maxSequenceLen; i++) {
    if (i < categorySeq.size()) {
    categoryIdxSeq[i] = categoryVocab.getIndex(categorySeq.get(i));
    } else {
    categoryIdxSeq[i] = 0; // 补齐用0(未知类目)
    }
    }
    Tensor categorySeqTensor = torch.tensor(categoryIdxSeq, torch.torch_int64()).reshape(new long[]{1, maxSequenceLen});

    // 4. 转换目标商品ID为索引张量(训练标签)
    long targetItemIdx = itemIdVocab.getIndex(feature.getTargetItemId());
    Tensor targetItemTensor = torch.tensor(new long[]{targetItemIdx}, torch.torch_int64());

    // 5. 转换目标商品价格为张量(辅助特征)
    Tensor targetPriceTensor = torch.tensor(new float[]{feature.getTargetItemPrice()}, torch.torch_float32());

    // 返回特征字典(适配 MIND 模型输入)
    return IValue.mapFrom(
    "userId", userIdTensor,
    "itemSeq", itemSeqTensor,
    "categorySeq", categorySeqTensor,
    "targetItem", targetItemTensor,
    "targetPrice", targetPriceTensor
    );
    }

    /**
    * 词汇表工具类:将离散特征(用户ID、商品ID等)映射为整数索引,用于嵌入层
    */

    public static class Vocab {
    private final java.util.Map<String, Long> vocabMap;
    private long nextIndex = 1; // 0 用于未知值

    public Vocab() {
    this.vocabMap = new java.util.HashMap<>();
    }

    // 获取特征对应的索引,不存在则新增
    public long getIndex(String key) {
    return vocabMap.computeIfAbsent(key, k -> nextIndex++);
    }

    // 获取词汇表大小
    public long size() {
    return nextIndex;
    }
    }
    }

    2. MIND 模型完整实现(严格遵循官方论文)

    基于 JavaCPP-PyTorch 2.10-1.5.13 API,严格实现 MIND 论文中的 4 个核心模块,每个模块对应论文中的具体逻辑,注释详细,可直接复用:

    import org.bytedeco.pytorch.*;
    import org.bytedeco.pytorch.global.torch;
    import java.io.File;

    /**
    * MIND 模型(Multi-Interest Network with Dynamic Routing)
    * 严格遵循官方论文实现,基于 JavaCPP-PyTorch 2.10-1.5.13
    * 核心模块:嵌入层 → 行为胶囊聚合 → 动态路由 → 标签感知注意力
    */

    public class MINDModel extends Module {
    // 模型核心参数(与 MIND 论文一致)
    private final int embeddingDim; // 嵌入向量维度(64维)
    private final int maxSequenceLen; // 用户行为序列最大长度(10)
    private final int numInterest; // 兴趣胶囊数量 K(4个,论文默认)
    private final int routingIterations; // 动态路由迭代次数(3轮,论文默认)
    private final long userIdVocabSize; // 用户ID词汇表大小
    private final long itemIdVocabSize; // 商品ID词汇表大小
    private final long categoryVocabSize; // 类目词汇表大小

    // 1. 嵌入层(Embedding Layer):离散特征 → 稠密向量
    private Embedding userIdEmbedding; // 用户ID嵌入
    private Embedding itemIdEmbedding; // 商品ID嵌入
    private Embedding categoryEmbedding; // 类目嵌入

    // 2. 行为胶囊聚合层(Behavior Capsule Aggregation)
    private Linear behaviorCapsuleLinear; // 行为胶囊线性变换
    private float epsilon = 1e-8f; // Squash 激活函数防止除零

    // 3. 动态路由层(Dynamic Routing)
    private Linear routingLinear; // 路由线性变换(行为胶囊 → 兴趣胶囊)

    // 4. 标签感知注意力层(Label-Aware Attention)
    private Linear attentionLinear; // 注意力线性变换

    // 5. 输出层(推荐召回:兴趣向量 → 商品评分)
    private Linear outputLinear;

    /**
    * 模型初始化(严格遵循论文参数设置)
    * @param embeddingDim 嵌入维度(64)
    * @param maxSequenceLen 行为序列长度(10)
    * @param numInterest 兴趣胶囊数量(4)
    * @param routingIterations 路由迭代次数(3)
    * @param vocabSizes 词汇表大小(userIdVocabSize, itemIdVocabSize, categoryVocabSize)
    */

    public MINDModel(int embeddingDim, int maxSequenceLen, int numInterest, int routingIterations, long... vocabSizes) {
    super("MINDModel");
    this.embeddingDim = embeddingDim;
    this.maxSequenceLen = maxSequenceLen;
    this.numInterest = numInterest;
    this.routingIterations = routingIterations;
    this.userIdVocabSize = vocabSizes[0];
    this.itemIdVocabSize = vocabSizes[1];
    this.categoryVocabSize = vocabSizes[2];

    // 初始化各模块(严格遵循论文逻辑)
    initEmbeddingLayer();
    initBehaviorCapsuleLayer();
    initDynamicRoutingLayer();
    initLabelAwareAttentionLayer();
    initOutputLayer();

    // 注册模型参数(JavaCPP-PyTorch 必须注册,否则无法训练)
    register_module("userIdEmbedding", userIdEmbedding);
    register_module("itemIdEmbedding", itemIdEmbedding);
    register_module("categoryEmbedding", categoryEmbedding);
    register_module("behaviorCapsuleLinear", behaviorCapsuleLinear);
    register_module("routingLinear", routingLinear);
    register_module("attentionLinear", attentionLinear);
    register_module("outputLinear", outputLinear);
    }

    /**
    * 1. 嵌入层初始化(论文核心:离散特征转换为稠密向量)
    * 用户ID、商品ID、类目分别通过独立嵌入矩阵,最终拼接商品与类目嵌入作为商品特征
    */

    private void initEmbeddingLayer() {
    // 用户ID嵌入:(vocabSize, embeddingDim)
    userIdEmbedding = new Embedding((int) userIdVocabSize, embeddingDim);
    // 商品ID嵌入:(vocabSize, embeddingDim/2),与类目嵌入拼接后为 embeddingDim
    itemIdEmbedding = new Embedding((int) itemIdVocabSize, embeddingDim / 2);
    // 类目嵌入:(vocabSize, embeddingDim/2)
    categoryEmbedding = new Embedding((int) categoryVocabSize, embeddingDim / 2);

    // 初始化嵌入矩阵(论文建议:正态分布初始化)
    torch.nn_init_normal_(userIdEmbedding.weight(), 0.0f, 0.01f);
    torch.nn_init_normal_(itemIdEmbedding.weight(), 0.0f, 0.01f);
    torch.nn_init_normal_(categoryEmbedding.weight(), 0.0f, 0.01f);
    }

    /**
    * 2. 行为胶囊聚合层初始化(论文核心:将商品特征聚合为行为胶囊)
    * 输入:商品嵌入向量(embeddingDim)
    * 输出:行为胶囊(embeddingDim),通过 Squash 激活函数
    */

    private void initBehaviorCapsuleLayer() {
    // 线性变换:将商品嵌入(embeddingDim)转换为行为胶囊(embeddingDim)
    behaviorCapsuleLinear = new Linear(embeddingDim, embeddingDim);
    torch.nn_init_xavier_uniform_(behaviorCapsuleLinear.weight());
    torch.nn_init_zeros_(behaviorCapsuleLinear.bias());
    }

    /**
    * 3. 动态路由层初始化(论文核心:行为胶囊 → 兴趣胶囊)
    * 输入:行为胶囊(embeddingDim)
    * 输出:兴趣胶囊(numInterest, embeddingDim)
    */

    private void initDynamicRoutingLayer() {
    // 线性变换:行为胶囊 → 兴趣胶囊的预测向量
    routingLinear = new Linear(embeddingDim, embeddingDim * numInterest);
    torch.nn_init_xavier_uniform_(routingLinear.weight());
    torch.nn_init_zeros_(routingLinear.bias());
    }

    /**
    * 4. 标签感知注意力层初始化(论文核心:选择与目标商品最相关的兴趣向量)
    * 输入:兴趣胶囊(numInterest, embeddingDim)、目标商品嵌入(embeddingDim)
    * 输出:注意力权重(numInterest, 1)
    */

    private void initLabelAwareAttentionLayer() {
    // 线性变换:兴趣向量 + 目标商品向量 → 注意力得分
    attentionLinear = new Linear(2 * embeddingDim, 1);
    torch.nn_init_xavier_uniform_(attentionLinear.weight());
    torch.nn_init_zeros_(attentionLinear.bias());
    }

    /**
    * 5. 输出层初始化(论文核心:兴趣向量 → 商品推荐评分)
    * 输入:注意力加权后的兴趣向量(embeddingDim)
    * 输出:商品评分(itemIdVocabSize)
    */

    private void initOutputLayer() {
    outputLinear = new Linear(embeddingDim, (int) itemIdVocabSize);
    torch.nn_init_xavier_uniform_(outputLinear.weight());
    torch.nn_init_zeros_(outputLinear.bias());
    }

    /**
    * Squash 激活函数(胶囊网络核心,论文公式)
    * 作用:归一化胶囊向量,确保胶囊长度在 [0,1] 之间
    */

    private Tensor squash(Tensor x) {
    // 计算向量的 L2 范数(按最后一维)
    Tensor norm = torch.norm(x, 2, 1, true);
    // 论文公式:squash(x) = (norm² / (1 + norm²)) * (x / norm)
    Tensor normSquared = torch.pow(norm, 2);
    Tensor squashFactor = normSquared.div(normSquared.add(1.0f)).div(norm.add(epsilon));
    return x.mul(squashFactor);
    }

    /**
    * 动态路由算法(严格遵循论文迭代逻辑,3轮迭代)
    * @param behaviorCapsules 行为胶囊 (batchSize, maxSequenceLen, embeddingDim)
    * @return 兴趣胶囊 (batchSize, numInterest, embeddingDim)
    */

    private Tensor dynamicRouting(Tensor behaviorCapsules) {
    long batchSize = behaviorCapsules.size(0);

    // 1. 行为胶囊线性变换:(batchSize, maxSequenceLen, embeddingDim) → (batchSize, maxSequenceLen, numInterest*embeddingDim)
    Tensor predictedCapsules = routingLinear.forward(behaviorCapsules);
    // 重塑为:(batchSize, maxSequenceLen, numInterest, embeddingDim)
    predictedCapsules = predictedCapsules.view(new long[]{batchSize, maxSequenceLen, numInterest, embeddingDim});

    // 2. 初始化耦合系数 b_ij(batchSize, maxSequenceLen, numInterest),初始为0
    Tensor b = torch.zeros(new long[]{batchSize, maxSequenceLen, numInterest}, torch.get_default_dtype());
    if (torch.cuda_is_available()) {
    b = b.to(torch.device(torch.CUDA));
    }

    // 3. 迭代路由(论文默认3轮)
    for (int iter = 0; iter < routingIterations; iter++) {
    // 耦合系数 softmax 归一化:c_ij = softmax(b_ij)
    Tensor c = torch.softmax(b, 1); // (batchSize, maxSequenceLen, numInterest)

    // 计算兴趣胶囊 s_j = sum(c_ij * predictedCapsules_ij)
    Tensor cUnsqueezed = c.unsqueeze(1); // (batchSize, maxSequenceLen, numInterest, 1)
    Tensor s = torch.sum(predictedCapsules.mul(cUnsqueezed), 1); // (batchSize, numInterest, embeddingDim)

    // 兴趣胶囊 Squash 激活:v_j = squash(s_j)
    Tensor v = squash(s); // (batchSize, numInterest, embeddingDim)

    // 非最后一轮迭代:更新耦合系数 b_ij = b_ij + predictedCapsules_ij · v_j
    if (iter < routingIterations 1) {
    // 计算预测胶囊与兴趣胶囊的点积:(batchSize, maxSequenceLen, numInterest)
    Tensor dotProduct = torch.sum(predictedCapsules.mul(v.unsqueeze(1)), 1);
    b = b.add(dotProduct);
    }
    }

    return v; // 返回最终兴趣胶囊
    }

    /**
    * 标签感知注意力(论文核心:选择最相关的兴趣向量)
    * @param interestCapsules 兴趣胶囊 (batchSize, numInterest, embeddingDim)
    * @param targetItemEmbedding 目标商品嵌入 (batchSize, embeddingDim)
    * @return 注意力加权后的兴趣向量 (batchSize, embeddingDim)
    */

    private Tensor labelAwareAttention(Tensor interestCapsules, Tensor targetItemEmbedding) {
    long batchSize = interestCapsules.size(0);

    // 目标商品嵌入扩展维度:(batchSize, 1, embeddingDim)
    Tensor targetUnsqueezed = targetItemEmbedding.unsqueeze(1);
    // 拼接兴趣胶囊与目标商品嵌入:(batchSize, numInterest, 2*embeddingDim)
    Tensor concat = torch.cat(new Tensor[]{interestCapsules, targetUnsqueezed.expand(new long[]{batchSize, numInterest, embeddingDim})}, 1);

    // 计算注意力得分:(batchSize, numInterest, 1)
    Tensor attentionScores = attentionLinear.forward(concat);
    // 注意力权重归一化(softmax)
    Tensor attentionWeights = torch.softmax(attentionScores, 1);

    // 注意力加权:sum(weight_j * interest_j) → (batchSize, embeddingDim)
    Tensor weightedInterest = torch.sum(interestCapsules.mul(attentionWeights), 1);

    return weightedInterest;
    }

    /**
    * 模型前向传播(严格遵循论文流程:嵌入层 → 行为胶囊 → 动态路由 → 注意力 → 输出)
    */

    @Override
    public IValue forward(IValue input) {
    // 解析输入特征(来自 Streaming Dataset)
    java.util.Map<String, IValue> inputMap = input.toMap();
    Tensor userId = inputMap.get("userId").toTensor(); // (batchSize, 1)
    Tensor itemSeq = inputMap.get("itemSeq").toTensor(); // (batchSize, maxSequenceLen)
    Tensor categorySeq = inputMap.get("categorySeq").toTensor(); // (batchSize, maxSequenceLen)
    Tensor targetItem = inputMap.get("targetItem").toTensor(); // (batchSize, 1)
    Tensor targetPrice = inputMap.get("targetPrice").toTensor(); // (batchSize, 1)

    // 1. 嵌入层:离散特征 → 稠密向量
    // 用户嵌入:(batchSize, 1, embeddingDim)
    Tensor userEmbedding = userIdEmbedding.forward(userId).unsqueeze(1);
    // 商品嵌入:(batchSize, maxSequenceLen, embeddingDim/2)
    Tensor itemEmbedding = itemIdEmbedding.forward(itemSeq);
    // 类目嵌入:(batchSize, maxSequenceLen, embeddingDim/2)
    Tensor categoryEmbedding = this.categoryEmbedding.forward(categorySeq);
    // 拼接商品与类目嵌入:(batchSize, maxSequenceLen, embeddingDim)
    Tensor itemCategoryEmbedding = torch.cat(new Tensor[]{itemEmbedding, categoryEmbedding}, 1);

    // 2. 行为胶囊聚合:商品嵌入 → 行为胶囊(Squash 激活)
    Tensor behaviorCapsules = behaviorCapsuleLinear.forward(itemCategoryEmbedding); // (batchSize, maxSequenceLen, embeddingDim)
    behaviorCapsules = squash(behaviorCapsules); // 激活后,胶囊长度归一化

    // 3. 动态路由:行为胶囊 → 兴趣胶囊(K个)
    Tensor interestCapsules = dynamicRouting(behaviorCapsules); // (batchSize, numInterest, embeddingDim)

    // 4. 目标商品嵌入(用于注意力和损失计算)
    Tensor targetItemEmbedding = itemIdEmbedding.forward(targetItem).unsqueeze(1); // (batchSize, 1, embeddingDim/2)
    Tensor targetCategory = categorySeq.select(1, maxSequenceLen 1).unsqueeze(1); // 目标商品对应的类目
    Tensor targetCategoryEmbedding = this.categoryEmbedding.forward(targetCategory); // (batchSize, 1, embeddingDim/2)
    Tensor targetEmbedding = torch.cat(new Tensor[]{targetItemEmbedding, targetCategoryEmbedding}, 1).squeeze(1); // (batchSize, embeddingDim)

    // 5. 标签感知注意力:选择最相关的兴趣向量
    Tensor weightedInterest = labelAwareAttention(interestCapsules, targetEmbedding); // (batchSize, embeddingDim)

    // 6. 输出层:兴趣向量 → 商品推荐评分
    Tensor outputScores = outputLinear.forward(weightedInterest); // (batchSize, itemIdVocabSize)

    // 返回输出(推荐评分、兴趣胶囊、目标商品嵌入,

    赞(0)
    未经允许不得转载:171主机测评 » 【Java PyTorch深度学习】PyTorch On Java 进阶课程 Flink特征工程 与PyTorch实时特征工程与流式推荐系统[PyTorch Java 硕士研一课程]
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址