欢迎光临
我们一直在努力

Spark 从筑基到化神

master阅读(28)

Spark 从零到进阶:数据开发工程师的系统学习指南

写在前面:如果你是一名刚接触 Spark 的数据开发工程师,面对网上零散的教程和概念感到无从下手,那么这篇文章就是为你准备的。我们将从"Spark 是什么"出发,一路走到性能调优与生产踩坑,配合大量代码示例和架构图,帮你建立完整的知识体系。本文基于 Spark 3.5.x 版本编写,并会提及 Spark 4.0 的预览方向。


目录

  • Spark 概述
  • Spark 架构与运行原理
  • 环境搭建
  • RDD 编程
  • Spark SQL
  • Spark Streaming vs Structured Streaming
  • Spark 数据湖与 Lakehouse
  • Spark 性能调优
  • Spark 常见问题与踩坑
  • 端到端实战项目
  • Spark 3.x 新特性
  • 学习路线与实战建议

  • 1. Spark 概述

    1.1 什么是 Spark

    Apache Spark 是一个开源的统一分析引擎,专为大规模数据处理而设计。它最初于 2009 年在加州大学伯克利分校 AMPLab 诞生,2010 年开源,2014 年成为 Apache 顶级项目。Spark 提供了 SQL、流计算、机器学习和图计算等一整套大数据处理能力,可以运行在 Hadoop YARN、Kubernetes、Standalone 等多种集群管理器上。

    1.2 发展历史

    时间里程碑
    2009 Spark 在 UC Berkeley AMPLab 诞生
    2010 BSD 许可开源
    2013 捐赠给 Apache 软件基金会
    2014 Apache 顶级项目;Spark 1.0 发布
    2016 Spark 2.0:DataFrame/Dataset 成为主流 API,Structured Streaming
    2020 Spark 3.0:AQE、Dynamic Partition Pruning
    2023 Spark 3.4 / 3.5:Spark Connect 稳定、Pandas API 完善
    2024+ Spark 4.0 预览:更好的 ANSI SQL 兼容、性能持续优化

    1.3 Spark vs Hadoop MapReduce

    很多初学者会问"Spark 会替代 Hadoop 吗?"准确地说,Spark 替代的是 MapReduce 计算引擎,而 HDFS、YARN 依然是重要的存储和资源管理组件。

    对比维度MapReduceSpark
    中间结果 落盘(HDFS) 优先内存,可溢写磁盘
    编程模型 Map + Reduce 两阶段 DAG(有向无环图),多阶段
    延迟 高(批处理) 低(批流统一)
    迭代计算 每次迭代读写 HDFS 内存缓存,迭代效率高 10-100x
    生态 仅 MapReduce SQL/Streaming/ML/Graph 一体化
    容错 基于磁盘复制 RDD 血缘(Lineage)重算

    1.4 核心特性

    • 速度快:DAG 执行引擎 + 内存计算,比 MapReduce 快 10-100 倍
    • 易用性:支持 Scala、Python、Java、R 四种语言 API,代码简洁
    • 通用性:一站式覆盖批处理、SQL、流计算、ML、图计算
    • 兼容性:可读写 HDFS、HBase、Cassandra、S3 等数据源,运行在 YARN/K8s/Standalone/Mesos 上

    1.5 生态组件全景

    在这里插入图片描述

    • Spark Core:底层引擎,提供 RDD 抽象、任务调度、内存管理、故障恢复
    • Spark SQL:结构化数据处理,DataFrame/Dataset API,兼容 Hive
    • Structured Streaming:基于 Spark SQL 的流计算引擎(原 Spark Streaming 已进入维护模式)
    • MLlib:分布式机器学习库
    • GraphX:图计算引擎

    2. Spark 架构与运行原理

    2.1 核心组件

    在这里插入图片描述

    组件职责
    Driver 运行 main() 函数,创建 SparkContext,将用户程序转化为 DAG,调度 Task
    Executor 在 Worker 节点上启动的 JVM 进程,执行 Task 并缓存数据
    Cluster Manager 集群资源管理器,负责分配 CPU/内存(YARN、K8s、Standalone)
    Worker 集群中可运行 Application 代码的节点

    2.2 部署模式对比

    模式特点适用场景
    Local 单机运行,多线程模拟分布式 开发调试、学习
    Standalone Spark 自带集群管理器 小规模独立集群
    On YARN 复用 Hadoop YARN 资源管理 企业最常用,与 Hadoop 生态融合
    On Kubernetes 容器化部署,弹性扩缩容 云原生环境,近年增长迅速

    YARN 模式又分为 yarn-cluster(Driver 运行在 AM 中,适合生产)和 yarn-client(Driver 在客户端,适合交互调试)。

    2.3 Job、Stage、Task 的划分

    当一个 Action 算子被触发时,Spark 会提交一个 Job。DAGScheduler 根据 RDD 的依赖关系将 Job 划分为多个 Stage,每个 Stage 内创建一组 Task,由 TaskScheduler 分发到 Executor 执行。

    在这里插入图片描述

    划分规则:

    • 遇到宽依赖(Shuffle)就切分 Stage
    • 一个 Stage 内的所有 Task 执行完全相同的代码,只是处理的数据分区不同
    • Task 数量 = Stage 最后一个 RDD 的分区数

    2.4 宽窄依赖与 Shuffle

    类型说明典型算子
    窄依赖 父 RDD 每个分区最多被子 RDD 一个分区使用 map、filter、union、mapPartitions
    宽依赖 父 RDD 每个分区被子 RDD 多个分区使用,需要 Shuffle reduceByKey、groupByKey、join、distinct、repartition

    在这里插入图片描述

    Shuffle 过程:Map 端将数据按 Key 写入磁盘缓冲区,经过排序、合并、溢写后产生 Shuffle 文件;Reduce 端通过 BlockManager 拉取属于自己的数据。Shuffle 涉及磁盘 I/O、网络 I/O 和数据序列化,是 Spark 性能瓶颈的主要来源。

    2.5 统一内存模型(Spark 1.6+)

    Spark Executor 的内存分为以下区域: 在这里插入图片描述

    • Storage Memory:缓存 RDD、Broadcast 变量
    • Execution Memory:Shuffle、Join、Sort 等执行过程中的临时数据
    • 两者之间可以动态借用:Execution 可以借 Storage 的空闲内存,Storage 也可以借 Execution 的,但 Execution 优先级更高,Storage 被借走的部分在需要时会被淘汰

    3. 环境搭建

    3.1 Local 模式(最快上手)

    # 下载(以 3.5.1 为例)
    wget https://archive.apache.org/dist/spark/spark-3.5.1/spark-3.5.1-bin-hadoop3.tgz
    tar -xzf spark-3.5.1-bin-hadoop3.tgz
    cd spark-3.5.1-bin-hadoop3

    # 启动 Spark Shell(Scala)
    ./bin/spark-shell –master local[*]

    # 启动 PySpark
    ./bin/pyspark –master local[*]

    local[*] 表示使用所有 CPU 核心。也可以用 local[2] 指定 2 个核心。

    3.2 Standalone 模式

    # 1. 配置 conf/spark-env.sh
    cp conf/spark-env.sh.template conf/spark-env.sh
    cat >> conf/spark-env.sh << 'EOF'
    export SPARK_MASTER_HOST=master
    export SPARK_WORKER_CORES=4
    export SPARK_WORKER_MEMORY=8g
    EOF

    # 2. 配置 workers
    echo "worker1" > conf/workers
    echo "worker2" >> conf/workers

    # 3. 启动集群
    ./sbin/start-all.sh

    # 4. 提交应用
    ./bin/spark-submit \\
    –master spark://master:7077 \\
    –class com.example.MyApp \\
    myapp.jar

    3.3 On YARN 模式

    确保 HADOOP_CONF_DIR 或 YARN_CONF_DIR 指向 Hadoop 配置目录:

    export HADOOP_CONF_DIR=/etc/hadoop/conf

    # 提交到 YARN(cluster 模式)
    ./bin/spark-submit \\
    –master yarn \\
    –deploy-mode cluster \\
    –driver-memory 2g \\
    –executor-memory 4g \\
    –executor-cores 2 \\
    –num-executors 10 \\
    myapp.py

    3.4 Spark Shell 与 spark-submit

    工具用途
    spark-shell Scala 交互式 REPL,适合探索和调试
    pyspark Python 交互式 REPL
    spark-sql SQL 交互式命令行
    spark-submit 提交应用到集群(生产环境)

    spark-submit 常用参数:

    参数说明示例
    –master 集群管理器 URL yarn / spark://host:7077 / local[*]
    –deploy-mode Driver 运行位置 cluster / client
    –name 应用名称 MySparkApp
    –class Java/Scala 主类 com.example.MyApp
    –jars 额外依赖 JAR –jars lib/mysql-connector.jar
    –packages Maven 依赖自动下载 –packages org.apache.spark:spark-sql-kafka-0-10_2.12:3.5.1
    –files 分发文件到工作目录 –files config.properties
    –conf Spark 配置项 –conf spark.sql.shuffle.partitions=200
    –driver-memory Driver 内存 4g
    –executor-memory 每个 Executor 内存 8g
    –executor-cores 每个 Executor CPU 核数 4
    –num-executors Executor 数量(YARN) 20
    –total-executor-cores 总核数(Standalone) 100
    –archives 分发归档文件并解压 –archives env.tar.gz#env
    –py-files 额外 Python 文件/压缩包 –py-files utils.zip

    3.5 PySpark 环境配置

    # 方式一:pip 安装
    pip install pyspark==3.5.1

    # 方式二:使用 Spark 自带的 PySpark
    export SPARK_HOME=/opt/spark
    export PYTHONPATH=$SPARK_HOME/python:$SPARK_HOME/python/lib/py4j-0.10.9.7-src.zip:$PYTHONPATH
    export PATH=$SPARK_HOME/bin:$PATH

    # 验证
    pyspark –version

    在 Jupyter Notebook 中使用:

    pip install jupyter
    export PYSPARK_DRIVER_PYTHON=jupyter
    export PYSPARK_DRIVER_PYTHON_OPTS='notebook'
    pyspark –master local[*]

    3.6 Spark on Kubernetes

    Kubernetes(K8s)已成为大数据上云的主流部署方式。Spark on K8s 从 Spark 2.3 开始支持,Spark 3.1 进入 GA,目前在云原生场景下增长迅速。

    部署模式原理:

    Spark on K8s 将 Driver 和 Executor 直接作为 Pod 运行在 K8s 集群中:

    在这里插入图片描述

    • Driver Pod 启动后,通过 K8s API 动态申请 Executor Pod
    • Executor Pod 运行结束后自动销毁,资源归还 K8s
    • 支持动态资源分配(Dynamic Allocation),按需扩缩容

    与 YARN 模式对比:

    对比维度Spark on YARNSpark on K8s
    部署方式 依赖 Hadoop YARN 集群 容器化部署,镜像打包所有依赖
    弹性扩缩容 依赖 YARN 队列配置,扩容较慢 秒级 Pod 创建,弹性能力强
    资源隔离 队列级隔离,Container 共享 OS Pod 级隔离,容器边界清晰
    依赖管理 通过 –jars/–packages 分发 打包进 Docker 镜像,版本一致性好
    生态融合 与 HDFS/Hive/HBase 深度集成 云原生生态(Prometheus/Istio/Envoy)
    运维复杂度 Hadoop 运维体系成熟 需 K8s 运维能力,学习曲线稍陡
    多租户 YARN 队列 + Linux 用户 Namespace + RBAC,隔离更彻底
    适用场景 传统 Hadoop 集群、离线数仓 云原生、混合云、弹性计算、流批一体

    spark-submit on K8s 命令示例:

    ./bin/spark-submit \\
    –master k8s://https://<k8s-apiserver-host>:<port> \\
    –deploy-mode cluster \\
    –name spark-pi \\
    –class org.apache.spark.examples.SparkPi \\
    –conf spark.executor.instances=5 \\
    –conf spark.executor.memory=4g \\
    –conf spark.executor.cores=2 \\
    –conf spark.kubernetes.container.image=registry.example.com/spark:3.5.1 \\
    –conf spark.kubernetes.namespace=spark-jobs \\
    –conf spark.kubernetes.authenticate.driver.serviceAccountName=spark \\
    –conf spark.kubernetes.driver.pod.name=spark-pi-driver \\
    local:///opt/spark/examples/jars/spark-examples_2.12-3.5.1.jar

    Docker 镜像构建要点:

    # 基础镜像
    FROM openjdk:11-jre-slim

    # 安装 Spark
    ARG SPARK_VERSION=3.5.1
    ARG HADOOP_VERSION=3
    RUN apt-get update && apt-get install -y curl tini && \\
    curl -sL https://archive.apache.org/dist/spark/spark-${SPARK_VERSION}/spark-${SPARK_VERSION}-bin-hadoop${HADOOP_VERSION}.tgz | tar xz -C /opt && \\
    ln -s /opt/spark-${SPARK_VERSION}-bin-hadoop${HADOOP_VERSION} /opt/spark && \\
    apt-get clean

    # 安装 Python(PySpark 场景)
    RUN apt-get install -y python3 python3-pip && \\
    pip3 install pyspark==${SPARK_VERSION} pandas pyarrow

    # 拷贝自定义 JAR 依赖(如 JDBC 驱动、数据湖 SDK)
    COPY jars/mysql-connector-j-8.0.33.jar /opt/spark/jars/
    COPY jars/iceberg-spark-runtime-3.5_2.12-1.5.2.jar /opt/spark/jars/

    # 拷贝应用 JAR
    COPY app/your-app.jar /opt/spark/user-jars/

    ENTRYPOINT ["/usr/bin/tini", "–"]

    构建并推送镜像:

    docker build -t registry.example.com/spark:3.5.1-custom .
    docker push registry.example.com/spark:3.5.1-custom

    K8s 特有配置参数:

    参数说明示例
    spark.kubernetes.container.image Docker 镜像地址 registry/spark:3.5.1
    spark.kubernetes.namespace K8s 命名空间 spark-jobs
    spark.kubernetes.authenticate.driver.serviceAccountName Driver 使用的 ServiceAccount spark
    spark.kubernetes.driver.pod.name Driver Pod 名称 spark-app-driver
    spark.kubernetes.executor.podNamePrefix Executor Pod 名称前缀 spark-app-exec
    spark.kubernetes.driver.limit.cores Driver CPU 限制 2
    spark.kubernetes.executor.limit.cores Executor CPU 限制 4
    spark.kubernetes.driver.request.cores Driver CPU 请求 1
    spark.kubernetes.executor.request.cores Executor CPU 请求 2
    spark.kubernetes.memoryOverheadFactor 堆外内存比例 0.2
    spark.kubernetes.allocation.batch.size 每批申请 Pod 数 10
    spark.kubernetes.executor.deleteOnTermination 完成后删除 Executor Pod true
    spark.kubernetes.driver.podTemplateFile Driver Pod 模板文件 driver-pod.yaml
    spark.kubernetes.executor.podTemplateFile Executor Pod 模板文件 executor-pod.yaml
    spark.kubernetes.file.upload.path 应用依赖上传路径 s3://spark-uploads/
    spark.kubernetes.authenticate.caCertFile CA 证书路径 /var/run/secrets/…

    适用场景:

    • 云原生架构:与公司现有 K8s 平台统一运维,避免 Hadoop/YARN 独立集群
    • 弹性扩缩容:利用 K8s Cluster Autoscaler,根据负载自动增减节点,降低成本
    • 混合云/多云:同一镜像在不同云厂商 K8s 集群运行,避免厂商锁定
    • 流批一体:Streaming 长任务与离线批任务共享 K8s 资源池,通过 Namespace/Quota 隔离
    • 依赖隔离:不同应用使用不同镜像,彻底解决"一个集群多版本依赖冲突"问题

    💡 生产建议:K8s 上建议启用 External Shuffle Service(或使用 Spark 3.0+ 的 Shuffle Data on PVC 方案),开启 Dynamic Allocation,并配置 Pod 模板以挂载共享存储(如 HDFS/S3/OSS)。


    4. RDD 编程

    虽然现在 Spark SQL/DataFrame 是主流 API,但理解 RDD(Resilient Distributed Dataset,弹性分布式数据集)是掌握 Spark 底层原理的关键。

    4.1 RDD 五大特性

  • 分区列表(Partitions):数据被切分为多个分区,每个分区在一个节点上计算
  • 计算函数(Compute):每个分区都有一个计算函数来生成数据
  • 依赖关系(Dependencies):RDD 之间有血缘关系,用于故障恢复
  • 分区器(Partitioner):KV 类型 RDD 可选(Hash/Range),决定数据分布
  • 优先位置(Preferred Locations):每个分区的计算优先调度到数据所在节点
  • 4.2 创建 RDD

    from pyspark.sql import SparkSession

    spark = SparkSession.builder.appName("RDDDemo").master("local[*]").getOrCreate()
    sc = spark.sparkContext

    # 方式一:从集合创建
    rdd = sc.parallelize([1, 2, 3, 4, 5], numSlices=3)

    # 方式二:从外部存储创建
    rdd = sc.textFile("hdfs:///data/input.txt")
    rdd = sc.wholeTextFiles("hdfs:///data/logs/") # 读取目录下所有文件

    4.3 Transformation vs Action

    Transformation 是懒执行的,只记录转换逻辑,不立即计算;Action 触发真正的计算。

    类型常见算子
    Transformation map、filter、flatMap、mapPartitions、sample、union、intersection、distinct、groupByKey、reduceByKey、aggregateByKey、sortByKey、join、cogroup、cartesian、coalesce、repartition
    Action collect、count、take、first、takeOrdered、reduce、fold、aggregate、foreach、saveAsTextFile、countByKey、collectAsMap

    # Transformation 示例
    nums = sc.parallelize([1, 2, 3, 4, 5, 6])

    # map:一对一转换
    squares = nums.map(lambda x: x * x) # [1, 4, 9, 16, 25, 36]

    # filter:过滤
    evens = nums.filter(lambda x: x % 2 == 0) # [2, 4, 6]

    # flatMap:一对多展平
    lines = sc.parallelize(["hello world", "hello spark"])
    words = lines.flatMap(lambda line: line.split(" "))
    # ["hello", "world", "hello", "spark"]

    # reduceByKey:按 Key 聚合(Map 端预聚合,性能优于 groupByKey)
    pairs = words.map(lambda w: (w, 1))
    counts = pairs.reduceByKey(lambda a, b: a + b)
    # [("hello", 2), ("world", 1), ("spark", 1)]

    # join
    rdd1 = sc.parallelize([("a", 1), ("b", 2)])
    rdd2 = sc.parallelize([("a", "x"), ("b", "y"), ("b", "z")])
    rdd1.join(rdd2).collect()
    # [("a", (1, "x")), ("b", (2, "y")), ("b", (2, "z"))]

    # sortByKey
    sorted_counts = counts.sortByKey(ascending=True)

    # Action 示例
    print(counts.collect()) # 返回所有结果到 Driver
    print(counts.count()) # 元素数量
    print(counts.take(2)) # 取前 2 个
    print(counts.first()) # 第一个
    total = nums.reduce(lambda a, b: a + b) # 聚合
    counts.saveAsTextFile("hdfs:///output/wordcount")

    4.4 reduceByKey vs groupByKey

    这是面试高频考点:

    对比reduceByKeygroupByKey
    Map 端预聚合 ✅ 有(Combiner) ❌ 无
    Shuffle 数据量
    性能
    适用场景 聚合类操作 需要遍历所有值的场景

    优先使用 reduceByKey / aggregateByKey / foldByKey,它们会在 Map 端做预聚合。

    4.5 持久化:cache / persist / checkpoint

    from pyspark import StorageLevel

    # cache() = persist(MEMORY_ONLY)
    rdd.cache()

    # persist 支持多种存储级别
    rdd.persist(StorageLevel.MEMORY_AND_DISK) # 内存不够溢写磁盘
    rdd.persist(StorageLevel.MEMORY_ONLY_SER) # 序列化后存内存,节省空间
    rdd.persist(StorageLevel.MEMORY_AND_DISK_SER_2) # 序列化 + 磁盘 + 2 副本

    # 解除缓存
    rdd.unpersist()

    # checkpoint:将 RDD 写入 HDFS 等可靠存储,切断血缘
    sc.setCheckpointDir("hdfs:///checkpoint")
    rdd.checkpoint()

    方式存储位置血缘适用场景
    cache 内存 保留 频繁重用的小数据
    persist 内存/磁盘可选 保留 根据数据量选择
    checkpoint HDFS 切断 长血缘链、迭代计算

    4.6 共享变量

    默认情况下,Spark 会把函数中用到的变量拷贝到每个 Task,Task 对变量的修改不会回传 Driver。共享变量提供了两种特殊机制:

    Broadcast Variable(广播变量):将只读变量缓存到每个 Executor,而不是每个 Task 一份,大幅减少数据传输。

    # 大数据量的查找表
    lookup_table = {"US": "United States", "CN": "China", "JP": "Japan"}
    bc_var = sc.broadcast(lookup_table)

    rdd = sc.parallelize(["US", "CN", "JP"])
    result = rdd.map(lambda code: bc_var.value.get(code, code)).collect()
    # ["United States", "China", "Japan"]

    Accumulator(累加器):分布式计数器,只能在 Executor 端 add,在 Driver 端读取 value。

    accum = sc.accumulator(0)

    def process_line(line):
    global accum
    if "ERROR" in line:
    accum.add(1)
    return line

    sc.textFile("hdfs:///logs/app.log").map(process_line).count()
    print(f"Error count: {accum.value}")

    ⚠️ 注意:Accumulator 在 Transformation 中可能因 Task 重试导致重复计数,建议在 Action(如 foreach)中使用,或使用 Named Accumulator 并在 Driver 端确认只读取一次。


    5. Spark SQL

    Spark SQL 是当前 Spark 最核心、使用最广泛的模块。DataFrame/Dataset API 比 RDD 更高级、更优化。

    5.1 DataFrame 与 Dataset

    特性RDDDataFrameDataset
    数据模型 无结构 带 Schema 的行 带 Schema 的强类型对象
    类型安全 是(编译时) 否(运行时) 是(编译时)
    优化 Catalyst Catalyst
    语言 Scala/Java/Python/R 全部 主要 Scala/Java
    Tungsten

    在 PySpark 中,DataFrame 是主要 API(Python 没有 Dataset 的编译时类型检查)。

    5.2 创建 DataFrame

    from pyspark.sql import SparkSession
    from pyspark.sql.types import *
    from pyspark.sql.functions import col, sum as _sum, avg, count

    spark = SparkSession.builder \\
    .appName("SparkSQLDemo") \\
    .master("local[*]") \\
    .config("spark.sql.shuffle.partitions", "4") \\
    .getOrCreate()

    # 1. 从 RDD 转换
    rdd = sc.parallelize([(1, "Alice", 25), (2, "Bob", 30)])
    df = spark.createDataFrame(rdd, ["id", "name", "age"])

    # 2. 显式指定 Schema
    schema = StructType([
    StructField("id", IntegerType(), False),
    StructField("name", StringType(), True),
    StructField("age", IntegerType(), True),
    ])
    df = spark.createDataFrame(rdd, schema)

    # 3. 读取文件
    df = spark.read.csv("data/users.csv", header=True, inferSchema=True)
    df = spark.read.json("data/users.json")
    df = spark.read.parquet("data/users.parquet")
    df = spark.read.orc("data/users.orc")

    # 4. 读取 Hive 表
    df = spark.sql("SELECT * FROM default.users")

    # 5. JDBC
    df = spark.read.format("jdbc") \\
    .option("url", "jdbc:mysql://localhost:3306/mydb") \\
    .option("dbtable", "users") \\
    .option("user", "root") \\
    .option("password", "xxx") \\
    .load()

    # 6. toDF
    df = [(1, "Alice"), (2, "Bob")].toDF(["id", "name"]) # Scala 风格
    # PySpark 中:
    df = spark.createDataFrame([(1, "Alice"), (2, "Bob")], ["id", "name"])

    5.3 常用 Transformation

    # 选择列
    df.select("name", "age").show()
    df.select(col("name"), col("age") + 1).show()

    # 过滤
    df.filter(col("age") > 25).show()
    df.where("age > 25").show()

    # 分组聚合
    df.groupBy("department").agg(
    count("*").alias("emp_count"),
    _sum("salary").alias("total_salary"),
    avg("salary").alias("avg_salary")
    ).show()

    # 排序
    df.orderBy(col("age").desc()).show()

    # 去重
    df.select("department").distinct().show()

    # 列重命名
    df.withColumnRenamed("name", "employee_name")

    # 新增列
    df.withColumn("age_next_year", col("age") + 1)

    # Join
    df1.join(df2, on="id", how="inner") # inner/left/right/outer/semi/anti
    df1.join(df2, df1.id == df2.emp_id, "left_outer")

    # 窗口函数
    from pyspark.sql.window import Window
    from pyspark.sql.functions import row_number, rank, dense_rank

    window_spec = Window.partitionBy("department").orderBy(col("salary").desc())
    df.withColumn("rank", row_number().over(window_spec)).show()

    # 写数据
    df.write.mode("overwrite").parquet("hdfs:///output/users")
    df.write.partitionBy("department").parquet("hdfs:///output/users_by_dept")

    5.4 UDF / UDAF

    from pyspark.sql.functions import udf
    from pyspark.sql.types import StringType, IntegerType

    # 标量 UDF
    def age_group(age):
    if age < 25:
    return "Young"
    elif age < 40:
    return "Mid"
    else:
    return "Senior"

    age_group_udf = udf(age_group, StringType())
    df.withColumn("age_group", age_group_udf(col("age"))).show()

    # 更高效的方式:pandas UDF(向量化执行,基于 Apache Arrow)
    import pandas as pd
    from pyspark.sql.functions import pandas_udf

    @pandas_udf(StringType())
    def age_group_pandas(age: pd.Series) > pd.Series:
    return age.apply(age_group)

    df.withColumn("age_group", age_group_pandas(col("age"))).show()

    💡 性能提示:普通 Python UDF 每行一次 Python/JVM 序列化,性能差;Pandas UDF 以列式批量传输,性能提升 10-100 倍。Spark 3.5 还引入了 Python UDF profiling 和更高效的 Arrow 批处理。

    5.5 开窗函数

    from pyspark.sql.window import Window
    from pyspark.sql.functions import row_number, rank, dense_rank, lag, lead

    window_spec = Window.partitionBy("department").orderBy(col("salary").desc())

    # 排名
    df.select(
    "name", "department", "salary",
    row_number().over(window_spec).alias("rn"),
    rank().over(window_spec).alias("rank"),
    dense_rank().over(window_spec).alias("dr"),
    lag("salary", 1).over(window_spec).alias("prev_salary"),
    lead("salary", 1).over(window_spec).alias("next_salary"),
    ).show()

    5.6 执行计划与 Catalyst 优化器

    df.explain() # 简明物理计划
    df.explain(True) # 含解析/分析/优化/物理计划
    df.explain("formatted") # 格式化输出(Spark 3.0+)

    Catalyst 优化器是 Spark SQL 的核心,它将用户的 DataFrame/SQL 代码经过以下阶段优化:

    在这里插入图片描述

    常见优化规则:

    • 谓词下推(Predicate Pushdown):将过滤条件尽早下推到数据源
    • 列裁剪(Column Pruning):只读取需要的列
    • 常量折叠(Constant Folding):预计算常量表达式
    • 布尔简化:简化 AND/OR 逻辑

    5.7 Tungsten 优化

    Tungsten 是 Spark 的执行引擎优化项目,目标是最大化 CPU 效率和内存效率:

  • 内存管理:使用堆外内存(off-heap),避免 JVM GC 开销
  • 二进制处理:数据以二进制格式存储,避免 Java 对象的序列化/反序列化
  • Whole-Stage CodeGen:将整个 Stage 的多个算子融合为一个 Java 函数,消除虚函数调用
  • 向量化读取:Parquet/ORC 列式批量读取
  • 5.8 AQE(Adaptive Query Execution,自适应查询执行)

    Spark 3.0 引入的重大特性,在 Shuffle Map 阶段完成后根据运行时统计信息动态调整执行计划:

    spark.conf.set("spark.sql.adaptive.enabled", "true")
    spark.conf.set("spark.sql.adaptive.coalescePartitions.enabled", "true")
    spark.conf.set("spark.sql.adaptive.skewJoin.enabled", "true")
    spark.conf.set("spark.sql.adaptive.localShuffleReader.enabled", "true")

    三大核心能力:

    功能说明
    动态合并 Shuffle 分区 自动将过小的分区合并,减少 Task 数量
    动态处理数据倾斜 自动检测倾斜分区并拆分 Join
    动态切换 Join 策略 运行时发现小表可广播时,自动从 Sort-Merge Join 切换为 Broadcast Join

    6. Spark Streaming vs Structured Streaming

    6.1 Spark Streaming(DStream)

    Spark Streaming 是早期的流计算方案,基于微批处理(Micro-Batch),将实时数据流按时间间隔切分为小的 RDD 批处理。 在这里插入图片描述

    from pyspark.streaming import StreamingContext

    ssc = StreamingContext(sc, batchDuration=5) # 5 秒一个批次
    lines = ssc.socketTextStream("localhost", 9999)
    counts = lines.flatMap(lambda line: line.split(" ")) \\
    .map(lambda w: (w, 1)) \\
    .reduceByKey(lambda a, b: a + b)
    counts.pprint()
    ssc.start()
    ssc.awaitTermination()

    ⚠️ Spark Streaming(DStream)从 Spark 3.4 起已标记为维护模式,新项目应使用 Structured Streaming。

    6.2 Structured Streaming

    Structured Streaming 是基于 Spark SQL 引擎构建的端到端流计算,将实时数据流视为一张持续追加的"无界表"。

    from pyspark.sql.functions import *

    # 从 Kafka 读取
    df = spark.readStream \\
    .format("kafka") \\
    .option("kafka.bootstrap.servers", "localhost:9092") \\
    .option("subscribe", "orders") \\
    .load()

    # 解析 JSON
    parsed = df.selectExpr("CAST(value AS STRING) as json") \\
    .select(from_json("json", schema).alias("data")) \\
    .select("data.*")

    # 聚合
    agg = parsed.groupBy("product_id").agg(
    count("*").alias("order_count"),
    sum("amount").alias("total_amount")
    )

    # 输出到控制台
    query = agg.writeStream \\
    .outputMode("complete") \\
    .format("console") \\
    .trigger(processingTime="10 seconds") \\
    .start()

    query.awaitTermination()

    6.3 核心概念

    Event Time(事件时间):数据产生时自带的时间戳,而非 Spark 处理时间(Processing Time)。

    Watermark(水位线):定义了系统等待迟到数据的时间阈值,超过阈值的迟到数据将被丢弃。

    # 设置 10 分钟 watermark
    df_with_watermark = parsed \\
    .withWatermark("event_time", "10 minutes") \\
    .groupBy(
    window("event_time", "5 minutes"), # 5 分钟滚动窗口
    "product_id"
    ).agg(count("*").alias("cnt"))

    Window(窗口):

    窗口类型说明
    Tumbling Window 固定大小,不重叠(如每 5 分钟统计)
    Sliding Window 固定大小,可重叠(如每 1 分钟统计最近 5 分钟)
    Session Window 基于活动间隔动态划分(Spark 3.2+)

    Output Mode(输出模式):

    模式说明适用聚合
    Append 仅输出新行 无聚合
    Complete 输出全量结果 有聚合
    Update 仅输出更新行 有聚合(Spark 2.1+)

    6.4 Source 与 Sink

    Source说明
    Kafka 最常用,支持从 Kafka 读取消息
    File 监听目录中新文件
    Socket 测试用,从 TCP Socket 读取
    Rate 测试用,每秒生成指定行数
    Sink说明
    Kafka 写入 Kafka Topic
    File 写入文件(Parquet/JSON/CSV)
    Console 控制台(调试用)
    Foreach/ForeachBatch 自定义写入逻辑
    Memory 存储为内存表(调试用)

    6.5 Exactly-Once 语义

    Structured Streaming 通过以下机制保证精确一次(Exactly-Once):

  • 可重放的 Source:如 Kafka,记录 offset 可重新读取
  • 幂等的 Sink:或使用事务写入(如 Kafka 事务、文件原子写入)
  • Checkpoint + WAL:将 offset 和聚合状态持久化到可靠存储
  • query = agg.writeStream \\
    .format("kafka") \\
    .option("kafka.bootstrap.servers", "localhost:9092") \\
    .option("topic", "order_summary") \\
    .option("checkpointLocation", "hdfs:///checkpoint/orders") \\
    .outputMode("complete") \\
    .start()

    注意:foreachBatch 本身不保证 Exactly-Once,需要 Sink 端实现幂等或使用事务。

    7. Spark 数据湖与 Lakehouse

    随着大数据从"数据仓库"向"湖仓一体"演进,数据湖格式已成为 Spark 生态不可或缺的一环。本章介绍三大开源数据湖格式及 Paimon,并给出 Spark 实战代码。

    7.1 数据湖概念与三剑客对比

    数据湖(Data Lake):以原始格式存储海量结构化、半结构化、非结构化数据的集中式存储,支持多种计算引擎直接读写。传统数据湖缺乏事务支持,容易产生脏数据。

    Lakehouse(湖仓一体):在数据湖存储之上增加事务层、Schema 管理、ACID 语义,兼具数据湖的灵活性和数据仓库的管理能力。

    特性Delta LakeApache IcebergApache Hudi
    定位 Databricks 主导,强绑定 Spark 通用表格式,引擎无关 流式数据湖,Upsert 优先
    ACID 事务
    Schema 演化 ✅(新增/删除列) ✅(新增/删除/重命名/调序)
    分区演化 ✅(隐式分区演化)
    时间旅行 ✅(Snapshot ID/时间戳)
    Upsert/Merge ✅ MERGE INTO ✅ MERGE INTO ✅(原生 Upsert 是核心特性)
    删除更新 ✅(V2 Row-Level) ✅(COW/MOR)
    流写入 ✅(Spark Structured Streaming) ✅(深度流式支持)
    计算引擎 Spark 为主(Delta 2.0+ 支持 Flink/Presto) Spark/Flink/Trino/Presto/Hive Spark/Flink/Trino/Presto
    社区活跃度 Databricks 商业驱动 Apache 顶级项目,中立开放 Apache 顶级项目,Uber 起源
    典型用户 Databricks 客户 Apple/Netflix/LinkedIn Uber/Grab/字节跳动
    特色 Z-Order 优化、Delta UniForm 隐藏分区、分区演化、快照隔离 COW/MOR 双模式、Compaction/Clustering

    Apache Paimon 简介(原 Flink Table Store):

    Paimon 是 Flink 社区发起的流批一体存储项目,2024 年成为 Apache 顶级项目。它的核心定位是流式数据湖:

    • 原生支持 Flink CDC 摄入,毫秒级延迟 Upsert
    • 同时支持 Spark 读写,流批统一
    • LSM Tree 架构,高吞吐写入与查询兼顾
    • 适合实时数仓场景(Flink + Paimon + Spark 分析)

    选型提示:如果团队以 Spark 为核心且用 Databricks,Delta Lake 最省心;如果追求引擎中立和多引擎生态,Iceberg 是趋势;如果核心需求是流式 Upsert 和增量处理,Hudi/Paimon 更合适。

    7.2 Spark + Iceberg 实战

    Iceberg 核心概念:

    概念说明
    Snapshot(快照) 表的一次完整状态,每次写入生成新快照,旧快照保留用于时间旅行
    Manifest(清单) 记录数据文件路径、分区信息、列级统计,Manifest List 管理多个 Manifest
    Metadata File 表元数据入口,记录当前 Snapshot、Schema、Partition Spec 等
    Partition Spec 分区规则,支持隐藏分区(如按天分区但查询可按小时过滤)
    Schema Evolution 新增/删除/重命名/调整列顺序,无需重写数据文件
    Partition Evolution 修改分区规则后旧数据不受影响,新数据按新规则写入

    Spark 提交 Iceberg 依赖配置:

    spark-submit \\
    –packages org.apache.iceberg:iceberg-spark-runtime-3.5_2.12:1.5.2 \\
    –conf spark.sql.extensions=org.apache.iceberg.spark.extensions.IcebergSparkSessionExtensions \\
    –conf spark.sql.catalog.spark_catalog=org.apache.iceberg.spark.SparkSessionCatalog \\
    –conf spark.sql.catalog.spark_catalog.type=hive \\
    –conf spark.sql.catalog.local=org.apache.iceberg.spark.SparkCatalog \\
    –conf spark.sql.catalog.local.type=hadoop \\
    –conf spark.sql.catalog.local.warehouse=s3://my-bucket/iceberg-warehouse \\
    your_app.py

    建表与写入:

    from pyspark.sql import SparkSession
    from pyspark.sql.functions import col, current_date, date_format

    spark = SparkSession.builder.appName("IcebergDemo").getOrCreate()

    # 建表
    spark.sql("""
    CREATE TABLE IF NOT EXISTS local.db.orders (
    order_id BIGINT,
    user_id BIGINT,
    product_id BIGINT,
    amount DECIMAL(10,2),
    order_time TIMESTAMP,
    dt STRING
    ) USING iceberg
    PARTITIONED BY (dt)
    """
    )

    # 批量写入
    batch_df = spark.read.parquet("s3://raw-data/orders/2024-01-15/")
    batch_df.writeTo("local.db.orders").overwritePartitions()

    # 流式写入(Structured Streaming)
    from pyspark.sql.functions import from_json, col

    kafka_schema = "order_id BIGINT, user_id BIGINT, product_id BIGINT, " \\
    "amount DECIMAL(10,2), order_time TIMESTAMP, dt STRING"

    stream_df = spark.readStream \\
    .format("kafka") \\
    .option("kafka.bootstrap.servers", "localhost:9092") \\
    .option("subscribe", "orders") \\
    .load() \\
    .selectExpr("CAST(value AS STRING) as json") \\
    .select(from_json(col("json"), kafka_schema).alias("data")) \\
    .select("data.*")

    stream_df.writeStream \\
    .format("iceberg") \\
    .option("path", "local.db.orders") \\
    .option("checkpointLocation", "s3://checkpoints/orders") \\
    .trigger(processingTime="1 minute") \\
    .start()

    查询与时间旅行:

    # 查询当前快照
    spark.sql("SELECT * FROM local.db.orders WHERE dt = '2024-01-15'").show()

    # Time Travel:按 Snapshot ID
    spark.sql("""
    SELECT * FROM local.db.orders.snapshotId(12345678901234567)
    """
    ).show()

    # Time Travel:按时间戳
    spark.sql("""
    SELECT * FROM local.db.orders.timestampAsOf('2024-01-15 10:00:00')
    """
    ).show()

    # 查看快照历史
    spark.sql("SELECT * FROM local.db.orders.snapshots").show()

    # 回滚到指定快照
    spark.sql("CALL local.system.rollback_to_snapshot('db.orders', 12345678901234567)")

    Schema 演化:

    # 新增列
    spark.sql("ALTER TABLE local.db.orders ADD COLUMN status STRING")

    # 删除列
    spark.sql("ALTER TABLE local.db.orders DROP COLUMN status")

    # 重命名列
    spark.sql("ALTER TABLE local.db.orders RENAME COLUMN amount TO total_amount")

    # 调整列顺序
    spark.sql("ALTER TABLE local.db.orders ALTER COLUMN order_id AFTER order_time")

    分区演化:

    # 从按天分区改为按月分区(旧数据不受影响)
    spark.sql("ALTER TABLE local.db.orders ADD PARTITION FIELD months(order_time)")
    spark.sql("ALTER TABLE local.db.orders DROP PARTITION FIELD dt")

    MERGE INTO(Upsert):

    # 创建更新数据集
    updates_df = spark.createDataFrame([
    (1001, 2001, 3001, 99.99, "2024-01-15 14:30:00", "2024-01-15", "PAID"),
    (1002, 2002, 3002, 59.99, "2024-01-15 15:00:00", "2024-01-15", "PAID"),
    ], ["order_id", "user_id", "product_id", "amount", "order_time", "dt", "status"])

    updates_df.createOrReplaceTempView("updates")

    spark.sql("""
    MERGE INTO local.db.orders t
    USING updates s
    ON t.order_id = s.order_id
    WHEN MATCHED THEN
    UPDATE SET t.amount = s.amount, t.status = s.status, t.order_time = s.order_time
    WHEN NOT MATCHED THEN
    INSERT (order_id, user_id, product_id, amount, order_time, dt, status)
    VALUES (s.order_id, s.user_id, s.product_id, s.amount, s.order_time, s.dt, s.status)
    """
    )

    7.3 Spark + Hudi 实战

    Hudi 表类型对比:

    特性COW(Copy On Write)MOR(Merge On Read)
    写入方式 每次写入更新重写整个 Parquet 文件 增量写入 Log 文件,定期合并
    写入延迟 高(重写文件) 低(追加日志)
    读取延迟 低(直接读 Parquet) 稍高(需合并 Parquet + Log)
    Compaction 不需要 需要(同步/异步)
    适用场景 批量更新、读多写少 流式 Upsert、写多读多
    存储效率 高(列式压缩) 稍低(Log 行式存储)
    典型场景 离线 ETL、T+1 报表 实时数仓、近实时分析

    Hudi 核心概念:

    概念说明
    Timeline 表的操作历史,包含 Commit/DeltaCommit/Compaction/Clean,每个 Instant 记录状态
    File Group 一组文件,由 File ID 标识,包含 Base File(Parquet)和 Log File
    File Slice 某一时刻 File Group 的快照 = 一个 Base File + 对应 Log 文件
    Compaction MOR 表将 Log 文件合并到 Base File 的过程,可同步或异步执行
    Clustering 将小文件合并、按列聚簇重写,优化查询性能

    Spark 写入 Hudi 代码示例:

    from pyspark.sql import SparkSession

    spark = SparkSession.builder \\
    .appName("HudiDemo") \\
    .config("spark.jars.packages",
    "org.apache.hudi:hudi-spark3.5-bundle_2.12:0.15.0") \\
    .config("spark.sql.extensions",
    "org.apache.spark.sql.hudi.HoodieSparkSessionExtension") \\
    .config("spark.sql.catalog.spark_catalog",
    "org.apache.spark.sql.hudi.catalog.HoodieCatalog") \\
    .getOrCreate()

    table_path = "s3://my-bucket/hudi/orders"
    table_name = "orders"

    # 批量导入(bulk_insert,最高效,不做去重)
    df = spark.read.parquet("s3://raw-data/orders/")
    df.write.format("hudi") \\
    .option("hoodie.table.name", table_name) \\
    .option("hoodie.datasource.write.recordkey.field", "order_id") \\
    .option("hoodie.datasource.write.partitionpath.field", "dt") \\
    .option("hoodie.datasource.write.precombine.field", "order_time") \\
    .option("hoodie.datasource.write.operation", "bulk_insert") \\
    .mode("overwrite") \\
    .save(table_path)

    # Upsert(默认操作,按主键更新或插入)
    updates_df = spark.read.parquet("s3://raw-data/orders/incremental/")
    updates_df.write.format("hudi") \\
    .option("hoodie.table.name", table_name) \\
    .option("hoodie.datasource.write.recordkey.field", "order_id") \\
    .option("hoodie.datasource.write.partitionpath.field", "dt") \\
    .option("hoodie.datasource.write.precombine.field", "order_time") \\
    .option("hoodie.datasource.write.operation", "upsert") \\
    .option("hoodie.upsert.shuffle.parallelism", 200) \\
    .mode("append") \\
    .save(table_path)

    # MOR 表写入(流式场景推荐)
    stream_df.writeStream.format("hudi") \\
    .option("hoodie.table.name", table_name) \\
    .option("hoodie.datasource.write.recordkey.field", "order_id") \\
    .option("hoodie.datasource.write.partitionpath.field", "dt") \\
    .option("hoodie.datasource.write.precombine.field", "order_time") \\
    .option("hoodie.datasource.write.table.type", "MERGE_ON_READ") \\
    .option("hoodie.datasource.write.operation", "upsert") \\
    .option("checkpointLocation", "s3://checkpoints/hudi-orders") \\
    .trigger(processingTime="1 minute") \\
    .start(table_path)

    Clustering / Compaction 配置:

    # 异步 Clustering(写入时自动优化小文件)
    hudi_options = {
    "hoodie.table.name": table_name,
    "hoodie.datasource.write.recordkey.field": "order_id",
    "hoodie.datasource.write.partitionpath.field": "dt",
    # Clustering
    "hoodie.clustering.async.enabled": "true",
    "hoodie.clustering.async.max.commits": "4", # 每 4 次提交触发一次
    "hoodie.clustering.plan.strategy.target.file.max.bytes": "134217728", # 128MB
    "hoodie.clustering.plan.strategy.small.file.limit": "629145600", # 600MB
    # MOR Compaction
    "hoodie.compact.inline": "true",
    "hoodie.compact.inline.max.delta.commits": "5", # 每 5 次 Delta Commit 触发
    "hoodie.compact.inline.trigger.strategy":
    "org.apache.hudi.table.action.compact.strategy.NumCommitsAfterLastCompactionTriggerStrategy",
    }

    # 离线触发 Clustering
    spark.sql(f"""
    CALL run_clustering(
    table => '
    {table_name}',
    path => '
    {table_path}',
    options => 'hoodie.clustering.plan.strategy.target.file.max.bytes=134217728'
    )
    """
    )

    # 离线触发 Compaction
    spark.sql(f"""
    CALL run_compaction(
    table => '
    {table_name}',
    path => '
    {table_path}'
    )
    """
    )

    7.4 Spark + Delta Lake

    Delta Lake 核心特性:

    特性说明
    ACID 事务 基于事务日志(_delta_log),并发写入通过乐观并发控制保证一致性
    Time Travel 通过版本号或时间戳读取历史数据快照
    Schema Enforcement 写入时严格校验 Schema,类型不匹配直接拒绝
    Schema Evolution 支持新增列、删除列(Delta 1.2+)
    MERGE INTO 原生 Upsert 语法,支持复杂条件
    OPTIMIZE + Z-ORDER 合并小文件并按排序列聚簇,加速点查
    Change Data Feed 增量读取变更数据(类似 Hudi 增量视图)

    Spark 读写 Delta 代码示例:

    from pyspark.sql import SparkSession
    from pyspark.sql.functions import col

    spark = SparkSession.builder \\
    .appName("DeltaDemo") \\
    .config("spark.jars.packages",
    "io.delta:delta-spark_2.12:3.1.0") \\
    .config("spark.sql.extensions",
    "io.delta.sql.DeltaSparkSessionExtension") \\
    .config("spark.sql.catalog.spark_catalog",
    "org.apache.spark.sql.delta.catalog.DeltaCatalog") \\
    .getOrCreate()

    # 写入 Delta 表
    df = spark.range(0, 1000).withColumn("value", col("id") * 10)
    df.write.format("delta").mode("overwrite").save("s3://my-bucket/delta/events")

    # 读取
    spark.read.format("delta").load("s3://my-bucket/delta/events").show()

    # SQL 建表
    spark.sql("""
    CREATE TABLE IF NOT EXISTS delta_events (
    id LONG, value LONG
    ) USING delta LOCATION 's3://my-bucket/delta/events'
    """
    )

    Time Travel:

    # 按版本号
    df_v2 = spark.read.format("delta") \\
    .option("versionAsOf", 2) \\
    .load("s3://my-bucket/delta/events")

    # 按时间戳
    df_ts = spark.read.format("delta") \\
    .option("timestampAsOf", "2024-01-15T10:30:00Z") \\
    .load("s3://my-bucket/delta/events")

    # SQL 语法
    spark.sql("SELECT * FROM delta_events VERSION AS OF 2")
    spark.sql("SELECT * FROM delta_events TIMESTAMP AS OF '2024-01-15 10:30:00'")

    MERGE / OPTIMIZE / ZORDER:

    # MERGE INTO(Upsert)
    updates = spark.createDataFrame([(1, 999), (1001, 500)], ["id", "value"])
    updates.createOrReplaceTempView("updates")

    spark.sql("""
    MERGE INTO delta_events t
    USING updates s
    ON t.id = s.id
    WHEN MATCHED THEN UPDATE SET *
    WHEN NOT MATCHED THEN INSERT *
    """
    )

    # OPTIMIZE:合并小文件
    spark.sql("OPTIMIZE delta_events")

    # Z-ORDER:按高频过滤列聚簇排序,加速点查
    spark.sql("OPTIMIZE delta_events ZORDER BY (id)")

    # 清理旧快照(VACUUM,默认保留 7 天)
    spark.sql("VACUUM delta_events RETAIN 168 HOURS")

    7.5 选型建议

    场景推荐格式理由
    Databricks 生态 / Spark 重度用户 Delta Lake 官方原生支持,Z-Order + OPTIMIZE 成熟
    多引擎共享(Spark + Flink + Trino) Iceberg 引擎中立,社区开放,隐藏分区设计优秀
    流式 Upsert / 实时数仓 Hudi (MOR) 原生流式支持,Compaction/Clustering 完善
    Flink CDC 实时入湖 + Spark 分析 Paimon Flink 原生流批一体,LSM 高吞吐写入
    离线批处理 + 偶尔 Upsert Iceberg / Delta COW 模式即可满足,运维简单
    严格 Schema 管理 + 审计合规 Iceberg 完整的 Schema/分区演化,快照隔离可审计
    已投入 Hadoop 生态 / Hive 兼容 Hudi / Iceberg 均支持 Hive Catalog,迁移成本低

    总结:没有"最好"的数据湖格式,只有"最合适"的选型。建议根据团队的计算引擎、更新频率、查询模式和运维能力综合决策。新项目若以 Spark 为核心且需要多引擎支持,Iceberg 是当前最安全的中立选择。


    8. Spark 性能调优

    性能调优是 Spark 工程师的核心能力。以下从多个维度系统讲解。

    8.1 数据倾斜诊断与解决

    数据倾斜是最常见的性能问题,表现为:大部分 Task 很快完成,但少数 Task 执行极慢甚至 OOM。

    诊断方法:

    • Spark UI 的 Stages 页面查看 Task 的 Shuffle Read/Write 数据量分布
    • 观察是否有 Task 处理的数据量远超其他 Task
    • 在 Spark SQL 中查看 AQE 倾斜处理日志

    解决方案:

    方案说明适用场景
    加盐(Salting) 对热点 Key 加随机前缀,分散到不同 Task 聚合类倾斜
    广播 Join 小表广播,避免 Shuffle Join 倾斜(小表 < 100MB)
    MapJoin / BroadcastHashJoin AQE 自动或手动 hint 大小表 Join
    拆分热点 Key 将热点 Key 单独处理后合并 复杂聚合
    两阶段聚合 先加随机前缀局部聚合,再去前缀全局聚合 groupBy 倾斜
    AQE Skew Join Spark 3.0 自动检测并拆分 Join 倾斜(推荐)

    # 加盐示例:对热点 Key 打散
    from pyspark.sql.functions import concat, lit, rand, col

    # 第一步:加随机前缀
    salted = df.withColumn("salted_key", concat(col("key"), lit("_"), (rand() * 10).cast("int")))
    partial = salted.groupBy("salted_key").agg(sum("value").alias("partial_sum"))

    # 第二步:去掉前缀再聚合
    result = partial.withColumn("key", split(col("salted_key"), "_")[0]) \\
    .groupBy("key").agg(sum("partial_sum").alias("total"))

    # 广播 Join 提示
    from pyspark.sql.functions import broadcast
    df_large.join(broadcast(df_small), "key")

    # SQL Hint
    spark.sql("SELECT /*+ BROADCAST(s) */ * FROM large l JOIN small s ON l.key = s.key")

    8.2 Shuffle 调优

    # 1. 调整 Shuffle 分区数(默认 200,通常需调整)
    spark.conf.set("spark.sql.shuffle.partitions", "200") # 通用建议:总核数的 2-3 倍

    # 2. 启用 AQE 自动合并分区
    spark.conf.set("spark.sql.adaptive.enabled", "true")
    spark.conf.set("spark.sql.adaptive.coalescePartitions.enabled", "true")

    # 3. Shuffle 压缩
    spark.conf.set("spark.shuffle.compress", "true")
    spark.conf.set("spark.shuffle.spill.compress", "true")
    spark.conf.set("spark.io.compression.codec", "zstd") # zstd 比 snappy 压缩率更好

    # 4. Shuffle 文件缓冲区
    spark.conf.set("spark.shuffle.file.buffer", "1MB") # 默认 32KB
    spark.conf.set("spark.reducer.maxSizeInFlight", "96MB") # 默认 48MB

    # 5. RDD API 中使用 coalesce 减少分区
    small_df = large_df.coalesce(10) # 窄依赖,不 Shuffle
    # repartition 会 Shuffle
    repartitioned = df.repartition(200, col("date"))

    8.3 内存调优

    # spark-submit 内存参数
    –driver-memory 4g \\
    –executor-memory 8g \\
    –executor-cores 4 \\
    –num-executors 20 \\
    –conf spark.memory.fraction=0.6 \\ # Spark 内存占比(默认 0.6)
    –conf spark.memory.storageFraction=0.5 # Storage 占 Spark 内存比例(默认 0.5)

    Executor 内存规划原则:

    • 每个 Executor 建议 4-8GB,过大易导致 GC 压力
    • 每个 Executor 2-5 个 Core,过多 Core 导致 HDFS I/O 竞争
    • 预留系统内存:spark.memory.fraction 默认为 0.6,即 60% 给 Spark,40% 给用户和其他开销

    序列化选择:

    序列化方式优点缺点
    Java Serialization(默认) 兼容好 慢、体积大
    Kryo Serialization 快 10x、体积小 需注册类

    spark.conf.set("spark.serializer", "org.apache.spark.serializer.KryoSerializer")
    # Scala: spark.registerKryoClasses(Array(classOf[MyClass]))

    缓存策略:

    • 优先使用 MEMORY_AND_DISK_SER 而非 MEMORY_ONLY,避免 OOM
    • DataFrame 建议用 spark.catalog.cacheTable() 或 df.cache(),列式存储更高效
    • 及时 unpersist() 不再使用的缓存

    8.4 并行度调优

    • 合理设置分区数:分区太少 → 并行度不足;分区太多 → Task 调度开销大
    • 经验值:分区数 ≈ 总 Executor Core 数的 2-3 倍
    • 读取时控制并行度:
      • HDFS 文件的 InputSplit 数量
      • spark.sql.files.maxPartitionBytes(默认 128MB)
    • AQE 自动调整:Spark 3.0+ 优先使用 AQE

    8.5 数据本地性

    Spark 倾向于将 Task 调度到数据所在节点,减少网络传输:

    级别说明
    PROCESS_LOCAL 数据在同一 JVM 中(最快)
    NODE_LOCAL 数据在同一节点
    RACK_LOCAL 数据在同一机架
    ANY 数据在任意位置(最慢)

    通常不需要手动调整,但如果集群网络延迟高,可适当增大 spark.locality.wait(默认 3s)。

    8.6 小文件问题

    问题:大量小文件导致 Task 数量过多,调度开销巨大。

    # 1. 写入时合并小文件
    spark.conf.set("spark.sql.adaptive.coalescePartitions.enabled", "true")
    spark.conf.set("spark.sql.adaptive.advisoryPartitionSizeInBytes", "128MB")

    # 2. 读取时合并小文件
    df = spark.read.option("mergeSchema", "true").parquet("hdfs:///data/")

    # 3. 使用 Hive 的 CombineHiveInputFormat(对 Hive 表)
    spark.conf.set("spark.hadoop.hive.input.format",
    "org.apache.hadoop.hive.ql.io.CombineHiveInputFormat")

    # 4. 写入后用 DISTRIBUTE BY 控制输出文件数
    df.write.partitionBy("dt").saveAsTable("table")
    # 或在 SQL 中
    # INSERT OVERWRITE TABLE t PARTITION(dt) SELECT … DISTRIBUTE BY dt

    8.7 列式存储与数据格式

    格式类型压缩特点
    Text/CSV 行式 无/弱 通用但慢
    JSON 行式 解析开销大
    Parquet 列式 Snappy/Gzip/Zstd Spark 默认推荐,列裁剪高效
    ORC 列式 Zlib/Snappy/Zstd Hive 生态好,ACID 支持
    Avro 行式 Snappy Schema 演进好,Kafka 常用

    # 写入 Parquet(推荐格式)
    df.write.mode("overwrite") \\
    .option("compression", "zstd") \\
    .parquet("hdfs:///data/output")

    # 分区裁剪 + 谓词下推(Parquet/ORC 自动支持)
    spark.sql("SELECT name, age FROM users WHERE dt = '2024-01-01' AND age > 25")
    # 只读 dt=2024-01-01 分区,且只扫描 name、age 列

    8.8 Join 策略选择

    Spark 支持多种 Join 实现,理解其原理对性能调优至关重要:

    Join 策略触发条件特点
    Broadcast Hash Join (BHJ) 小表 < spark.sql.autoBroadcastJoinThreshold(默认 10MB) 无 Shuffle,最快
    Sort-Merge Join (SMJ) 大表 Join 大表(默认) 需 Shuffle + 排序,稳定但慢
    Shuffled Hash Join (SHJ) 配置开启且大小表均适合构建哈希表 Shuffle 但不排序,比 SMJ 快
    Broadcast Nested Loop Join (BNLJ) 无等值条件的 Join(如 CROSS JOIN) 无 Shuffle,但 O(n×m)
    Cartesian Product 无条件 Join 结果集可能极大,慎用

    # 手动指定 Join 策略 Hint
    df1.join(broadcast(df2), "key") # 推荐:广播小表
    spark.sql("SELECT /*+ MERGE(l, r) */ * FROM large l JOIN large r ON l.id = r.id")
    spark.sql("SELECT /*+ SHUFFLE_HASH(l, r) */ * FROM large l JOIN medium r ON l.id = r.id")

    调优建议:

    • 大小表 Join:优先 Broadcast Hash Join,可适当调大 autoBroadcastJoinThreshold(如 100MB)
    • 大表 Join 大表:确保 Join Key 数据类型一致,避免隐式转换导致 Shuffle
    • 多表 Join:注意 Join 顺序,小表尽量先参与
    • AQE 开启后,运行时自动切换为 Broadcast Join

    8.9 分区策略

    合理的分区设计是性能优化的基础:

    # 写入时按字段分区(目录分区)
    df.write.partitionBy("dt", "region").parquet("hdfs:///data/sales")
    # 生成目录结构:/data/sales/dt=2024-01-01/region=CN/part-xxx.parquet

    # 分桶(Bucket):相同 Key 的数据落在同一文件,避免 Join 时 Shuffle
    df.write.bucketBy(32, "user_id").sortBy("user_id").saveAsTable("bucketed_users")

    # 读取时分区裁剪自动生效
    spark.sql("SELECT * FROM sales WHERE dt = '2024-01-01'") # 只读一个分区目录

    策略适用场景注意事项
    分区(Partition) 高基数的查询过滤字段(如日期) 分区数不宜过多(< 10000)
    分桶(Bucket) 频繁 Join 或聚合的字段 桶数需为 2 的幂,与 Join 表一致
    聚簇(Clustering) 数据湖中自动优化文件布局 Iceberg/Hudi/Delta Lake 支持

    8.10 Spark UI 实战调试指南

    Spark UI(默认端口 4040)是性能调优最重要的工具,能直观展示作业执行细节。以下逐页讲解排查思路。

    Jobs 页面:

    • 展示所有 Job 的状态(Succeeded/Failed/Running)、耗时、Stage 数
    • 点击 Job 进入详情,查看 DAG 可视化图:每个方框代表一个 Stage,箭头表示 Shuffle 依赖
    • 任务耗时分布:关注 Duration 列,若某 Job 耗时远超预期,进入对应 Stage 排查
    • 失败任务定位:Failed Jobs 区域直接显示异常堆栈,点击 Failed Stage 查看 Task 失败日志

    Stages 页面(核心!):

    • 展示每个 Stage 的 Task 数、输入/输出数据量、Shuffle Read/Write
    • Shuffle Read 列:关注 Min / 25th / Median / 75th / Max,若 Max 远大于 Median(如 Max 是 Median 的 5 倍以上),说明数据倾斜
    • Shuffle Write 列:Map 端写出的数据量,倾斜同样会表现为分布不均
    • Task Time 分布:GC Time 占比过高说明内存压力大;Scheduler Delay 过高说明资源不足
    • 点击 Stage 详情可查看每个 Task 的指标,支持排序

    倾斜识别要点:在 Stages 页面查看 Summary Metrics 表,比较 Task Max 与 Median 的 Shuffle Read Size。如果 Max 是 Median 的数倍甚至数十倍,基本可以确认倾斜。

    Storage 页面:

    • 展示缓存的 RDD/DataFrame,包括存储级别(Memory/Deserialized 等)、缓存分区数、占用内存/磁盘大小
    • 缓存命中率:Fraction Cached 列显示已缓存分区比例,若偏低说明部分分区未缓存(Executor 丢失或内存不足被淘汰)
    • 分区大小:Size in Memory/ExternalBlockStore 列可判断分区是否均匀
    • 内存 vs 磁盘:若大量数据在 Disk 上,说明内存不足,需调整缓存策略或增大内存

    Environment 页面:

    • 展示所有 Spark 配置属性(System Properties、Classpath、Spark Properties)
    • 配置排查:确认提交参数是否生效(如 spark.sql.shuffle.partitions、spark.executor.memory)
    • 可对比运行配置与预期配置,排查"参数没生效"类问题

    Executors 页面:

    列排查意义
    Task Time (Total/GC) GC Time 占 Total 10%+ 说明内存压力大,需调内存或减少缓存
    Shuffle Read/Write 各 Executor 的 Shuffle 数据量,差异大说明倾斜
    Input Size/Records 读取数据量分布,不均匀可能是文件大小不均
    Storage Memory 已用/总缓存内存,满了会溢写磁盘
    Failed Tasks 非零需查看日志,可能 OOM 或数据异常
    Log 链接 点击 stdout/stderr 查看 Executor 日志(YARN 模式需通过 RM 代理)

    SQL 页面(Spark SQL / DataFrame 作业必看):

    • 展示 SQL 查询的完整执行计划 DAG,节点显示 Scan/Filter/Project/Join/Aggregate/Sort 等算子
    • 点击节点查看详细指标:行数、数据量、耗时、CodeGen 信息
    • Join 识别:节点显示 SortMergeJoin/BroadcastHashJoin/ShuffledHashJoin,确认是否走了预期的 Join 策略
    • Broadcast 识别:若看到 BroadcastExchange 节点,说明小表被广播;若大表被广播可能 OOM
    • AQE Skew Partition:开启 AQE 后,倾斜 Join 会显示 CustomShuffleReader 节点,标注拆分的分区数
    • Scan 节点:查看 number of files read、size of files read、partition filters、data filters,确认分区裁剪和谓词下推是否生效

    Structured Streaming 页面:

    • Input Rate:每秒输入数据条数,反映流量波动
    • Processing Rate:每秒处理数据条数,若持续低于 Input Rate 说明处理能力不足
    • Batch Duration:每个微批的处理耗时,若持续接近 Trigger Interval 说明背压严重
    • Watermark 信息:展示当前 Watermark 位置、迟到数据丢弃情况
    • State Operator:状态算子(聚合/Join)的状态行数、内存占用

    实战:通过 UI 定位数据倾斜的步骤:

  • 打开 Stages 页面,找到耗时最长的 Stage,点击 Description 进入详情
  • 查看 Summary Metrics 表,对比 Shuffle Read 的 Max 和 Median。例如 Median 为 128MB 而 Max 为 8.5GB,确认倾斜
  • 点击 Tasks 表,按 Shuffle Read Size 降序排列,找到处理最大数据量的 Task
  • 记录该 Task 的 Locality Level 和 Executor ID,排除节点本地性问题
  • 返回 SQL 页面,找到对应 Stage 的执行计划,确认是 Join 还是聚合导致倾斜
  • 若为 Join 倾斜:检查 Join Key 分布,对热点 Key 加盐或开启 AQE Skew Join
  • 若为聚合倾斜:使用两阶段聚合(加盐局部聚合 + 去盐全局聚合)
  • 重新提交作业,在 Stages 页面验证 Task 耗时是否均匀
  • 8.11 资源规划与配置参数

    Executor 资源规划公式:

    推荐配置:
    –num-executors = ceil(总数据量 / (executor-cores × 单核算力))
    –executor-cores = 3 ~ 5(建议 4)
    –executor-memory = 4g ~ 8g(建议 6g)
    –driver-memory = 2g ~ 4g(复杂作业建议 4g+)

    经验法则:
    1. 每个 Executor 分配 4 Core,6GB 内存为黄金配比
    2. YARN 总核数 = num-executors × executor-cores
    3. 预留 10%~20% 集群资源给系统和其他服务
    4. Executor 数量 = 集群总核数 / executor-cores × 0.8

    常见踩坑:

    问题原因解决方案
    executor-memory 过大导致 GC 单 Executor 超过 8GB,JVM GC 停顿时间长 控制在 4-8GB,增加 Executor 数量而非单 Executor 内存
    cores 过多导致 HDFS 并发 单 Executor 超过 5 Core,HDFS 客户端并发写导致超时 控制在 3-5 Core,HDFS 并发度 = executor-cores × num-executors
    Driver OOM collect() 大量数据或广播大表 增大 driver-memory;避免 collect 大数据集
    YARN Container 被 Kill 物理内存超 Container 限制 增大 spark.yarn.executor.memoryOverhead

    YARN 模式内存开销:

    YARN Container 内存 = spark.executor.memory + spark.executor.memoryOverhead

    memoryOverhead 默认值 = max(executor-memory × spark.kubernetes.memoryOverheadFactor, 384MB)
    YARN 模式默认 factor = 0.10

    示例:
    –executor-memory 6g
    默认 memoryOverhead = max(6g × 0.10, 384MB) ≈ 614MB
    Container 总内存 = 6g + 614MB ≈ 6.7GB

    调优建议:
    –conf spark.yarn.executor.memoryOverhead=1g # 显式指定,避免默认不足

    关键配置参数速查表:

    参数名默认值说明推荐值
    spark.sql.shuffle.partitions 200 Shuffle 后分区数 总核数 × 2-3(开启 AQE 后可适当调大)
    spark.executor.memory 1g 每个 Executor 堆内存 4g-8g
    spark.executor.cores 1 每个 Executor CPU 核数 4
    spark.executor.instances Executor 数量(YARN/K8s) 按集群资源计算
    spark.sql.adaptive.enabled false(3.x 建议 true) 开启 AQE 自适应执行 true
    spark.sql.adaptive.coalescePartitions.enabled true AQE 自动合并小分区 true
    spark.sql.adaptive.coalescePartitions.minPartitionNum 1 合并后最小分区数 总核数 × 1
    spark.sql.adaptive.advisoryPartitionSizeInBytes 64MB AQE 目标分区大小 128MB
    spark.sql.adaptive.skewJoin.enabled true AQE 自动处理倾斜 Join true
    spark.sql.adaptive.skewJoin.skewedPartitionFactor 5 倾斜分区判定因子(倍数) 5(默认即可)
    spark.sql.adaptive.skewJoin.skewedPartitionThresholdInBytes 256MB 倾斜分区最小数据量阈值 256MB
    spark.sql.autoBroadcastJoinThreshold 10MB 小表自动广播阈值 100MB(视内存调整)
    spark.serializer org.apache.spark.serializer.JavaSerializer 序列化器 org.apache.spark.serializer.KryoSerializer
    spark.memory.fraction 0.6 Spark 内存占 Executor 比例 0.6(默认,缓存多可降至 0.5)
    spark.memory.storageFraction 0.5 Storage 占 Spark 内存比例 0.5(默认)
    spark.speculation false 推测执行(慢 Task 备份) true(集群空闲时开启)
    spark.sql.files.maxPartitionBytes 128MB 读取文件时单个分区最大字节 128MB(大集群可调至 256MB)
    spark.sql.files.openCostInBytes 4MB 打开文件的估算开销(小文件合并用) 8MB(小文件多时增大)
    spark.sql.broadcastTimeout 300s 广播等待超时 600(大表广播时)
    spark.network.timeout 120s 网络通信超时 300s(大 Shuffle 时)
    spark.streaming.kafka.consumer.cache.enabled true 缓存 Kafka 消费者 true(默认,避免重复创建)
    spark.sql.crossJoin.enabled false 允许笛卡尔积 false(默认,禁止意外笛卡尔积)
    spark.shuffle.service.enabled false External Shuffle Service true(YARN 动态分配必须)
    spark.dynamicAllocation.enabled false 动态资源分配 true(配合 ESS)
    spark.sql.session.timeZone JVM 默认 会话时区 Asia/Shanghai

    9. Spark 常见问题与踩坑

    9.1 OOM(OutOfMemoryError)

    类型原因解决方案
    Driver OOM collect() 拉取大量数据到 Driver 用 take()/sample()/write 替代;增大 driver-memory
    Executor OOM 数据倾斜、缓存过多、大表广播 解决倾斜;调整缓存策略;增大 executor-memory
    GC Overhead JVM 堆外内存不足 增大 spark.executor.memoryOverhead(默认 executor-memory × 0.1)
    Python Worker OOM Pandas UDF 内存占用大 增大 Arrow 批量;分批处理

    # 堆外内存配置(重要!)
    –conf spark.executor.memoryOverhead=2g \\
    –conf spark.driver.memoryOverhead=1g \\
    –conf spark.python.worker.memory=2g

    9.2 数据倾斜

    见 8.1 节。典型症状:99% 的 Task 很快完成,1% 的 Task 卡住不动。

    9.3 Shuffle Fetch Failed

    org.apache.spark.shuffle.FetchFailedException:
    Failed to connect to /xxx:xxxx

    常见原因与解决:

    • Executor 内存不足导致被 YARN Kill:增大 executor-memory / memoryOverhead
    • 网络超时:增大 spark.shuffle.io.maxRetries(默认 3)和 spark.shuffle.io.retryWait(默认 5s)
    • Shuffle 文件过大:增加 Shuffle 分区数
    • 节点故障:开启 spark.shuffle.service.enabled=true(External Shuffle Service)

    9.4 序列化错误

    org.apache.spark.SparkException: Task not serializable

    原因:在 RDD/DataFrame 的闭包中引用了不可序列化的对象(如数据库连接、非序列化的外部类)。

    解决:

    • 在函数内部创建不可序列化对象(如数据库连接),不要在 Driver 端创建后传入
    • 使用 foreachPartition / mapPartitions,每个分区创建一次连接
    • 使用 Kryo 序列化并注册类

    # ✅ 正确:在分区内创建连接
    def process_partition(rows):
    conn = create_db_connection() # 每个分区创建一次
    for row in rows:
    conn.insert(row)
    conn.close()

    df.foreachPartition(process_partition)

    9.5 时区问题

    # Spark 默认使用 UTC 时区,可能导致时间差 8 小时
    spark.conf.set("spark.sql.session.timeZone", "Asia/Shanghai")

    # 读取时指定时区
    df = spark.read.option("timestampFormat", "yyyy-MM-dd HH:mm:ss") \\
    .option("timeZone", "Asia/Shanghai").csv("data.csv")

    9.6 UDF 性能问题

    • Python UDF 性能差,优先使用 Spark SQL 内置函数
    • 必须用 UDF 时,使用 Pandas UDF(向量化)
    • 避免在 UDF 中创建重型对象,用 mapPartitions 替代

    9.7 其他常见坑

    问题说明
    count() 触发重算 未 cache 的 RDD/DataFrame,每次 Action 都从头计算
    collect() 内存溢出 确保结果集不超过 Driver 内存
    隐式类型转换错误 数字与字符串比较时注意类型
    Hive 分区不生效 执行 MSCK REPAIR TABLE 或添加分区
    Spark UI 看不到 YARN cluster 模式下通过 ResourceManager 代理访问
    文件已存在报错 使用 .mode("overwrite") 或先删除

    10. 端到端实战项目

    10.1 项目背景:电商用户行为日志分析平台

    项目目标:构建一个流批一体的电商用户行为分析平台,实时采集用户点击流数据,结合业务库数据,完成实时指标统计和离线报表分析。

    数据源:

    数据源内容采集方式
    Kafka 用户点击流 页面浏览、点击、搜索、加购等行为事件 Flume/SDK → Kafka
    MySQL 订单库 订单主表、订单明细 Flink CDC / Spark JDBC 全量+增量
    MySQL 用户库 用户基本信息、地域信息 每日全量同步
    MySQL 商品库 商品分类、价格、品牌 每日全量同步

    技术栈:

    • 计算引擎:Spark 3.5(Structured Streaming 实时 + Spark SQL 离线)
    • 数据湖:Apache Iceberg(ODS/DWD/DWS/ADS 分层存储)
    • 消息队列:Kafka
    • 缓存:Redis(实时指标查询)
    • 业务库:MySQL(ADS 报表数据)
    • 资源调度:YARN / K8s
    • 监控:Prometheus + Grafana

    分层架构: 在这里插入图片描述

    10.2 项目架构图

    在这里插入图片描述

    10.3 核心代码

    1. Kafka 消费 + 数据清洗写入 Iceberg ODS:

    from pyspark.sql import SparkSession
    from pyspark.sql.functions import from_json, col, current_timestamp
    from pyspark.sql.types import StructType, StructField, StringType, LongType, DoubleType

    spark = SparkSession.builder \\
    .appName("UserBehaviorODS") \\
    .config("spark.sql.catalog.local", "org.apache.iceberg.spark.SparkCatalog") \\
    .config("spark.sql.catalog.local.type", "hadoop") \\
    .config("spark.sql.catalog.local.warehouse", "s3://lakehouse/warehouse") \\
    .getOrCreate()

    schema = StructType([
    StructField("user_id", LongType()),
    StructField("event_type", StringType()), # view/click/cart/search/buy
    StructField("product_id", LongType()),
    StructField("category_id", LongType()),
    StructField("event_time", StringType()),
    StructField("device", StringType()),
    StructField("ip", StringType()),
    ])

    kafka_df = spark.readStream \\
    .format("kafka") \\
    .option("kafka.bootstrap.servers", "kafka1:9092,kafka2:9092") \\
    .option("subscribe", "user_behavior") \\
    .option("startingOffsets", "latest") \\
    .option("maxOffsetsPerTrigger", 100000) \\
    .load()

    parsed = kafka_df.select(
    from_json(col("value").cast("string"), schema).alias("data")
    ).select("data.*") \\
    .withColumn("ingest_time", current_timestamp()) \\
    .withColumn("dt", col("event_time").substr(1, 10))

    # 写入 Iceberg ODS(追加模式)
    parsed.writeStream \\
    .format("iceberg") \\
    .option("path", "local.ods.user_behavior") \\
    .option("checkpointLocation", "s3://checkpoints/ods_user_behavior") \\
    .trigger(processingTime="30 seconds") \\
    .outputMode("append") \\
    .start()

    2. 维度 Join(Broadcast)写入 DWD:

    from pyspark.sql.functions import broadcast

    # 加载维度表(小表广播)
    dim_user = spark.read.format("iceberg").load("local.dim.dim_user")
    dim_product = spark.read.format("iceberg").load("local.dim.dim_product")

    # 从 ODS 流式读取
    ods_stream = spark.readStream \\
    .format("iceberg") \\
    .option("stream-from-timestamp", str(int((__import__('time').time() 86400) * 1000))) \\
    .load("local.ods.user_behavior")

    # 维度关联:广播小表避免 Shuffle
    dwd = ods_stream.alias("o") \\
    .join(broadcast(dim_user).alias("u"),
    col("o.user_id") == col("u.user_id"), "left") \\
    .join(broadcast(dim_product).alias("p"),
    col("o.product_id") == col("p.product_id"), "left") \\
    .select(
    col("o.user_id"),
    col("u.age_group"),
    col("u.city"),
    col("o.event_type"),
    col("o.product_id"),
    col("p.category_name"),
    col("p.brand"),
    col("p.price"),
    col("o.event_time"),
    col("o.dt"),
    )

    # Upsert 写入 DWD
    dwd.writeStream \\
    .format("iceberg") \\
    .option("path", "local.dwd.dwd_user_behavior") \\
    .option("checkpointLocation", "s3://checkpoints/dwd_user_behavior") \\
    .trigger(processingTime="1 minute") \\
    .start()

    3. 窗口聚合写入 DWS:

    from pyspark.sql.functions import window, count, sum as _sum, when, col

    # 实时窗口聚合:每 1 分钟滚动窗口,统计各分类 PV/UV/加购数
    dws_agg = dwd \\
    .withWatermark("event_time", "10 minutes") \\
    .groupBy(
    window(col("event_time"), "1 minute"),
    col("category_name"),
    col("dt"),
    ).agg(
    count("*").alias("pv"),
    count("user_id").alias("event_count"),
    _sum(when(col("event_type") == "cart", 1).otherwise(0)).alias("cart_count"),
    _sum(when(col("event_type") == "buy", col("price")).otherwise(0)).alias("gmv"),
    ).select(
    col("window.start").alias("window_start"),
    col("window.end").alias("window_end"),
    col("category_name"),
    col("pv"),
    col("event_count"),
    col("cart_count"),
    col("gmv"),
    col("dt"),
    )

    dws_agg.writeStream \\
    .format("iceberg") \\
    .option("path", "local.dws.dws_category_realtime") \\
    .option("checkpointLocation", "s3://checkpoints/dws_category_rt") \\
    .outputMode("append") \\
    .trigger(processingTime="1 minute") \\
    .start()

    4. ADS 结果输出到 MySQL/Redis:

    def write_to_mysql(batch_df, batch_id):
    """将每批次结果写入 MySQL(幂等:REPLACE INTO)"""
    batch_df.write \\
    .format("jdbc") \\
    .option("url", "jdbc:mysql://mysql:3306/ads") \\
    .option("dbtable", "ads_category_realtime") \\
    .option("user", "spark") \\
    .option("password", "xxx") \\
    .option("driver", "com.mysql.cj.jdbc.Driver") \\
    .mode("append") \\
    .save()

    def write_to_redis(batch_df, batch_id):
    """将实时指标写入 Redis,供大屏查询"""
    import redis
    r = redis.Redis(host='redis', port=6379, db=0)
    for row in batch_df.collect():
    key = f"ads:category:{row['category_name']}:{row['window_start']}"
    r.hset(key, mapping={
    "pv": row["pv"],
    "cart_count": row["cart_count"],
    "gmv": float(row["gmv"]),
    })
    r.expire(key, 86400)

    # 双流写入
    dws_agg.writeStream \\
    .foreachBatch(lambda df, id: (write_to_mysql(df, id), write_to_redis(df, id))) \\
    .option("checkpointLocation", "s3://checkpoints/ads_sink") \\
    .trigger(processingTime="1 minute") \\
    .start()

    5. 离线补数(批式回刷历史数据):

    from datetime import datetime, timedelta

    def backfill(start_date, end_date):
    """批量回刷指定日期范围的数据"""
    spark = SparkSession.builder.appName("Backfill").getOrCreate()

    current = datetime.strptime(start_date, "%Y-%m-%d")
    end = datetime.strptime(end_date, "%Y-%m-%d")

    while current <= end:
    dt = current.strftime("%Y-%m-%d")
    print(f"Backfilling {dt} …")

    # 读取 ODS 指定分区
    ods_df = spark.read.format("iceberg") \\
    .load("local.ods.user_behavior") \\
    .where(col("dt") == dt)

    # 关联维度
    dwd_df = ods_df.alias("o") \\
    .join(broadcast(dim_user), "user_id", "left") \\
    .join(broadcast(dim_product), "product_id", "left")

    # 覆盖写入 DWD 指定分区
    dwd_df.writeTo("local.dwd.dwd_user_behavior") \\
    .overwritePartitions()

    current += timedelta(days=1)

    spark.stop()

    # 执行回刷
    backfill("2024-01-01", "2024-01-15")

    10.4 项目要点

    幂等写入与 Checkpoint 管理:

    • Checkpoint 目录存储 Kafka offset、聚合状态和事务元数据,删除 Checkpoint = 从头消费
    • 生产环境 Checkpoint 应放在可靠存储(HDFS/S3),并设置生命周期管理
    • 幂等写入策略:Iceberg 通过 MERGE INTO 按主键 Upsert;MySQL 使用 REPLACE INTO 或 INSERT … ON DUPLICATE KEY UPDATE
    • 重新部署时保留 Checkpoint,仅在需要重置时手动删除

    延迟数据处理(Watermark + allowedLateness):

    # 设置 10 分钟 Watermark:允许事件时间最多比处理时间晚 10 分钟
    df.withWatermark("event_time", "10 minutes")

    # Iceberg 流式写入可配合 to-snapshot 处理迟到数据
    # 超过 Watermark 的数据会被丢弃,但 Iceberg 的 Time Travel 可追溯
    # 对于重要的迟到数据,可在 ODS 层保留全量,通过离线补数修复 DWS/DWS

    • Watermark 阈值根据业务延迟特征设置(一般 5-30 分钟)
    • 对于严重延迟的数据(如客户端断网数小时),建议通过离线批处理补数修正
    • 在 ODS 层保留原始数据(Append Only,不删不改),作为"真相源"

    监控告警(Streaming Query Listener):

    from pyspark.sql.streaming import StreamingQueryListener

    class MyListener(StreamingQueryListener):
    def onQueryStarted(self, event):
    print(f"Query started: {event.id}")

    def onQueryProgress(self, event):
    progress = event.progress
    print(f"Batch {progress.batchId}: "
    f"input={progress.numInputRows} rows, "
    f"rate={progress.inputRowsPerSecond:.1f}/s, "
    f"duration={progress.batchDuration}ms")

    # 告警条件:处理速率低于输入速率(消费滞后)
    if progress.inputRowsPerSecond > progress.processedRowsPerSecond * 1.5:
    send_alert(f"Consumer lag! input={progress.inputRowsPerSecond}, "
    f"processed={progress.processedRowsPerSecond}")

    # 告警:批次耗时异常
    if progress.batchDuration > 120000: # 超过 2 分钟
    send_alert(f"Batch duration too long: {progress.batchDuration}ms")

    def onQueryTerminated(self, event):
    print(f"Query terminated: {event.id}, exception={event.exception}")
    if event.exception:
    send_alert(f"Streaming query failed: {event.exception}")

    spark.streams.addListener(MyListener())

    def send_alert(msg):
    """对接 Prometheus AlertManager / 钉钉 / 企业微信"""
    import requests
    requests.post("https://alert-webhook.example.com/", json={"text": msg})

    常见问题:

    问题原因解决方案
    小文件过多 每批次产生大量小 Parquet 文件 开启 Iceberg write.distribution-mode=hash;定期执行 rewrite_da

    (课程设计)基于java+springboot+vue前后端分离摩托车销售商城系统设计与实现-计算机毕设 附源码 91694

    master阅读(17)

                 

    基于java+springboot+vue前后端分离摩托车销售商城系统设计与实现

    摘  要

    本文设计并实现了一个基于Java、Spring Boot与Vue的前后端分离摩托车销售商城平台。后端采用Spring Boot构建RESTful服务,负责用户管理、商品信息维护、订单处理及权限控制等核心业务逻辑;前端使用Vue框架开发用户界面,实现商品浏览、购物车操作、订单提交、二手交易、社区交流及个人中心等功能模块。系统通过角色划分支持注册用户、商家用户和管理员三类身份,各自具备相应的操作权限与功能视图,保障数据安全与业务流程有序进行。开发过程中遵循软件工程规范,完成需求分析、架构设计、编码实现及功能测试等环节,验证了平台在实际运行中的稳定性与可用性。结果表明,所采用的技术方案能够有效支撑摩托车垂直电商场景下的多样化需求,具备良好的可维护性与扩展潜力,为同类应用提供了可行的技术参考。

    关键词:Spring Boot;Vue;前后端分离;摩托车商城;

    Abstract

    This paper designs and implements a front-end and back-end separated motorcycle sales mall platform based on Java, Spring Boot, and Vue. The backend employs Spring Boot to construct RESTful services, handling core business logic such as user management, product information maintenance, order processing, and permission control. The frontend utilizes the Vue framework to develop the user interface, enabling features like product browsing, shopping cart operations, order submission, second-hand trading, community interaction, and personal center modules. The system supports three user roles—registered users, merchant users, and administrators—each with corresponding operational permissions and functional views, ensuring data security and orderly business processes. During development, software engineering standards were followed, completing requirements analysis, architectural design, coding implementation, and functional testing, validating the platform's stability and usability in practical operation. The results demonstrate that the adopted technical solution effectively supports diverse requirements in motorcycle vertical e-commerce scenarios, offering good maintainability and scalability potential, providing a feasible technical reference for similar applications.

    Keywords: Spring Boot; Vue; Front-end and Back-end Separation; Motorcycle Mall;

    目  录

    摘  要

    Abstract

    1 绪论

    1.1 研究背景

    1.2 研究意义

    1.3 国内外研究现状

    2 相关技术介绍

    2.1 Java语言

    2.2 SpringBoot框架

    2.3 MySQL数据库

    2.4 Vue框架

    2.5 B/S模式

    3 系统分析

    3.1 系统可行性

    3.1.1 技术可行性

    3.1.2 经济可行性

    3.1.3 操作可行性

    3.2 系统功能需求

    3.2.1 注册用户功能需求分析

    3.2.2 前台商家用户功能需求分析

    3.2.3 后台商家用户功能需求分析

    3.2.4 管理员功能需求分析

    3.3 非功能性需求

    3.4 用户用例模型

    3.4.1 注册用户用例图

    3.4.2 (前台)商家用户用例图

    3.4.3 (后台)商家用户用例图

    3.4.4 管理员用例图

    3.5 系统流程分析

    3.5.1 注册登录流程图

    3.5.2 数据添加流程图

    3.5.3 数据修改流程图

    3.5.4 数据删除流程图

    3.5.5 数据搜索流程图

    4 系统设计

    4.1 系统架构设计

    4.2 功能模块设计

    4.3 数据库设计

    4.3.1 概念结构设计

    4.3.2 物理结构设计

    5 系统实现

    5.1 注册用户功能实现

    5.1.1 用户注册模块

    5.1.2 用户登录模块

    5.1.3 首页模块

    5.1.4 个人中心模块

    5.1.5 线上商城模块

    5.1.6 二手商品模块

    5.2 (前台)商家用户功能实现

    5.2.1 首页模块

    5.3 (后台)商家用户功能实现

    5.3.1 首页模块

    5.3.2 商城管理模块

    5.4 管理员功能实现

    5.4.1 首页模块

    5.4.2 二手商品管理模块

    5.4.3 购买记录管理模块

    5.4.4 退货记录管理模块

    5.4.5 通知发布模块

    5.5 测试目的方法

    5.6 功能测试用例

    5.7 测试结果分析

    结  论

    参考文献

    致  谢

    附  录

    1绪论

    1.1研究背景

    随着互联网技术的持续演进,电子商务已深度融入人们日常生活,成为商品交易的重要渠道。传统线下摩托车销售模式受限于地域、信息传播效率及服务响应速度,在满足用户多样化需求方面面临挑战。与此同时,消费者对购车体验、信息透明度及售后服务的要求不断提高,推动行业向数字化、平台化方向转型。在此背景下,构建一个功能完善、操作便捷、安全可靠的线上摩托车销售平台,成为连接用户与商家、提升交易效率的关键路径。

    近年来,前后端分离架构因其高内聚、低耦合的特性,被广泛应用于现代Web应用开发中。以Java语言为基础的Spring Boot框架凭借其快速开发、配置简化和生态成熟等优势,成为后端服务构建的主流选择;而Vue作为轻量级前端框架,以其组件化思想和高效渲染能力,显著提升了用户界面的交互体验。将二者结合,可为垂直领域电商平台提供稳定、灵活且易于维护的技术支撑,也为摩托车销售行业的数字化升级提供了可行的技术实现路径。

    1.2研究意义

    本研究聚焦于摩托车垂直电商场景,通过构建集新车销售、二手交易、资讯发布、社区互动与订单管理于一体的综合性平台,有效整合分散的市场资源,打破信息壁垒,提升交易透明度与用户参与度。平台不仅为普通用户提供便捷的选购与交流空间,也为商家提供商品展示、订单处理与客户管理的统一入口,有助于优化运营流程、降低沟通成本,促进供需高效匹配。

    从技术实践角度看,采用前后端分离架构进行开发,有助于明确职责边界,提升开发效率与系统可维护性。后端专注业务逻辑与数据安全,前端聚焦用户体验与界面表现,二者通过标准接口协同工作,既保障了系统的稳定性,也为后续功能扩展奠定良好基础。研究成果不仅验证了主流Web技术在特定行业应用中的适应性,也为其他垂直领域电商平台的建设提供了可复用的设计思路与实施经验,具有一定的推广价值与现实意义。

    1.3国内外研究现状

    在电子商务领域,国内外对垂直类交易平台的研究与实践已取得显著进展。国外较早起步的电商平台如eBay、Amazon等,在摩托车及相关配件销售方面积累了丰富经验。eBay自20世纪90年代起便支持用户发布二手摩托车信息,其成熟的信用评价机制、支付保障体系和全球物流网络,为用户提供了安全便捷的交易环境。近年来,专门聚焦于摩托车领域的平台如RevZilla、Motorcycle.com等进一步细化服务,不仅提供新车与配件在线购买,还整合维修指南、车型评测、社区论坛等内容,形成“交易+内容+服务”的生态闭环。这些平台普遍采用模块化架构,后端以Java或.NET为基础支撑高并发访问,前端则通过现代JavaScript框架提升交互体验,体现出对用户需求深度挖掘与技术实现高度融合的特点。

    国内电子商务发展虽起步稍晚,但增速迅猛,尤其在垂直电商方向展现出强大活力。以京东、淘宝为代表的综合平台早已开设摩托车及骑行装备专区,引入品牌官方旗舰店,保障正品与售后。与此同时,一批专注于机车文化的本土平台如“摩托邦”“哈罗摩托”等应运而生,除商品销售外,更强调用户社群运营与骑行生活方式的传播。例如,“哈罗摩托”通过整合车型数据库、用户口碑、线下活动报名等功能,构建起围绕摩托车用户的活跃社区,并逐步拓展至新车导购与二手车交易服务。在技术实现上,国内项目普遍采用Spring Boot构建后端服务,结合Vue或React开发前端界面,实现前后端解耦,提升开发效率与系统可维护性。尽管部分平台在订单履约、售后服务标准化等方面仍有提升空间,但整体已形成较为完整的线上摩托车消费生态,反映出市场对专业化、集成化数字平台的迫切需求与持续探索。

    2 相关技术介绍

    2.1Java语言

    Java是一种广泛应用于企业级软件开发的面向对象编程语言,具有平台无关性、安全性高、稳定性强等优点[1]。通过JVM(Java虚拟机)机制,Java实现了“一次编写,到处运行”的特性,极大地提升了代码的可移植性和复用性,其丰富的类库支持和成熟的生态系统,使得Java在Web应用、分布式系统及大型后台服务开发中占据主导地位[2]。在面向摩托车销售商城系统中,Java作为主要开发语言,能够实现高效的服务器端逻辑处理和数据管理,为系统的功能实现提供了坚实的基础,同时Java的多线程和网络编程能力有助于系统的高并发处理,为多个用户提供服务。

    2.2SpringBoot框架

    SpringBoot是一个开源框架,可以快速构建基于Spring的应用程序,通过简化配置和提供开箱即用的特性,极大地提升了开发效率[3]。SpringBoot框架安排了一系列默认的配置,支持开展自动化配置,减少了应用启动的复杂性,并集成嵌入式Web服务器,让开发者可以不依赖外部独立运行Java应用,无需借助外部容器[4]。SpringBoot框架的应用为面向摩托车销售商城系统的后端开发提供了强大的支持。开发者使用SpringBoot框架可以轻松创建RESTful风格的API,与前端进行数据交互,支持各种业务操作。数据访问层使用Spring Data JPA简单集成数据库,与MySQL进行操作,实现数据的持久化。安全管理方面利用Spring Security提供用户身份验证和授权功能,确保数据访问的安全性。

    2.3MySQL数据库

    MySQL是一款流行的开源关系型数据库管理系统,因其性能稳定、易于使用、社区活跃而被广泛应用于各类Web应用系统中,支持标准SQL语句,即以结构化查询语言为基础,通过表格的形式存储数据,具备良好的事务处理能力和多用户并发访问能力[5]。同时MySQL占用资源较少,适合中小型项目的数据库需求。本系统采用MySQL作为核心数据存储工具,负责管理面向摩托车销售商城系统中的用户信息、业务数据和操作日志等关键信息,通过合理的数据库设计和优化,保障数据的安全性和一致性,并且MySQL支持对多用户的高并发操作,能够快速处理系统中的各种数据查询需求,为后端提供所需的数据服务[6]。

    2.4Vue框架

    Vue 是一款轻量级、渐进式的前端JavaScript框架,专注于构建用户界面,具有学习成本低、开发效率高、组件化设计良好等特点,能够快速实现响应式数据绑定和动态页面交互[7]。Vue支持模块化开发,并与现代前端构建工具,如Webpack、Vite,兼容良好,适用于开发高性能、易维护的单页面应用(SPA)。在面向摩托车销售商城系统中,Vue框架被用于实现系统的前端界面,提供友好的交互体验,提升系统的交互体验与开发效率[8]。

    2.5B/S模式

    B/S(Browser/Server)架构是一种基于浏览器的三层体系结构,用户通过浏览器即可访问服务器端的应用程序和数据资源,无需安装额外客户端,具有部署简单、维护方便、跨平台兼容性强等优势[9]。B/S架构降低了用户的客户端配置要求,用户只需要使用浏览器即可访问应用。该架构下,业务逻辑主要集中在服务器端处理,前端仅负责展示与交互,有利于提升系统的可扩展性和安全性[10]。本面向摩托车销售商城系统基于B/S架构进行设计与开发,使得系统能够在不同设备和操作系统上便捷访问,增强系统的可用性与灵活性。

    3 系统分析

    3.1系统可行性

    3.1.1技术可行性

    系统采用Java语言结合Spring Boot框架构建后端服务,技术成熟、生态完善,具备良好的稳定性与可扩展性;前端使用Vue框架,支持组件化开发,提升界面交互体验与开发效率。前后端通过RESTful API进行通信,结构清晰、耦合度低,便于维护与迭代。相关开发工具、数据库(如MySQL)、中间件及部署环境均为主流且开源,技术门槛适中,开发资源丰富,能够有效支撑项目顺利实施。

    3.1.2经济可行性

    项目所需软硬件资源均为通用配置,无需高昂的专用设备投入。开发过程主要依赖开源技术栈,大幅降低授权与许可成本。平台上线后可通过商品交易佣金、广告位出租、商家入驻服务等方式实现盈利,具备清晰的商业化路径。同时,线上运营模式减少了实体门店租金与人力开支,长期来看具有良好的成本控制优势和投资回报潜力。

    3.1.3操作可行性

    系统功能设计贴合摩托车销售业务实际流程,界面简洁直观,用户学习成本低。注册用户、商家用户及管理员三类角色权限分明,操作逻辑符合各自使用习惯。后台管理功能集中、流程规范,便于商家高效处理订单与商品信息。整体交互流程顺畅,关键操作均有提示与反馈,保障用户在不同场景下均可顺利完成目标任务。

    3.2系统功能需求

    根据用户在系统中的操作权限及功能需求的差异,本摩托车销售商城系统在设计过程中将用户角色划分为三大主要类别:注册用户、商家用户和管理员。其中,商家用户进一步细分为前台商家用户与后台商家用户——前者面向平台前端,主要进行商品展示、社区互动及基础信息维护;后者进入专属管理后台,负责商品发布、订单处理、售后管理等核心运营操作。注册用户可浏览商品、下单购买、参与社区交流并管理个人资料;管理员则拥有全局管控权限,负责用户管理、内容审核、公告发布及系统配置等整体运维工作。针对每种角色的具体功能职责,其详细的功能模块划分如下所述。

    3.2.1注册用户功能需求分析

  • 首页:提供平台最新动态、推荐商品及活动入口,是用户访问的起点。
  • 交流社区:支持用户发布话题、参与讨论,增强社区互动性与用户粘性。
  • 网站公告:展示平台重要通知与更新信息,确保用户及时了解最新政策。
  • 摩托车资讯:发布行业新闻、车型评测等资讯内容,丰富用户体验。
  • 聊天室:为用户提供即时通讯工具,便于用户间或用户与商家间的沟通。
  • 线上商城:集中展示各类摩托车及其配件,支持商品搜索与筛选。
  • 商城管理:
  • 我的购物车:允许用户收藏感兴趣的商品以便后续购买。

    我的订单:查看历史订单详情及状态,支持订单跟踪。

    我的地址:维护收货地址信息,简化下单流程。

  • 文心一言:集成智能对话模块,解答用户关于摩托车的相关疑问。
  • 二手商品:浏览或发布二手摩托车信息,促进资源循环利用。
  • 我的账户:管理个人信息、修改密码、查看积分等。
  • 我的主页:个性化展示个人资料、收藏夹、历史记录等内容。
  • 个人中心:
  • 个人首页:用户个性化的信息汇总页。

    用户反馈:提交对平台的意见和建议,帮助改进服务。

    订单配送:追踪订单物流状态,掌握配送进度。

    交流社区:参与社区互动,分享经验。

    收藏:保存感兴趣的帖子或商品链接。

    评论管理:查看并回复自己发布的评论。

    3.2.2前台商家用户功能需求分析

  • 首页:同注册用户,获取平台最新动态。
  • 交流社区:与普通用户一样参与社区互动,提升品牌知名度。
  • 网站公告:了解平台规则变动,确保业务合规。
  • 摩托车资讯:关注市场趋势,调整销售策略。
  • 聊天室:直接与潜在客户沟通,解答疑问,促进成交。
  • 文心一言:利用智能对话工具,快速响应用户咨询。
  • 二手商品:发布自家库存中的二手摩托车信息,拓宽销售渠道。
  • 我的账户:管理商家账号信息,保障信息安全。
  • 我的主页:展示商家简介、联系方式等基本信息。
  • 个人中心:
  • 个人首页:商家信息概览。

    用户反馈:接收顾客评价与建议,持续优化服务质量。

    订单配送:实时更新订单发货状态,提高客户满意度。

    交流社区:作为行业专家分享专业知识,吸引目标客户。

    收藏:保存感兴趣的行业资讯或竞争对手动态。

    评论管理:管理顾客留言,维护品牌形象。

    3.2.3后台商家用户功能需求分析

  • 后台首页:提供店铺运营数据概览,辅助决策制定。
  • 交流管理:监控并处理用户在交流社区内的发言,维护良好氛围。
  • 商城管理:
  • 线上商城:管理店铺内所有商品的信息更新、上下架操作。

    分类列表:设置商品分类,方便顾客查找所需商品。

    订单列表:查看所有订单详情,进行批量处理。

    订单配送:更新订单的物流信息,确保顾客能及时收到商品。

    订单售后:处理退货退款请求,保证售后服务质量。

    3.2.4管理员功能需求分析

  • 后台首页:展示全站关键指标,如用户增长数、销售额等,帮助管理员全面了解平台状况。
  • 系统用户:管理平台上所有用户的账号信息,包括注册用户、商家用户等。
  • 二手商品管理:审核二手商品发布内容,确保符合规定。
  • 购买记录管理:查看并导出用户的购买记录,用于数据分析或争议解决。
  • 退货记录管理:处理退货事务,协调商家与买家之间的纠纷。
  • 用户反馈管理:收集并分析来自用户的反馈,指导产品改进方向。
  • 系统管理:执行系统配置、权限分配等高级操作,维护平台稳定运行。
  • 网站公告管理:发布官方公告,传达重要信息给全体用户。
  • 资源管理:管理图片、视频等多媒体资源,保证内容库的有序性。
  • 交流管理:监督交流社区的运作,营造健康和谐的讨论环境。
  • 商城管理:统筹线上商城的整体布局与规则设定,促进商业繁荣。
  • 通知发布:向特定用户群体发送消息通知,提高信息传达效率。
  • 3.3非功能性需求

    在面向摩托车销售商城系统的开发与设计过程中,除了满足基本的功能性需求外,还需要从多个维度考虑系统的非功能性要求,以保证系统的整体质量和用户体验。下面将从可靠性、可用性、性能、安全性等方面进行详细分析。

    可靠性:系统应具备较高的稳定性和容错能力,在面对高并发访问或异常操作时能够保持正常运行。通过合理的异常处理机制和日志记录功能,确保系统在发生错误时可以快速定位问题并恢复服务,从而保障用户操作的连续性和数据的一致性。

    可用性:系统界面应设计简洁、直观,操作流程清晰,便于不同层次的用户理解和使用。对于不同角色用户,应提供相应的引导提示和帮助信息以提升用户的操作效率和满意度。

    性能需求:系统需支持多用户并发访问,并能在合理的时间内响应用户的请求。在典型负载条件下,页面加载时间应控制在合理范围内,关键业务操作的响应延迟不宜过高。

    安全性:系统需具备完善的身份认证和权限控制机制,防止未经授权的访问和敏感信息泄露,并应对用户密码进行加密存储,对关键操作进行审计日志记录,确保系统具备一定的抗攻击能力和数据保护能力。

    3.4用户用例模型

    3.4.1注册用户用例图

    注册用户作为核心参与者,可访问首页、交流社区、网站公告、摩托车资讯、聊天室、线上商城及二手商城等公共模块,实现信息浏览与互动;通过“商城管理”可操作购物车、订单和收货地址,完成购物流程;进入“个人中心”后,可进行个人主页维护、提交用户反馈、查看订单配送状态、参与社区讨论以及管理收藏与评论内容,全面支持用户的个性化使用需求。注册用户角色用例图如图3-1所示。

    图3-1 注册用户用例图

    3.4.2(前台)商家用户用例图

    前台商家用户在系统中的功能交互关系。作为平台内容提供者与服务参与者,前台商家用户可访问首页、交流社区、网站公告、摩托车资讯、聊天室、文心一言及二手商城等公共模块,实现信息获取与用户互动;同时具备发布和管理二手商品的能力,支持与注册用户进行沟通交流;通过“我的账户”和“我的主页”维护自身身份信息,并可在“个人中心”中查看个人首页、提交用户反馈、跟踪订单配送状态、参与社区讨论以及管理收藏与评论内容,全面支撑其在平台上的运营与社交需求。(前台)商家用户角色用例图如如图3-2所示。

    图3-2 (前台)商家用户用例图

    3.4.3(后台)商家用户用例图

    后台商家用户在系统中的功能交互关系。作为店铺运营主体,后台商家用户可访问后台首页以查看经营数据概览;通过交流管理模块对社区内容进行监控与维护;在商城管理模块中实现对线上商品的发布与维护、商品分类的查看、订单列表的查询与处理、订单配送状态的更新以及售后事务的管理,全面支持其在平台上的商品运营与订单服务流程。(后台)商家用户角色用例图如如图3-3所示。

    图3-3(后台)商家用户用例图

    3.4.4管理员用例图

    管理员在系统中的功能交互关系。作为平台的最高权限角色,管理员可访问后台首页以掌握系统整体运行状态;负责系统用户账号的管理与维护,对二手商品、购买记录、退货记录及用户反馈等关键业务数据进行统一监管;具备系统配置、网站公告发布、资源上传与管理、交流内容审核以及商城全局设置的权限,并可通过通知发布功能向用户推送重要信息,全面保障平台的安全性、规范性与高效运营。管理员角色用例图如如图3-4所示。

    图3-4管理员用例图

    3.5系统流程分析

    3.5.1注册登录流程图

    用户注册登录模块主要是为了方便用户和管理员能够安全地访问系统并管理自己的信息,用户根据提示输入注册信息进行注册操作以获得个人账号,在登录界面输入账号密码等信息进行系统登录,用户注册登录流程如下图所示。

    图3-5 系统操作流程图

    3.5.2数据添加流程图

    用户成功登录系统后,即可进行添加数据操作。添加的数据具有一个由系统自动生成的特定编号,用户可根据提示输入其余信息并提交,系统会对提交信息进行验证,验证通过则显示添加数据成功。数据添加流程如下图所示。

    图3-6 数据添加流程图

    3.5.3数据修改流程图

    用户成功登录系统后,可进行修改数据操作,流程与添加数据操作相似,数据修改流程如下图所示。

    图3-7 数据修改流程图

    3.5.4数据删除流程图

    当系统里面存在一些无效或过期的数据信息,系统支持相关的管理人员对这些数据进行删除操作,数据修改流程如下图所示。

    图3-8 数据删除流程图

    3.5.5数据搜索流程图

    用户可以通过输入关键字方式在系统大量的数据中检索所需的信息,在搜索框输入关键字确认查询后,系统会自动检索数据库并显示特定的数据信息,数据搜索流程如下图所示。

    图3-9 数据搜索流程图

    4 系统设计

    4.1系统架构设计

    本系统在架构设计上采用了经典的三层体系结构,分别为表现层、业务逻辑层以及数据访问层,并通过集群化管理方式实现高效的数据处理与并发访问。其中表现层主要承担用户交互功能,负责信息的展示与用户输入的接收,将用户操作传达给业务逻辑层处理,并最终将处理的结果反馈给用户,采用JavaScript等技术进行构建简洁和直观前端页面,以提升界面响应速度与用户体验。业务逻辑层作为系统的核心处理单元,主要负责各类核心业务规则的封装与执行,按需与数据访问层交互,以便获取或更新所需的数据,该层基于SpringBoot框架进行开发,借助其良好的模块化设计和自动化配置能力来保障系统业务流程的高效运行与稳定性。数据访问层则专注于与数据库之间的交互操作,包括数据的存储、查询与更新等功能,系统中采用JPA或MyBatis等持久层框架来实现对MySQL数据库的访问,确保数据操作的准确性、一致性与安全性。系统架构图如图4-1所示,展示了各层次之间的功能划分与数据流向关系。

    图4-1 系统架构图

    4.2功能模块设计

    系统以用户角色为基础,划分为注册用户、前台商家用户、后台商家用户和管理员四大模块,各角色拥有独立的功能分支。注册用户可访问首页、交流社区、资讯浏览、线上商城及个人中心等基础功能,支持商品购买与互动;前台商家用户在具备普通用户功能的同时,可发布二手商品并参与社区运营;后台商家用户专注于店铺管理,涵盖商品上架、订单处理与售后维护等核心业务;管理员则负责全局管控,包括用户管理、交易记录监管、公告发布、资源维护及通知推送等,实现平台的高效协同与统一运维。不同角色对应的具体功能模块如图4-2所示,面向摩托车销售商城系统功能设计能够确保各角色能够负责其特定职责。

    图4-2 系统功能结构图

    4.3数据库设计

    4.3.1概念结构设计

    通过提供清晰的系统总E-R图,可以使其他用户快速理解和分析复杂的系统结构,更加轻松地掌握了解系统的整体架构和各功能组件之间的联系。根据对面向摩托车销售商城系统中各类实体及其属性的分析,本面向摩托车销售商城系统总体E-R图如图4-3所示,以直观地展示各实体之间的关系。

    图4-3 系统总体ER图

    4.3.2物理结构设计

    依据前一节对面向摩托车销售商城系统的整体E-R关系图的分析,为了满足系统功能需求,需要创建多个数据表。下面将着重介绍几个核心数据库表的设计结构,详细阐述这些关键数据库表的设计细节,包括但不限于字段定义、数据类型及其相互间的关系,从而为系统的稳定运行提供坚实的基础。

    表 4-1-access_token(登陆访问时长)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    create_time

    timestamp

    创建时间

    2

    info

    text

    65535

    信息

    3

    maxage

    int

    最大寿命:默认2小时

    4

    token

    varchar

    64

    临时访问牌

    5

    token_id

    int

    临时访问牌ID

    6

    update_time

    timestamp

    更新时间

    7

    user_id

    int

    用户编号

    表 4-2-address(收货地址)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    address

    varchar

    255

    地址

    2

    address_id

    int

    收货地址

    3

    create_time

    timestamp

    创建时间

    4

    default

    tinyint

    默认判断

    5

    name

    varchar

    32

    姓名

    6

    phone

    varchar

    13

    手机

    7

    postcode

    varchar

    8

    邮编

    8

    update_time

    timestamp

    更新时间

    9

    user_id

    mediumint

    用户ID

    表 4-3-article(文章)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    article_id

    mediumint

    文章id

    2

    content

    longtext

    4294967295

    正文

    3

    create_time

    timestamp

    创建时间

    4

    description

    text

    65535

    文章描述

    5

    hits

    int

    点击数

    6

    img

    varchar

    255

    封面图

    7

    praise_len

    int

    点赞数

    8

    source

    varchar

    255

    来源

    9

    tag

    varchar

    255

    标签

    10

    title

    varchar

    125

    标题

    11

    type

    varchar

    64

    文章分类

    12

    update_time

    timestamp

    更新时间

    13

    url

    varchar

    255

    来源地址

    表 4-4-article_type(文章分类)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    create_time

    timestamp

    创建时间

    2

    description

    varchar

    255

    描述

    3

    display

    smallint

    显示顺序

    4

    father_id

    smallint

    上级分类ID

    5

    icon

    text

    65535

    分类图标

    6

    name

    varchar

    16

    分类名称

    7

    type_id

    smallint

    分类ID

    8

    update_time

    timestamp

    更新时间

    9

    url

    varchar

    255

    外链地址

    表 4-5-auth(用户权限管理)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    add

    tinyint

    是否可增加

    2

    auth_id

    int

    授权ID

    3

    create_time

    timestamp

    创建时间

    4

    del

    tinyint

    是否可删除

    5

    field_add

    text

    65535

    添加字段

    6

    field_get

    text

    65535

    查询字段

    7

    field_set

    text

    65535

    修改字段

    8

    get

    tinyint

    是否可查看

    9

    mod_name

    varchar

    64

    模块名

    10

    mode

    varchar

    32

    跳转方式

    11

    option

    text

    65535

    配置

    12

    page_title

    varchar

    255

    页面标题

    13

    parent

    varchar

    64

    父级菜单

    14

    parent_sort

    int

    父级菜单排序

    15

    path

    varchar

    255

    路由路径

    16

    position

    varchar

    32

    位置

    17

    set

    tinyint

    是否可修改

    18

    table_name

    varchar

    64

    表名

    19

    table_nav

    varchar

    500

    跨表导航

    20

    table_nav_name

    varchar

    500

    跨表导航名称

    21

    update_time

    timestamp

    更新时间

    22

    user_group

    varchar

    64

    用户组

    表 4-6-business_user(商家用户)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    business_user_id

    int

    商家用户ID

    2

    create_by

    int

    创建用户ID

    3

    create_time

    datetime

    创建时间

    4

    examine_state

    varchar

    16

    审核状态

    5

    merchant_mobile_phone

    varchar

    16

    商家手机

    6

    merchant_name

    varchar

    64

    商家名称

    7

    update_time

    timestamp

    更新时间

    8

    user_id

    int

    用户ID

    表 4-7-cart(购物车)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    cart_id

    int

    购物车ID

    2

    create_time

    timestamp

    创建时间

    3

    description

    varchar

    255

    描述

    4

    goods_id

    mediumint

    商品id

    5

    img

    varchar

    255

    图片

    6

    norms

    varchar

    64

    规格

    7

    num

    int

    数量

    8

    price

    double

    单价

    9

    price_ago

    double

    原价

    10

    price_count

    double

    总价

    11

    state

    int

    状态:使用中,已失效

    12

    title

    varchar

    64

    标题

    13

    type

    varchar

    64

    商品分类

    14

    update_time

    timestamp

    更新时间

    15

    user_id

    int

    用户ID

    表 4-8-code_token(验证码)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    code

    varchar

    255

    验证码

    2

    code_token_id

    int

    验证码ID

    3

    create_time

    timestamp

    创建时间

    4

    expire_time

    timestamp

    失效时间

    5

    token

    varchar

    255

    令牌

    6

    update_time

    timestamp

    更新时间

    表 4-9-collect(收藏)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    collect_id

    int

    收藏ID

    2

    create_time

    timestamp

    创建时间

    3

    img

    varchar

    255

    封面

    4

    source_field

    varchar

    255

    来源字段

    5

    source_id

    int

    来源ID

    6

    source_table

    varchar

    255

    来源表

    7

    title

    varchar

    255

    标题

    8

    update_time

    timestamp

    更新时间

    9

    user_id

    int

    收藏人ID

    表 4-10-comment(评论)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    avatar

    varchar

    255

    头像地址

    2

    comment_id

    int

    评论ID

    3

    content

    longtext

    4294967295

    内容

    4

    create_time

    timestamp

    创建时间

    5

    hidden

    tinyint

    是否隐藏

    6

    nickname

    varchar

    255

    昵称

    7

    reply_to_id

    int

    回复评论ID

    8

    source_field

    varchar

    255

    来源字段

    9

    source_id

    int

    来源ID

    10

    source_table

    varchar

    255

    来源表

    11

    sticky

    tinyint

    是否置顶

    12

    update_time

    timestamp

    更新时间

    13

    user_id

    int

    评论人ID

    表 4-11-follow(用户关注)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    create_time

    timestamp

    创建时间

    2

    follow_id

    int

    用户关注ID

    3

    followed_avatar

    varchar

    255

    被关注人头像

    4

    followed_id

    int

    被关注人ID

    5

    followed_nickname

    varchar

    255

    被关注人昵称

    6

    follower_avatar

    varchar

    255

    关注人头像

    7

    follower_id

    int

    关注人ID

    8

    follower_nickname

    varchar

    255

    关注人昵称

    9

    update_time

    timestamp

    更新时间

    表 4-12-forum(论坛)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    avatar

    varchar

    255

    发帖人头像

    2

    content

    longtext

    4294967295

    正文

    3

    create_time

    timestamp

    创建时间

    4

    description

    varchar

    255

    描述

    5

    display

    smallint

    排序

    6

    forum_id

    mediumint

    论坛ID

    7

    hits

    int

    访问数

    8

    img

    text

    65535

    封面图

    9

    istop

    int

    是否置顶

    10

    keywords

    varchar

    125

    关键词

    11

    nickname

    varchar

    16

    昵称

    12

    praise_len

    int

    点赞数

    13

    tag

    varchar

    255

    标签

    14

    title

    varchar

    125

    标题

    15

    type

    varchar

    64

    论坛分类

    16

    update_time

    timestamp

    更新时间

    17

    url

    varchar

    255

    来源地址

    18

    user_id

    mediumint

    用户ID

    表 4-13-forum_type(论坛分类)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    create_time

    timestamp

    创建时间

    2

    description

    varchar

    255

    描述

    3

    father_id

    smallint

    上级分类ID

    4

    icon

    varchar

    255

    分类图标

    5

    name

    varchar

    16

    分类名称

    6

    type_id

    smallint

    分类ID

    7

    update_time

    timestamp

    更新时间

    8

    url

    varchar

    255

    外链地址

    表 4-14-goods(商品信息)

    编号

    字段名

    类型

    长度

    是否非空

    是否主键

    注释

    1

    content

    longtext

    4294967295

    正文

    2

    create_time

    timestamp

    创建时间

    3

    customize_field

    text

    65535

    自定义字段

    4

    description

    varchar

    255

    描述

    5

    goods_id

    mediumint

    产品ID

    6

    hits

    int

    点击量

    7

    img

    text

    65535

    封面图:用于显示于产品列表页

    8

    img_1

    text

    65535

    主图1

    9

    img_2

    text

    65535

    主图2

    10

    img_3

    text

    65535

    主图3

    11

    img_4

    text

    65535

    主图4

    12

    img_5

    text

    65535

    主图5

    13

    inventory

    int

    商品库存

    14

    list_status

    smallint

    上架状态(0:下架;1上架)

    15

    price

    double

    卖价

    16

    price_ago

    double

    原价

    17

    sales

    int

    销量

    18

    source_field

    varchar

    255

    来源字段

    19

    source_id

    int

    来源ID

    20

    source_table

    varchar

    255

    来源表

    21

    title

    varchar

    125

    标题

    22

    type

    varchar

    64

    商品分类

    23

    update_time

    timestamp

    更新时间

    24

    user_id

    int

    添加人

    表 4-15-goods_type(商品类型)

    编号

    字段名

    类型<

    第28章-实战项目一:博客系统

    master阅读(35)

    第28章 实战项目一:博客系统

    章节摘要

    本章通过构建一个完整的博客系统,综合运用前面所学的 ASP.NET Core、Entity Framework Core、身份认证、授权等技术,实现用户注册登录、文章发布、评论互动、分类标签、搜索等功能,掌握真实项目的开发流程和最佳实践。

    本章目录

    • 28.1 项目需求与架构设计
    • 28.2 数据访问层实现
    • 28.3 业务逻辑层实现
    • 28.4 API 控制器层实现
    • 28.5 Program.cs 配置与启动
    • 28.6 单元测试与集成测试
    • 28.7 部署配置与常见误区

    28.1 项目需求与架构设计

    项目背景

    构建一个现代化的个人博客系统,支持多用户协作、文章管理、评论互动、分类标签等功能。

    目标用户:

    • 博主:发布和管理文章
    • 读者:浏览文章、发表评论、搜索内容
    • 管理员:管理用户、审核内容

    功能需求

    核心功能:

    1. 用户系统

    • 用户注册与登录(本地账号 + OAuth)
    • 用户角色(Admin、Author、Reader)
    • 用户资料管理
    • 密码修改与重置

    2. 文章系统

    • 发布文章(支持 Markdown)
    • 编辑和删除文章
    • 文章草稿保存
    • 文章分类和标签
    • 文章搜索
    • 文章分页浏览
    • 文章阅读统计

    3. 评论系统

    • 发表评论
    • 评论回复(嵌套评论)
    • 评论点赞
    • 评论审核(管理员)

    4. 分类与标签

    • 创建和管理分类
    • 创建和管理标签
    • 按分类/标签筛选文章

    5. 管理后台

    • 用户管理
    • 文章管理
    • 评论管理
    • 系统统计

    非功能需求:

    • 响应式设计(支持移动端)
    • SEO 优化
    • 性能优化(缓存、分页)
    • 安全性(XSS、CSRF、SQL注入防护)
    • 代码可维护性

    技术选型

    后端技术栈:

    • ASP.NET Core 8.0 Web API
    • Entity Framework Core 8.0
    • SQL Server / PostgreSQL
    • ASP.NET Core Identity(用户认证)
    • JWT Bearer 认证
    • AutoMapper(对象映射)
    • FluentValidation(数据验证)
    • Serilog(日志记录)

    前端技术栈(可选):

    • Blazor / React / Vue.js
    • Bootstrap / Tailwind CSS
    • Markdown 编辑器

    项目架构

    采用分层架构(Clean Architecture):

    BlogSystem/
    ├── BlogSystem.Api/ # API 层(控制器、中间件)
    │ ├── Controllers/
    │ │ ├── AuthController.cs
    │ │ ├── PostsController.cs
    │ │ ├── CommentsController.cs
    │ │ ├── CategoriesController.cs
    │ │ └── TagsController.cs
    │ ├── Middlewares/
    │ └── Program.cs
    ├── BlogSystem.Core/ # 核心层(领域实体、接口)
    │ ├── Entities/
    │ │ ├── User.cs
    │ │ ├── Post.cs
    │ │ ├── Comment.cs
    │ │ ├── Category.cs
    │ │ └── Tag.cs
    │ ├── Interfaces/
    │ │ ├── IRepository.cs
    │ │ ├── IPostService.cs
    │ │ └── ICommentService.cs
    │ └── DTOs/
    ├── BlogSystem.Infrastructure/ # 基础设施层(数据访问、外部服务)
    │ ├── Data/
    │ │ ├── BlogDbContext.cs
    │ │ └── Repositories/
    │ ├── Services/
    │ └── Extensions/
    └── BlogSystem.Tests/ # 测试层
    ├── UnitTests/
    └── IntegrationTests/

    架构说明:

    • API 层:处理 HTTP 请求,调用服务层
    • Core 层:领域模型、业务接口、DTO
    • Infrastructure 层:数据访问、第三方服务
    • Tests 层:单元测试和集成测试

    依赖关系:

    Api → Core ← Infrastructure

    Tests

    数据库设计

    ER 图概览:

    ┌─────────────┐ ┌─────────────┐
    │ Users │ │ Posts │
    ├─────────────┤ ├─────────────┤
    │ Id (PK) │────────<│ AuthorId(FK)│
    │ UserName │ 1:N │ Title │
    │ Email │ │ Content │
    │ PasswordHash│ │ Slug │
    │ Role │ │ CategoryId │
    │ CreatedAt │ │ Status │
    └─────────────┘ │ ViewCount │
    │ CreatedAt │
    │ PublishedAt │
    └─────────────┘
    │ M:N

    ┌─────────────┐
    │ PostTags │
    ├─────────────┤
    │ PostId (FK) │
    │ TagId (FK) │
    └─────────────┘

    │ N:1
    ┌─────────────┐
    │ Tags │
    ├─────────────┤
    │ Id (PK) │
    │ Name │
    │ Slug │
    └─────────────┘

    ┌─────────────┐ ┌─────────────┐
    │ Categories │ │ Comments │
    ├─────────────┤ ├─────────────┤
    │ Id (PK) │────────<│ PostId (FK) │
    │ Name │ 1:N │ AuthorId(FK)│
    │ Slug │ │ Content │
    │ Description │ │ ParentId(FK)│
    └─────────────┘ │ IsApproved │
    │ CreatedAt │
    └─────────────┘

    表设计详情:

    1. Users(用户表)

    CREATE TABLE Users (
    Id NVARCHAR(450) PRIMARY KEY,
    UserName NVARCHAR(256) NOT NULL UNIQUE,
    Email NVARCHAR(256) NOT NULL UNIQUE,
    PasswordHash NVARCHAR(MAX) NOT NULL,
    FullName NVARCHAR(100),
    Bio NVARCHAR(500),
    AvatarUrl NVARCHAR(500),
    Role NVARCHAR(50) NOT NULL DEFAULT 'Reader',
    EmailConfirmed BIT NOT NULL DEFAULT 0,
    CreatedAt DATETIME2 NOT NULL DEFAULT GETUTCDATE(),
    UpdatedAt DATETIME2
    );

    2. Categories(分类表)

    CREATE TABLE Categories (
    Id INT PRIMARY KEY IDENTITY(1,1),
    Name NVARCHAR(100) NOT NULL UNIQUE,
    Slug NVARCHAR(100) NOT NULL UNIQUE,
    Description NVARCHAR(500),
    CreatedAt DATETIME2 NOT NULL DEFAULT GETUTCDATE()
    );

    3. Posts(文章表)

    CREATE TABLE Posts (
    Id INT PRIMARY KEY IDENTITY(1,1),
    Title NVARCHAR(200) NOT NULL,
    Slug NVARCHAR(200) NOT NULL UNIQUE,
    Content NVARCHAR(MAX) NOT NULL,
    Excerpt NVARCHAR(500),
    FeaturedImageUrl NVARCHAR(500),
    AuthorId NVARCHAR(450) NOT NULL,
    CategoryId INT NOT NULL,
    Status NVARCHAR(20) NOT NULL DEFAULT 'Draft', — Draft, Published, Archived
    ViewCount INT NOT NULL DEFAULT 0,
    CreatedAt DATETIME2 NOT NULL DEFAULT GETUTCDATE(),
    UpdatedAt DATETIME2,
    PublishedAt DATETIME2,
    FOREIGN KEY (AuthorId) REFERENCES Users(Id),
    FOREIGN KEY (CategoryId) REFERENCES Categories(Id)
    );

    CREATE INDEX IX_Posts_AuthorId ON Posts(AuthorId);
    CREATE INDEX IX_Posts_CategoryId ON Posts(CategoryId);
    CREATE INDEX IX_Posts_Status ON Posts(Status);
    CREATE INDEX IX_Posts_PublishedAt ON Posts(PublishedAt DESC);

    4. Tags(标签表)

    CREATE TABLE Tags (
    Id INT PRIMARY KEY IDENTITY(1,1),
    Name NVARCHAR(50) NOT NULL UNIQUE,
    Slug NVARCHAR(50) NOT NULL UNIQUE,
    CreatedAt DATETIME2 NOT NULL DEFAULT GETUTCDATE()
    );

    5. PostTags(文章-标签关联表)

    CREATE TABLE PostTags (
    PostId INT NOT NULL,
    TagId INT NOT NULL,
    PRIMARY KEY (PostId, TagId),
    FOREIGN KEY (PostId) REFERENCES Posts(Id) ON DELETE CASCADE,
    FOREIGN KEY (TagId) REFERENCES Tags(Id) ON DELETE CASCADE
    );

    CREATE INDEX IX_PostTags_TagId ON PostTags(TagId);

    6. Comments(评论表)

    CREATE TABLE Comments (
    Id INT PRIMARY KEY IDENTITY(1,1),
    PostId INT NOT NULL,
    AuthorId NVARCHAR(450) NOT NULL,
    ParentCommentId INT NULL, — 用于嵌套评论
    Content NVARCHAR(1000) NOT NULL,
    IsApproved BIT NOT NULL DEFAULT 0,
    LikeCount INT NOT NULL DEFAULT 0,
    CreatedAt DATETIME2 NOT NULL DEFAULT GETUTCDATE(),
    UpdatedAt DATETIME2,
    FOREIGN KEY (PostId) REFERENCES Posts(Id) ON DELETE CASCADE,
    FOREIGN KEY (AuthorId) REFERENCES Users(Id),
    FOREIGN KEY (ParentCommentId) REFERENCES Comments(Id)
    );

    CREATE INDEX IX_Comments_PostId ON Comments(PostId);
    CREATE INDEX IX_Comments_AuthorId ON Comments(AuthorId);
    CREATE INDEX IX_Comments_ParentCommentId ON Comments(ParentCommentId);
    CREATE INDEX IX_Comments_CreatedAt ON Comments(CreatedAt DESC);

    实体类定义

    Core/Entities/User.cs:

    using Microsoft.AspNetCore.Identity;

    namespace BlogSystem.Core.Entities
    {
    public class User : IdentityUser
    {
    public string FullName { get; set; } = "";
    public string? Bio { get; set; }
    public string? AvatarUrl { get; set; }
    public DateTime CreatedAt { get; set; } = DateTime.UtcNow;
    public DateTime? UpdatedAt { get; set; }

    // 导航属性
    public ICollection<Post> Posts { get; set; } = new List<Post>();
    public ICollection<Comment> Comments { get; set; } = new List<Comment>();
    }
    }

    Core/Entities/Category.cs:

    namespace BlogSystem.Core.Entities
    {
    public class Category
    {
    public int Id { get; set; }
    public string Name { get; set; } = "";
    public string Slug { get; set; } = "";
    public string? Description { get; set; }
    public DateTime CreatedAt { get; set; } = DateTime.UtcNow;

    // 导航属性
    public ICollection<Post> Posts { get; set; } = new List<Post>();
    }
    }

    Core/Entities/Post.cs:

    namespace BlogSystem.Core.Entities
    {
    public class Post
    {
    public int Id { get; set; }
    public string Title { get; set; } = "";
    public string Slug { get; set; } = "";
    public string Content { get; set; } = "";
    public string? Excerpt { get; set; }
    public string? FeaturedImageUrl { get; set; }

    // 外键
    public string AuthorId { get; set; } = "";
    public int CategoryId { get; set; }

    // 状态
    public PostStatus Status { get; set; } = PostStatus.Draft;
    public int ViewCount { get; set; } = 0;

    // 时间戳
    public DateTime CreatedAt { get; set; } = DateTime.UtcNow;
    public DateTime? UpdatedAt { get; set; }
    public DateTime? PublishedAt { get; set; }

    // 导航属性
    public User Author { get; set; } = null!;
    public Category Category { get; set; } = null!;
    public ICollection<Tag> Tags { get; set; } = new List<Tag>();
    public ICollection<Comment> Comments { get; set; } = new List<Comment>();
    }

    public enum PostStatus
    {
    Draft, // 草稿
    Published, // 已发布
    Archived // 已归档
    }
    }

    Core/Entities/Tag.cs:

    namespace BlogSystem.Core.Entities
    {
    public class Tag
    {
    public int Id { get; set; }
    public string Name { get; set; } = "";
    public string Slug { get; set; } = "";
    public DateTime CreatedAt { get; set; } = DateTime.UtcNow;

    // 导航属性
    public ICollection<Post> Posts { get; set; } = new List<Post>();
    }
    }

    Core/Entities/Comment.cs:

    namespace BlogSystem.Core.Entities
    {
    public class Comment
    {
    public int Id { get; set; }
    public int PostId { get; set; }
    public string AuthorId { get; set; } = "";
    public int? ParentCommentId { get; set; } // 父评论ID(用于嵌套评论)
    public string Content { get; set; } = "";
    public bool IsApproved { get; set; } = false;
    public int LikeCount { get; set; } = 0;
    public DateTime CreatedAt { get; set; } = DateTime.UtcNow;
    public DateTime? UpdatedAt { get; set; }

    // 导航属性
    public Post Post { get; set; } = null!;
    public User Author { get; set; } = null!;
    public Comment? ParentComment { get; set; }
    public ICollection<Comment> Replies { get; set; } = new List<Comment>();
    }
    }

    DTOs 定义

    Core/DTOs/PostDTOs.cs:

    using System.ComponentModel.DataAnnotations;

    namespace BlogSystem.Core.DTOs
    {
    // 创建文章请求
    public class CreatePostRequest
    {
    [Required(ErrorMessage = "标题不能为空")]
    [StringLength(200, ErrorMessage = "标题长度不能超过200字符")]
    public string Title { get; set; } = "";

    [Required(ErrorMessage = "内容不能为空")]
    public string Content { get; set; } = "";

    [StringLength(500, ErrorMessage = "摘要长度不能超过500字符")]
    public string? Excerpt { get; set; }

    public string? FeaturedImageUrl { get; set; }

    [Required(ErrorMessage = "分类不能为空")]
    public int CategoryId { get; set; }

    public List<string> Tags { get; set; } = new List<string>();

    public PostStatus Status { get; set; } = PostStatus.Draft;
    }

    // 更新文章请求
    public class UpdatePostRequest
    {
    [Required]
    [StringLength(200)]
    public string Title { get; set; } = "";

    [Required]
    public string Content { get; set; } = "";

    [StringLength(500)]
    public string? Excerpt { get; set; }

    public string? FeaturedImageUrl { get; set; }

    [Required]
    public int CategoryId { get; set; }

    public List<string> Tags { get; set; } = new List<string>();

    public PostStatus Status { get; set; }
    }

    // 文章响应
    public class PostResponse
    {
    public int Id { get; set; }
    public string Title { get; set; } = "";
    public string Slug { get; set; } = "";
    public string Content { get; set; } = "";
    public string? Excerpt { get; set; }
    public string? FeaturedImageUrl { get; set; }
    public string Status { get; set; } = "";
    public int ViewCount { get; set; }
    public DateTime CreatedAt { get; set; }
    public DateTime? PublishedAt { get; set; }

    // 作者信息
    public string AuthorId { get; set; } = "";
    public string AuthorName { get; set; } = "";
    public string? AuthorAvatarUrl { get; set; }

    // 分类信息
    public int CategoryId { get; set; }
    public string CategoryName { get; set; } = "";

    // 标签
    public List<TagResponse> Tags { get; set; } = new List<TagResponse>();

    // 评论数
    public int CommentCount { get; set; }
    }

    // 文章列表响应(简化版)
    public class PostListResponse
    {
    public int Id { get; set; }
    public string Title { get; set; } = "";
    public string Slug { get; set; } = "";
    public string? Excerpt { get; set; }
    public string? FeaturedImageUrl { get; set; }
    public int ViewCount { get; set; }
    public DateTime PublishedAt { get; set; }
    public string AuthorName { get; set; } = "";
    public string CategoryName { get; set; } = "";
    public List<string> Tags { get; set; } = new List<string>();
    public int CommentCount { get; set; }
    }
    }

    Core/DTOs/CommentDTOs.cs:

    using System.ComponentModel.DataAnnotations;

    namespace BlogSystem.Core.DTOs
    {
    // 创建评论请求
    public class CreateCommentRequest
    {
    [Required(ErrorMessage = "评论内容不能为空")]
    [StringLength(1000, MinimumLength = 1, ErrorMessage = "评论长度必须在1-1000字符之间")]
    public string Content { get; set; } = "";

    public int? ParentCommentId { get; set; }
    }

    // 评论响应
    public class CommentResponse
    {
    public int Id { get; set; }
    public int PostId { get; set; }
    public string Content { get; set; } = "";
    public bool IsApproved { get; set; }
    public int LikeCount { get; set; }
    public DateTime CreatedAt { get; set; }

    // 作者信息
    public string AuthorId { get; set; } = "";
    public string AuthorName { get; set; } = "";
    public string? AuthorAvatarUrl { get; set; }

    // 父评论ID(用于嵌套显示)
    public int? ParentCommentId { get; set; }

    // 回复列表
    public List<CommentResponse> Replies { get; set; } = new List<CommentResponse>();
    }
    }

    Core/DTOs/CategoryDTOs.cs:

    using System.ComponentModel.DataAnnotations;

    namespace BlogSystem.Core.DTOs
    {
    public class CreateCategoryRequest
    {
    [Required(ErrorMessage = "分类名称不能为空")]
    [StringLength(100, ErrorMessage = "分类名称长度不能超过100字符")]
    public string Name { get; set; } = "";

    [StringLength(500, ErrorMessage = "描述长度不能超过500字符")]
    public string? Description { get; set; }
    }

    public class CategoryResponse
    {
    public int Id { get; set; }
    public string Name { get; set; } = "";
    public string Slug { get; set; } = "";
    public string? Description { get; set; }
    public int PostCount { get; set; } // 该分类下的文章数
    }
    }

    Core/DTOs/TagResponse.cs:

    namespace BlogSystem.Core.DTOs
    {
    public class TagResponse
    {
    public int Id { get; set; }
    public string Name { get; set; } = "";
    public string Slug { get; set; } = "";
    public int PostCount { get; set; } // 该标签下的文章数
    }
    }


    28.2 数据访问层实现

    DbContext 配置

    Infrastructure/Data/BlogDbContext.cs:

    using BlogSystem.Core.Entities;
    using Microsoft.AspNetCore.Identity.EntityFrameworkCore;
    using Microsoft.EntityFrameworkCore;

    namespace BlogSystem.Infrastructure.Data
    {
    public class BlogDbContext : IdentityDbContext<User>
    {
    public BlogDbContext(DbContextOptions<BlogDbContext> options)
    : base(options)
    {
    }

    public DbSet<Post> Posts { get; set; }
    public DbSet<Category> Categories { get; set; }
    public DbSet<Tag> Tags { get; set; }
    public DbSet<Comment> Comments { get; set; }

    protected override void OnModelCreating(ModelBuilder builder)
    {
    base.OnModelCreating(builder);

    // 配置 User
    builder.Entity<User>(entity =>
    {
    entity.Property(e => e.FullName).HasMaxLength(100);
    entity.Property(e => e.Bio).HasMaxLength(500);
    entity.Property(e => e.AvatarUrl).HasMaxLength(500);
    entity.HasIndex(e => e.Email).IsUnique();
    });

    // 配置 Category
    builder.Entity<Category>(entity =>
    {
    entity.Property(e => e.Name).IsRequired().HasMaxLength(100);
    entity.Property(e => e.Slug).IsRequired().HasMaxLength(100);
    entity.Property(e => e.Description).HasMaxLength(500);
    entity.HasIndex(e => e.Name).IsUnique();
    entity.HasIndex(e => e.Slug).IsUnique();
    });

    // 配置 Post
    builder.Entity<Post>(entity =>
    {
    entity.Property(e => e.Title).IsRequired().HasMaxLength(200);
    entity.Property(e => e.Slug).IsRequired().HasMaxLength(200);
    entity.Property(e => e.Content).IsRequired();
    entity.Property(e => e.Excerpt).HasMaxLength(500);
    entity.Property(e => e.FeaturedImageUrl).HasMaxLength(500);
    entity.Property(e => e.Status).HasConversion<string>();

    entity.HasIndex(e => e.Slug).IsUnique();
    entity.HasIndex(e => e.AuthorId);
    entity.HasIndex(e => e.CategoryId);
    entity.HasIndex(e => e.Status);
    entity.HasIndex(e => e.PublishedAt);

    entity.HasOne(e => e.Author)
    .WithMany(u => u.Posts)
    .HasForeignKey(e => e.AuthorId)
    .OnDelete(DeleteBehavior.Restrict);

    entity.HasOne(e => e.Category)
    .WithMany(c => c.Posts)
    .HasForeignKey(e => e.CategoryId)
    .OnDelete(DeleteBehavior.Restrict);

    // 配置多对多关系(Post <-> Tag)
    entity.HasMany(e => e.Tags)
    .WithMany(t => t.Posts)
    .UsingEntity<Dictionary<string, object>>(
    "PostTags",
    j => j.HasOne<Tag>().WithMany().HasForeignKey("TagId"),
    j => j.HasOne<Post>().WithMany().HasForeignKey("PostId"));
    });

    // 配置 Tag
    builder.Entity<Tag>(entity =>
    {
    entity.Property(e => e.Name).IsRequired().HasMaxLength(50);
    entity.Property(e => e.Slug).IsRequired().HasMaxLength(50);
    entity.HasIndex(e => e.Name).IsUnique();
    entity.HasIndex(e => e.Slug).IsUnique();
    });

    // 配置 Comment
    builder.Entity<Comment>(entity =>
    {
    entity.Property(e => e.Content).IsRequired().HasMaxLength(1000);

    entity.HasIndex(e => e.PostId);
    entity.HasIndex(e => e.AuthorId);
    entity.HasIndex(e => e.ParentCommentId);
    entity.HasIndex(e => e.CreatedAt);

    entity.HasOne(e => e.Post)
    .WithMany(p => p.Comments)
    .HasForeignKey(e => e.PostId)
    .OnDelete(DeleteBehavior.Cascade);

    entity.HasOne(e => e.Author)
    .WithMany(u => u.Comments)
    .HasForeignKey(e => e.AuthorId)
    .OnDelete(DeleteBehavior.Restrict);

    entity.HasOne(e => e.ParentComment)
    .WithMany(c => c.Replies)
    .HasForeignKey(e => e.ParentCommentId)
    .OnDelete(DeleteBehavior.Restrict);
    });

    // 添加种子数据
    SeedData(builder);
    }

    private void SeedData(ModelBuilder builder)
    {
    // 添加默认分类
    builder.Entity<Category>().HasData(
    new Category { Id = 1, Name = "技术", Slug = "tech", Description = "技术相关文章" },
    new Category { Id = 2, Name = "生活", Slug = "life", Description = "生活随笔" },
    new Category { Id = 3, Name = "教程", Slug = "tutorial", Description = "教程文章" }
    );

    // 添加默认标签
    builder.Entity<Tag>().HasData(
    new Tag { Id = 1, Name = "C#", Slug = "csharp" },
    new Tag { Id = 2, Name = "ASP.NET Core", Slug = "aspnetcore" },
    new Tag { Id = 3, Name = "EF Core", Slug = "efcore" }
    );
    }
    }
    }

    Repository 接口定义

    Core/Interfaces/IRepository.cs:

    using System.Linq.Expressions;

    namespace BlogSystem.Core.Interfaces
    {
    public interface IRepository<T> where T : class
    {
    // 查询
    Task<T?> GetByIdAsync(int id);
    Task<IEnumerable<T>> GetAllAsync();
    Task<IEnumerable<T>> FindAsync(Expression<Func<T, bool>> predicate);

    // 分页查询
    Task<(IEnumerable<T> Items, int TotalCount)> GetPagedAsync(
    int page,
    int pageSize,
    Expression<Func<T, bool>>? filter = null,
    Func<IQueryable<T>, IOrderedQueryable<T>>? orderBy = null);

    // 添加
    Task<T> AddAsync(T entity);
    Task AddRangeAsync(IEnumerable<T> entities);

    // 更新
    void Update(T entity);
    void UpdateRange(IEnumerable<T> entities);

    // 删除
    void Remove(T entity);
    void RemoveRange(IEnumerable<T> entities);

    // 其他
    Task<bool> ExistsAsync(Expression<Func<T, bool>> predicate);
    Task<int> CountAsync(Expression<Func<T, bool>>? predicate = null);
    }
    }

    Infrastructure/Data/Repositories/Repository.cs:

    using System.Linq.Expressions;
    using BlogSystem.Core.Interfaces;
    using Microsoft.EntityFrameworkCore;

    namespace BlogSystem.Infrastructure.Data.Repositories
    {
    public class Repository<T> : IRepository<T> where T : class
    {
    protected readonly BlogDbContext _context;
    protected readonly DbSet<T> _dbSet;

    public Repository(BlogDbContext context)
    {
    _context = context;
    _dbSet = context.Set<T>();
    }

    public virtual async Task<T?> GetByIdAsync(int id)
    {
    return await _dbSet.FindAsync(id);
    }

    public virtual async Task<IEnumerable<T>> GetAllAsync()
    {
    return await _dbSet.ToListAsync();
    }

    public virtual async Task<IEnumerable<T>> FindAsync(Expression<Func<T, bool>> predicate)
    {
    return await _dbSet.Where(predicate).ToListAsync();
    }

    public virtual async Task<(IEnumerable<T> Items, int TotalCount)> GetPagedAsync(
    int page,
    int pageSize,
    Expression<Func<T, bool>>? filter = null,
    Func<IQueryable<T>, IOrderedQueryable<T>>? orderBy = null)
    {
    IQueryable<T> query = _dbSet;

    // 应用筛选条件
    if (filter != null)
    {
    query = query.Where(filter);
    }

    // 获取总数
    int totalCount = await query.CountAsync();

    // 应用排序
    if (orderBy != null)
    {
    query = orderBy(query);
    }

    // 应用分页
    var items = await query
    .Skip((page 1) * pageSize)
    .Take(pageSize)
    .ToListAsync();

    return (items, totalCount);
    }

    public virtual async Task<T> AddAsync(T entity)
    {
    await _dbSet.AddAsync(entity);
    return entity;
    }

    public virtual async Task AddRangeAsync(IEnumerable<T> entities)
    {
    await _dbSet.AddRangeAsync(entities);
    }

    public virtual void Update(T entity)
    {
    _dbSet.Update(entity);
    }

    public virtual void UpdateRange(IEnumerable<T> entities)
    {
    _dbSet.UpdateRange(entities);
    }

    public virtual void Remove(T entity)
    {
    _dbSet.Remove(entity);
    }

    public virtual void RemoveRange(IEnumerable<T> entities)
    {
    _dbSet.RemoveRange(entities);
    }

    public virtual async Task<bool> ExistsAsync(Expression<Func<T, bool>> predicate)
    {
    return await _dbSet.AnyAsync(predicate);
    }

    public virtual async Task<int> CountAsync(Expression<Func<T, bool>>? predicate = null)
    {
    if (predicate == null)
    {
    return await _dbSet.CountAsync();
    }

    return await _dbSet.CountAsync(predicate);
    }
    }
    }

    专用 Repository 接口

    Core/Interfaces/IPostRepository.cs:

    using BlogSystem.Core.Entities;

    namespace BlogSystem.Core.Interfaces
    {
    public interface IPostRepository : IRepository<Post>
    {
    Task<Post?> GetBySlugAsync(string slug);
    Task<Post?> GetByIdWithDetailsAsync(int id);
    Task<(IEnumerable<Post> Items, int TotalCount)> GetPublishedPostsAsync(
    int page,
    int pageSize,
    int? categoryId = null,
    string? tagSlug = null);
    Task<IEnumerable<Post>> GetPopularPostsAsync(int count);
    Task IncrementViewCountAsync(int postId);
    }
    }

    Core/Interfaces/ICommentRepository.cs:

    using BlogSystem.Core.Entities;

    namespace BlogSystem.Core.Interfaces
    {
    public interface ICommentRepository : IRepository<Comment>
    {
    Task<IEnumerable<Comment>> GetCommentsByPostIdAsync(int postId);
    Task<IEnumerable<Comment>> GetApprovedCommentsWithRepliesAsync(int postId);
    }
    }

    Core/Interfaces/ICategoryRepository.cs:

    using BlogSystem.Core.Entities;

    namespace BlogSystem.Core.Interfaces
    {
    public interface ICategoryRepository : IRepository<Category>
    {
    Task<Category?> GetBySlugAsync(string slug);
    Task<Category?> GetWithPostsAsync(int id);
    }
    }

    Core/Interfaces/ITagRepository.cs:

    using BlogSystem.Core.Entities;

    namespace BlogSystem.Core.Interfaces
    {
    public interface ITagRepository : IRepository<Tag>
    {
    Task<Tag?> GetBySlugAsync(string slug);
    Task<IEnumerable<Tag>> GetPopularTagsAsync(int count);
    Task<Tag> GetOrCreateAsync(string name, string slug);
    }
    }

    专用 Repository 实现

    Infrastructure/Data/Repositories/PostRepository.cs:

    using BlogSystem.Core.Entities;
    using BlogSystem.Core.Interfaces;
    using Microsoft.EntityFrameworkCore;

    namespace BlogSystem.Infrastructure.Data.Repositories
    {
    public class PostRepository : Repository<Post>, IPostRepository
    {
    public PostRepository(BlogDbContext context) : base(context)
    {
    }

    public async Task<Post?> GetBySlugAsync(string slug)
    {
    return await _dbSet
    .Include(p => p.Author)
    .Include(p => p.Category)
    .Include(p => p.Tags)
    .FirstOrDefaultAsync(p => p.Slug == slug);
    }

    public async Task<Post?> GetByIdWithDetailsAsync(int id)
    {
    return await _dbSet
    .Include(p => p.Author)
    .Include(p => p.Category)
    .Include(p => p.Tags)
    .Include(p => p.Comments.Where(c => c.IsApproved && c.ParentCommentId == null))
    .ThenInclude(c => c.Author)
    .Include(p => p.Comments)
    .ThenInclude(c => c.Replies.Where(r => r.IsApproved))
    .ThenInclude(r => r.Author)
    .FirstOrDefaultAsync(p => p.Id == id);
    }

    public async Task<(IEnumerable<Post> Items, int TotalCount)> GetPublishedPostsAsync(
    int page,
    int pageSize,
    int? categoryId = null,
    string? tagSlug = null)
    {
    IQueryable<Post> query = _dbSet
    .Include(p => p.Author)
    .Include(p => p.Category)
    .Include(p => p.Tags)
    .Where(p => p.Status == PostStatus.Published);

    // 按分类筛选
    if (categoryId.HasValue)
    {
    query = query.Where(p => p.CategoryId == categoryId.Value);
    }

    // 按标签筛选
    if (!string.IsNullOrEmpty(tagSlug))
    {
    query = query.Where(p => p.Tags.Any(t => t.Slug == tagSlug));
    }

    // 获取总数
    int totalCount = await query.CountAsync();

    // 应用分页和排序
    var items = await query
    .OrderByDescending(p => p.PublishedAt)
    .Skip((page 1) * pageSize)
    .Take(pageSize)
    .ToListAsync();

    return (items, totalCount);
    }

    public async Task<IEnumerable<Post>> GetPopularPostsAsync(int count)
    {
    return await _dbSet
    .Include(p => p.Author)
    .Include(p => p.Category)
    .Where(p => p.Status == PostStatus.Published)
    .OrderByDescending(p => p.ViewCount)
    .Take(count)
    .ToListAsync();
    }

    public async Task IncrementViewCountAsync(int postId)
    {
    var post = await _dbSet.FindAsync(postId);
    if (post != null)
    {
    post.ViewCount++;
    _dbSet.Update(post);
    }
    }
    }
    }

    Infrastructure/Data/Repositories/CommentRepository.cs:

    using BlogSystem.Core.Entities;
    using BlogSystem.Core.Interfaces;
    using Microsoft.EntityFrameworkCore;

    namespace BlogSystem.Infrastructure.Data.Repositories
    {
    public class CommentRepository : Repository<Comment>, ICommentRepository
    {
    public CommentRepository(BlogDbContext context) : base(context)
    {
    }

    public async Task<IEnumerable<Comment>> GetCommentsByPostIdAsync(int postId)
    {
    return await _dbSet
    .Include(c => c.Author)
    .Where(c => c.PostId == postId)
    .OrderByDescending(c => c.CreatedAt)
    .ToListAsync();
    }

    public async Task<IEnumerable<Comment>> GetApprovedCommentsWithRepliesAsync(int postId)
    {
    return await _dbSet
    .Include(c => c.Author)
    .Include(c => c.Replies.Where(r => r.IsApproved))
    .ThenInclude(r => r.Author)
    .Where(c => c.PostId == postId && c.IsApproved && c.ParentCommentId == null)
    .OrderBy(c => c.CreatedAt)
    .ToListAsync();
    }
    }
    }

    Infrastructure/Data/Repositories/CategoryRepository.cs:

    using BlogSystem.Core.Entities;
    using BlogSystem.Core.Interfaces;
    using Microsoft.EntityFrameworkCore;

    namespace BlogSystem.Infrastructure.Data.Repositories
    {
    public class CategoryRepository : Repository<Category>, ICategoryRepository
    {
    public CategoryRepository(BlogDbContext context) : base(context)
    {
    }

    public async Task<Category?> GetBySlugAsync(string slug)
    {
    return await _dbSet
    .FirstOrDefaultAsync(c => c.Slug == slug);
    }

    public async Task<Category?> GetWithPostsAsync(int id)
    {
    return await _dbSet
    .Include(c => c.Posts.Where(p => p.Status == PostStatus.Published))
    .ThenInclude(p => p.Author)
    .FirstOrDefaultAsync(c => c.Id == id);
    }
    }
    }

    Infrastructure/Data/Repositories/TagRepository.cs:

    using BlogSystem.Core.Entities;
    using BlogSystem.Core.Interfaces;
    using Microsoft.EntityFrameworkCore;

    namespace BlogSystem.Infrastructure.Data.Repositories
    {
    public class TagRepository : Repository<Tag>, ITagRepository
    {
    public TagRepository(BlogDbContext context) : base(context)
    {
    }

    public async Task<Tag?> GetBySlugAsync(string slug)
    {
    return await _dbSet
    .FirstOrDefaultAsync(t => t.Slug == slug);
    }

    public async Task<IEnumerable<Tag>> GetPopularTagsAsync(int count)
    {
    return await _dbSet
    .Include(t => t.Posts)
    .OrderByDescending(t => t.Posts.Count)
    .Take(count)
    .ToListAsync();
    }

    public async Task<Tag> GetOrCreateAsync(string name, string slug)
    {
    var tag = await GetBySlugAsync(slug);

    if (tag == null)
    {
    tag = new Tag
    {
    Name = name,
    Slug = slug
    };
    await AddAsync(tag);
    }

    return tag;
    }
    }
    }

    Unit of Work 模式

    Core/Interfaces/IUnitOfWork.cs:

    namespace BlogSystem.Core.Interfaces
    {
    public interface IUnitOfWork : IDisposable
    {
    IPostRepository Posts { get; }
    ICommentRepository Comments { get; }
    ICategoryRepository Categories { get; }
    ITagRepository Tags { get; }

    Task<int> SaveChangesAsync();
    Task BeginTransactionAsync();
    Task CommitTransactionAsync();
    Task RollbackTransactionAsync();
    }
    }

    Infrastructure/Data/UnitOfWork.cs:

    using BlogSystem.Core.Interfaces;
    using BlogSystem.Infrastructure.Data.Repositories;
    using Microsoft.EntityFrameworkCore.Storage;

    namespace BlogSystem.Infrastructure.Data
    {
    public class UnitOfWork : IUnitOfWork
    {
    private readonly BlogDbContext _context;
    private IDbContextTransaction? _transaction;

    // 仓储实例
    private IPostRepository? _posts;
    private ICommentRepository? _comments;
    private ICategoryRepository? _categories;
    private ITagRepository? _tags;

    public UnitOfWork(BlogDbContext context)
    {
    _context = context;
    }

    // 延迟初始化仓储
    public IPostRepository Posts
    {
    get
    {
    _posts ??= new PostRepository(_context);
    return _posts;
    }
    }

    public ICommentRepository Comments
    {
    get
    {
    _comments ??= new CommentRepository(_context);
    return _comments;
    }
    }

    public ICategoryRepository Categories
    {
    get
    {
    _categories ??= new CategoryRepository(_context);
    return _categories;
    }
    }

    public ITagRepository Tags
    {
    get
    {
    _tags ??= new TagRepository(_context);
    return _tags;
    }
    }

    public async Task<

    天赐范式第117天:N reset 的灵敏度分析与延髓节律的不可约个体差异

    master阅读(43)

    天赐范式第117天第二篇:Nreset\\mathcal{N}_{\\text{reset}}Nreset的灵敏度分析与延髓节律的不可约个体差异

    免责声明:本文纯属理论推演与生活科学探讨,不构成任何医疗、诊断或治疗建议。文中涉及的神经生理学参数均为公开文献估计值或明确标注为待校准的理论启发式参数。读者如有持续性打嗝或其他健康问题,应咨询专业医疗人员。

    版本:v1.1
    日期:2026-07-28
    关联:第117天第一篇《τ=rollback的神经节律——从打嗝反射看延髓信号的重写》v1.2、第115天第二篇《CΨ\\mathcal{C}_{\\Psi}CΨ灵敏度分析与Moravec粒度下界》v1.1、第108天《方法论的自证》v1.4
    新公式:2个(Smargin\\mathcal{S}_{\\text{margin}}Smargin 临界裕度系数、Ibio\\mathcal{I}_{\\text{bio}}Ibio 不可约个体差异下界)+ 1个模型假设参数(κ\\kappaκ 延髓反射中枢响应效率)
    一句话:第117天第一篇给出了Nreset=0.73\\mathcal{N}_{\\text{reset}} = 0.73Nreset=0.73(对瓶吹)的量化结论,但0.73刚过0.7阈值——裕度只有0.03。第二篇追问:这个裕度够不够?哪一项信号是"不能没有"的?存在一类人,任何外部动作都无法达到延髓节律重置阈值——这个极限在哪?

    v1.0→v1.1 修正记录:
    H1 引入κ\\kappaκ(延髓反射中枢响应效率)替代直接使用ηbio\\eta_{\\text{bio}}ηbio作为信号衰减因子,明确κ≈1−ηbio\\kappa \\approx 1 – \\eta_{\\text{bio}}κ1ηbio为模型假设,消除判定层与监察层混用;
    H2 删除"不是巧合,是模型自洽"的循环论证,改为"参数设定导致的数值巧合";
    H3 §3.2 CCO2\\mathcal{C}_{\\text{CO2}}CCO2扫描表格中C=0.2C=0.2C=0.2改为C=0.22C=0.22C=0.22Nreset=0.695N_{\\text{reset}}=0.695Nreset=0.695判L2;
    H4 R状态机补充EMERGENCY实例化说明,与117-1衔接;
    H5 P-117-2-1成功率区间改为定性表述;
    H6 P-117-2-4"成功率接近零"改为"成功率显著低于对瓶吹且极不稳定";
    M1 O2-5补充ηbio\\eta_{\\text{bio}}ηbio语义扩展标注;
    M2 §7临界等值面距离0.050修正为0.055;
    O1 DRR-R-2补充κ\\kappaκ与膈神经传导速度的生理学锚定(45-75 m/s→v<30 m/s对应κ下降20%-40%);
    O2 §5.1补注±0.2\\pm 0.2±0.2为经验性波动估计非统计标准差;
    O3 P-117-2-2实验设计补充"控制吞咽频率一致",解决屏气时无法连续吞咽的可比性问题。

    原则声明:本文基于v7.0弹药库(129算子/42+公式),引入2个新公式(Smargin\\mathcal{S}_{\\text{margin}}Smargin 临界裕度系数、Ibio\\mathcal{I}_{\\text{bio}}Ibio 不可约个体差异下界)和1个模型假设参数(κ\\kappaκ 延髓反射中枢响应效率),为Σ(#12)不确定性算子与Λ(#10)偏离预警算子在神经反射节律领域的灵敏度分析实例化。每个参数有物理来源或明确标注为待校准。

    来源说明:本文算子全部逐条对照v1.0最小化技术特征(精华版129算子),确认覆盖当前推演所需。其中ξ(#0)/Ξ(#1)/Θ(#2)/Γ(#9)/Σ(#12)/Λ(#10)/τ(#11)/Φ(#13)/R(#59)为判定层与调节层算子,MΣ(#29)/Con(#32)/ρ(#30)/δ(#31)/λ(#33)/C²(#34)为监察层算子。ξ/Ξ严格分离(ξ=#0初始化,Ξ=#1锚定),Σ=#12不确定性算子,MΣ=#29元不确定性算子。Smargin\\mathcal{S}_{\\text{margin}}Smargin为Σ(#12)与Λ(#10)联合实例化,Ibio\\mathcal{I}_{\\text{bio}}Ibio为Σ(#12)与MΣ(#29)联合实例化,均不新增算子。κ\\kappaκ为本文引入的模型假设参数,不属于算子编号体系;117-1中ηbio\\eta_{\\text{bio}}ηbio是Σ的监察层不确定性分量,本文不直接使用其作为衰减因子,通过假设κ≈1−ηbio\\kappa \\approx 1 – \\eta_{\\text{bio}}κ1ηbio建立两者联系。


    〇、引言:0.73够不够

    第117天第一篇的结论是:对瓶吹矿泉水的Nreset≈0.73\\mathcal{N}_{\\text{reset}} \\approx 0.73Nreset0.73,刚过0.7阈值,判L1-pass。小口喝Nreset≈0.22\\mathcal{N}_{\\text{reset}} \\approx 0.22Nreset0.22,判L3。

    但0.73只比阈值高0.03。这意味着:

  • 裕度极薄——参数稍有波动就可能跌破阈值
  • 哪一项最关键?——Sswallow\\mathcal{S}_{\\text{swallow}}Sswallow/Estretch\\mathcal{E}_{\\text{stretch}}Estretch/CCO2\\mathcal{C}_{\\text{CO2}}CCO2三项中,去掉哪一项会跌破阈值?去掉哪一项还能撑住?
  • 个体差异的极限——Σ=1.00提示Nreset\\mathcal{N}_{\\text{reset}}Nreset可能偏离±0.2\\pm 0.2±0.2。如果某个体的响应效率κ\\kappaκ使Nreset\\mathcal{N}_{\\text{reset}}Nreset的实际上界低于0.7,则任何吞咽动作都无法止嗝。这个极限在哪?
  • 第二篇用算子流回答这三个问题。与115天双篇的节奏对称:115-1建CΨ\\mathcal{C}_{\\Psi}CΨ模型,115-2做灵敏度与Moravec粒度下界;117-1建Nreset\\mathcal{N}_{\\text{reset}}Nreset模型,117-2做灵敏度与不可约个体差异下界。


    一、判定层与监察层各归其位

    算子编号在本文中的角色
    ξ #0 初始化:Domain=神经反射灵敏度场\\text{Domain} = \\text{神经反射灵敏度场}Domain=神经反射灵敏度场GridN={(Sswallow,Estretch,CCO2)∈[0,1]3}\\text{Grid}_N = \\{(\\mathcal{S}_{\\text{swallow}}, \\mathcal{E}_{\\text{stretch}}, \\mathcal{C}_{\\text{CO2}}) \\in [0,1]^3\\}GridN={(Sswallow,Estretch,CCO2)[0,1]3}
    Ξ #1 锚定:Nreset≥0.7\\mathcal{N}_{\\text{reset}} \\ge 0.7Nreset0.7为L1-pass阈值,Nreset<0.3\\mathcal{N}_{\\text{reset}} < 0.3Nreset<0.3为L3失败阈值
    Θ #2 感知:单因子扫描中每项信号从0到1变化时Nreset\\mathcal{N}_{\\text{reset}}Nreset的响应曲线
    Γ #9 度量:Nreset\\mathcal{N}_{\\text{reset}}Nreset的偏导数∂Nreset/∂xi\\partial \\mathcal{N}_{\\text{reset}} / \\partial x_iNreset/xi度量单因子灵敏度
    Σ #12 不确定性:Σ≈1.00\\Sigma \\approx 1.00Σ1.00(继承117-1),Nreset\\mathcal{N}_{\\text{reset}}Nreset偏离±0.2\\pm 0.2±0.2
    Λ #10 偏离预警:临界裕度Smargin\\mathcal{S}_{\\text{margin}}Smargin度量Nreset\\mathcal{N}_{\\text{reset}}Nreset与阈值的距离
    τ #11 回滚:Nreset<0.7\\mathcal{N}_{\\text{reset}} < 0.7Nreset<0.7时τ=rollback,节律未重置
    Φ #13 门控:L1/L2/L3三层语义不变
    R #59 状态机:{PASSIVE, ACTIVE, EMERGENCY}。本文补充实例化:当Nreset\\mathcal{N}_{\\text{reset}}Nreset持续低于0.3(L3失败区)且Σ\\SigmaΣ顶格时,R进入EMERGENCY,提示该个体可能超出C5生理性反射边界(117-1仅定义PASSIVE↔ACTIVE,本文扩展EMERGENCY)
    #29 元不确定性:Ibio\\mathcal{I}_{\\text{bio}}Ibio量化个体差异的不可约下界
    δ #31 领域饱和:N=43(灵敏度分析为第43个领域实例化),δ=1−e−43/35≈0.71\\delta = 1 – e^{-43/35} \\approx 0.71δ=1e43/350.71
    #34 表达力:不触发

    二、特征向量判定(TDP-CP审查)

    Checkpoint A:TDP(特征向量)

    特征物理含义范围
    x1=Sswallowx_1 = \\mathcal{S}_{\\text{swallow}}x1=Sswallow 归一化吞咽密度 [0,1][0, 1][0,1]
    x2=Estretchx_2 = \\mathcal{E}_{\\text{stretch}}x2=Estretch 食管壁牵张感受器激活度 [0,1][0, 1][0,1]
    x3=CCO2x_3 = \\mathcal{C}_{\\text{CO2}}x3=CCO2 CO₂抢权系数(累积呼吸暂停占比) [0,1][0, 1][0,1]

    Checkpoint B:CP(DAG拓扑)

    ξ→初始化Ξ→锚定阈值Θ→扫描感知Γ→偏导度量Σ→不确定性Λ→裕度预警τ→判定Φ→门控R\\xi \\xrightarrow{\\text{初始化}} \\Xi \\xrightarrow{\\text{锚定阈值}} \\Theta \\xrightarrow{\\text{扫描感知}} \\Gamma \\xrightarrow{\\text{偏导度量}} \\Sigma \\xrightarrow{\\text{不确定性}} \\Lambda \\xrightarrow{\\text{裕度预警}} \\tau \\xrightarrow{\\text{判定}} \\Phi \\xrightarrow{\\text{门控}} Rξ初始化Ξ锚定阈值Θ扫描感知Γ偏导度量Σ不确定性Λ裕度预警τ判定Φ门控R

    MΣ在Σ之后并行运行,提供个体差异的不可约下界。

    Checkpoint C:接口签名

    输入域:S0={(α,β,γ),(x1对瓶吹,x2对瓶吹,x3对瓶吹),(x1小口,x2小口,x3小口)}S_0 = \\{(\\alpha, \\beta, \\gamma), (x_1^{\\text{对瓶吹}}, x_2^{\\text{对瓶吹}}, x_3^{\\text{对瓶吹}}), (x_1^{\\text{小口}}, x_2^{\\text{小口}}, x_3^{\\text{小口}})\\}S0={(α,β,γ),(x1对瓶吹,x2对瓶吹,x3对瓶吹),(x1小口,x2小口,x3小口)}

    继承117-1参数:α=0.45\\alpha = 0.45α=0.45, β=0.3\\beta = 0.3β=0.3, γ=0.25\\gamma = 0.25γ=0.25

    输出域:GridN={Smargin,Ibio,Ψ指令}\\text{Grid}_N = \\{\\mathcal{S}_{\\text{margin}}, \\mathcal{I}_{\\text{bio}}, \\Psi_{\\text{指令}}\\}GridN={Smargin,Ibio,Ψ指令}


    三、单因子灵敏度分析

    3.1 偏导数与边际贡献

    Nreset\\mathcal{N}_{\\text{reset}}Nreset是三项的线性组合,偏导数即为权重:

    ∂Nreset∂x1=α=0.45,∂Nreset∂x2=β=0.3,∂Nreset∂x3=γ=0.25\\frac{\\partial \\mathcal{N}_{\\text{reset}}}{\\partial x_1} = \\alpha = 0.45, \\quad \\frac{\\partial \\mathcal{N}_{\\text{reset}}}{\\partial x_2} = \\beta = 0.3, \\quad \\frac{\\partial \\mathcal{N}_{\\text{reset}}}{\\partial x_3} = \\gamma = 0.25x1Nreset=α=0.45,x2Nreset=β=0.3,x3Nreset=γ=0.25

    信号权重边际贡献(xix_ixi从0→1的ΔNreset\\Delta \\mathcal{N}_{\\text{reset}}ΔNreset)排序
    Sswallow\\mathcal{S}_{\\text{swallow}}Sswallow α=0.45\\alpha = 0.45α=0.45 0.45 1(主导)
    Estretch\\mathcal{E}_{\\text{stretch}}Estretch β=0.3\\beta = 0.3β=0.3 0.30 2
    CCO2\\mathcal{C}_{\\text{CO2}}CCO2 γ=0.25\\gamma = 0.25γ=0.25 0.25 3(辅助)

    关键发现:Sswallow\\mathcal{S}_{\\text{swallow}}Sswallow的边际贡献最大(0.45),但没有任何单项能独立达到0.7阈值。α×1.0=0.45<0.7\\alpha \\times 1.0 = 0.45 < 0.7α×1.0=0.45<0.7——即使吞咽密度拉满,没有其余两项辅助也过不了阈值。这是"三重信号链"设计的必然结果:α+β+γ=1\\alpha + \\beta + \\gamma = 1α+β+γ=1且最大权重<0.7< 0.7<0.7

    3.2 单因子扫描

    固定其余两项在对瓶吹基线值,单独扫描一项:

    扫描Sswallow\\mathcal{S}_{\\text{swallow}}Sswallow(固定Estretch=0.8\\mathcal{E}_{\\text{stretch}}=0.8Estretch=0.8, CCO2=0.35\\mathcal{C}_{\\text{CO2}}=0.35CCO2=0.35):

    Nreset=0.45⋅Sswallow+0.3×0.8+0.25×0.35=0.45⋅Sswallow+0.3275\\mathcal{N}_{\\text{reset}} = 0.45 \\cdot \\mathcal{S}_{\\text{swallow}} + 0.3 \\times 0.8 + 0.25 \\times 0.35 = 0.45 \\cdot \\mathcal{S}_{\\text{swallow}} + 0.3275Nreset=0.45Sswallow+0.3×0.8+0.25×0.35=0.45Sswallow+0.3275

    Sswallow\\mathcal{S}_{\\text{swallow}}SswallowNreset\\mathcal{N}_{\\text{reset}}Nreset判定
    0.0 0.33 L2
    0.2 0.42 L2
    0.4 0.51 L2
    0.6 0.60 L2
    0.8 0.69 L2(差0.01!)
    0.9 0.73 L1-pass
    1.0 0.78 L1-pass

    Sswallow\\mathcal{S}_{\\text{swallow}}Sswallow的L1-pass临界值:x1∗=(0.7−0.3275)/0.45=0.828x_1^* = (0.7 – 0.3275) / 0.45 = 0.828x1=(0.70.3275)/0.45=0.828。即吞咽密度需超过∼0.83\\sim 0.830.83(对应n˙≈4.1\\dot{n} \\approx 4.1n˙4.1/s)才能跨过阈值。

    扫描Estretch\\mathcal{E}_{\\text{stretch}}Estretch(固定Sswallow=0.9\\mathcal{S}_{\\text{swallow}}=0.9Sswallow=0.9, CCO2=0.35\\mathcal{C}_{\\text{CO2}}=0.35CCO2=0.35):

    Nreset=0.45×0.9+0.3⋅Estretch+0.25×0.35=0.4925+0.3⋅Estretch\\mathcal{N}_{\\text{reset}} = 0.45 \\times 0.9 + 0.3 \\cdot \\mathcal{E}_{\\text{stretch}} + 0.25 \\times 0.35 = 0.4925 + 0.3 \\cdot \\mathcal{E}_{\\text{stretch}}Nreset=0.45×0.9+0.3Estretch+0.25×0.35=0.4925+0.3Estretch

    Estretch\\mathcal{E}_{\\text{stretch}}EstretchNreset\\mathcal{N}_{\\text{reset}}Nreset判定
    0.0 0.49 L2
    0.3 0.58 L2
    0.5 0.64 L2
    0.7 0.70 L1-pass(刚过)
    0.8 0.73 L1-pass
    1.0 0.79 L1-pass

    Estretch\\mathcal{E}_{\\text{stretch}}Estretch的L1-pass临界值:x2∗=(0.7−0.4925)/0.3=0.692x_2^* = (0.7 – 0.4925) / 0.3 = 0.692x2=(0.70.4925)/0.3=0.692。即牵张信号需超过∼0.69\\sim 0.690.69才能跨过阈值。

    扫描CCO2\\mathcal{C}_{\\text{CO2}}CCO2(固定Sswallow=0.9\\mathcal{S}_{\\text{swallow}}=0.9Sswallow=0.9, Estretch=0.8\\mathcal{E}_{\\text{stretch}}=0.8Estretch=0.8):

    Nreset=0.45×0.9+0.3×0.8+0.25⋅CCO2=0.645+0.25⋅CCO2\\mathcal{N}_{\\text{reset}} = 0.45 \\times 0.9 + 0.3 \\times 0.8 + 0.25 \\cdot \\mathcal{C}_{\\text{CO2}} = 0.645 + 0.25 \\cdot \\mathcal{C}_{\\text{CO2}}Nreset=0.45×0.9+0.3×0.8+0.25CCO2=0.645+0.25CCO2

    CCO2\\mathcal{C}_{\\text{CO2}}CCO2Nreset\\mathcal{N}_{\\text{reset}}Nreset判定
    0.0 0.65 L2
    0.1 0.67 L2
    0.22 0.70 L1-pass(刚过)
    0.35 0.73 L1-pass
    0.5 0.77 L1-pass
    1.0 0.90 L1-pass

    CCO2\\mathcal{C}_{\\text{CO2}}CCO2的L1-pass临界值:x3∗=(0.7−0.645)/0.25=0.22x_3^* = (0.7 – 0.645) / 0.25 = 0.22x3=(0.70.645)/0.25=0.22。即CO₂抢权只需超过∼0.22\\sim 0.220.22就能跨过阈值——但前提是另外两项已在对瓶吹基线。

    3.3 临界值汇总

    信号基线值(对瓶吹)L1-pass临界值裕度(基线−临界)
    Sswallow\\mathcal{S}_{\\text{swallow}}Sswallow 0.90 0.83 0.07
    Estretch\\mathcal{E}_{\\text{stretch}}Estretch 0.80 0.69 0.11
    CCO2\\mathcal{C}_{\\text{CO2}}CCO2 0.35 0.22 0.13

    Sswallow\\mathcal{S}_{\\text{swallow}}Sswallow的裕度最薄(0.07)——吞咽密度是对瓶吹止嗝的"瓶颈信号"。如果吞咽密度从0.90降到0.82,即使其余两项保持基线值,Nreset\\mathcal{N}_{\\text{reset}}Nreset也会跌破0.7。


    四、双因子替代分析

    核心问题:去掉哪一项还能撑住L1-pass?

    固定第三项为0(模拟该信号完全缺失),看另外两项能否达到0.7:

    4.1 去CCO2\\mathcal{C}_{\\text{CO2}}CCO2CCO2=0\\mathcal{C}_{\\text{CO2}} = 0CCO2=0

    Nreset=0.45⋅Sswallow+0.3⋅Estretch\\mathcal{N}_{\\text{reset}} = 0.45 \\cdot \\mathcal{S}_{\\text{swallow}} + 0.3 \\cdot \\mathcal{E}_{\\text{stretch}}Nreset=0.45Sswallow+0.3Estretch

    最大值(x1=x2=1x_1 = x_2 = 1x1=x2=1):0.45+0.3=0.75>0.70.45 + 0.3 = 0.75 > 0.70.45+0.3=0.75>0.7

    → 可以没有CO₂——只要吞咽密度和牵张信号足够强。这验证了117-1中"正常呼吸下对瓶吹"退化行为的正确性。

    4.2 去Estretch\\mathcal{E}_{\\text{stretch}}EstretchEstretch=0\\mathcal{E}_{\\text{stretch}} = 0Estretch=0

    Nreset=0.45⋅Sswallow+0.25⋅CCO2\\mathcal{N}_{\\text{reset}} = 0.45 \\cdot \\mathcal{S}_{\\text{swallow}} + 0.25 \\cdot \\mathcal{C}_{\\text{CO2}}Nreset=0.45Sswallow+0.25CCO2

    最大值(x1=x3=1x_1 = x_3 = 1x1=x3=1):0.45+0.25=0.70=0.70.45 + 0.25 = 0.70 = 0.70.45+0.25=0.70=0.7 ✓(刚好)

    → 可以没有牵张信号——但仅在吞咽密度和CO₂均拉满时刚好达到阈值。对应"干咽+屏气"的极端组合,裕度为零,实践中几乎不可行。

    4.3 去Sswallow\\mathcal{S}_{\\text{swallow}}SswallowSswallow=0\\mathcal{S}_{\\text{swallow}} = 0Sswallow=0

    Nreset=0.3⋅Estretch+0.25⋅CCO2\\mathcal{N}_{\\text{reset}} = 0.3 \\cdot \\mathcal{E}_{\\text{stretch}} + 0.25 \\cdot \\mathcal{C}_{\\text{CO2}}Nreset=0.3Estretch+0.25CCO2

    最大值(x2=x3=1x_2 = x_3 = 1x2=x3=1):0.3+0.25=0.55<0.70.3 + 0.25 = 0.55 < 0.70.3+0.25=0.55<0.7

    → 不能没有吞咽密度——即使牵张信号和CO₂均拉满,Nreset\\mathcal{N}_{\\text{reset}}Nreset最大只有0.55,远低于0.7。这从数学上证明了"吞咽密度是不可替代的主导信号"。

    4.4 替代矩阵

    去掉的信号Nresetmax⁡\\mathcal{N}_{\\text{reset}}^{\\max}Nresetmax能否达到L1-pass?实践含义
    CCO2\\mathcal{C}_{\\text{CO2}}CCO2 0.75 ✅ 可以 正常呼吸下对瓶吹仍有效
    Estretch\\mathcal{E}_{\\text{stretch}}Estretch 0.70 ⚠️ 刚好 干咽+屏气,裕度为零,不可行
    Sswallow\\mathcal{S}_{\\text{swallow}}Sswallow 0.55 ❌ 不行 不吞咽就无法止嗝(仅靠牵张+CO₂)

    关键结论:三重信号链中,Sswallow\\mathcal{S}_{\\text{swallow}}Sswallow是不可替代项,Estretch\\mathcal{E}_{\\text{stretch}}Estretch是强辅助项,CCO2\\mathcal{C}_{\\text{CO2}}CCO2是可省略项。这与117-1的权重设计(α>β>γ\\alpha > \\beta > \\gammaα>β>γ)一致——权重最大的项恰好是不可替代项。


    五、新公式

    5.1 Smargin\\mathcal{S}_{\\text{margin}}Smargin 临界裕度系数

    Smargin=Nreset−Nthreshold\\boxed{\\mathcal{S}_{\\text{margin}} = \\mathcal{N}_{\\text{reset}} – \\mathcal{N}_{\\text{threshold}}}Smargin=NresetNthreshold

    其中Nthreshold=0.7\\mathcal{N}_{\\text{threshold}} = 0.7Nthreshold=0.7为L1-pass阈值。

    符号含义单位/来源
    Nreset\\mathcal{N}_{\\text{reset}}Nreset 神经节律重置系数(继承117-1) 无量纲
    Nthreshold\\mathcal{N}_{\\text{threshold}}Nthreshold L1-pass阈值,Nthreshold=0.7\\mathcal{N}_{\\text{threshold}} = 0.7Nthreshold=0.7 无量纲
    Smargin\\mathcal{S}_{\\text{margin}}Smargin 临界裕度:Smargin>0\\mathcal{S}_{\\text{margin}} > 0Smargin>0为pass,Smargin<0\\mathcal{S}_{\\text{margin}} < 0Smargin<0为rollback 无量纲

    物理含义:Smargin\\mathcal{S}_{\\text{margin}}Smargin度量"对瓶吹止嗝的安全余量"。Smargin\\mathcal{S}_{\\text{margin}}Smargin越大,参数波动空间越大;Smargin→0\\mathcal{S}_{\\text{margin}} \\to 0Smargin0意味着"刚好过线",任何微小扰动都可能翻车。

    当前估计:

    场景Nreset\\mathcal{N}_{\\text{reset}}NresetSmargin\\mathcal{S}_{\\text{margin}}Smargin含义
    对瓶吹 0.73 +0.03 刚过阈值,裕度极薄
    小口喝 0.22 −0.48 远低于阈值
    对瓶吹(Σ\\SigmaΣ修正下界) 0.73−0.2=0.530.73 – 0.2 = 0.530.730.2=0.53 −0.17 个体差异使实际值可能跌破阈值

    → 核心发现:对瓶吹的名义裕度仅+0.03,考虑Σ=1.00的±0.2\\pm 0.2±0.2偏离后,实际裕度可能为负。这意味着对瓶吹并非对所有人有效——约有一部分个体的Nreset\\mathcal{N}_{\\text{reset}}Nreset实际值低于0.7。(此处±0.2\\pm 0.2±0.2为经验性波动估计,非Σ的统计标准差,待实验校准。)

    退化行为:

    • Smargin≫0\\mathcal{S}_{\\text{margin}} \\gg 0Smargin0:止嗝方法对参数波动鲁棒
    • Smargin→0+\\mathcal{S}_{\\text{margin}} \\to 0^+Smargin0+:方法有效但脆弱,参数稍有波动即翻车
    • Smargin<0\\mathcal{S}_{\\text{margin}} < 0Smargin<0:方法对该个体无效

    5.2 Ibio\\mathcal{I}_{\\text{bio}}Ibio 不可约个体差异下界

    模型假设:引入参数κ∈[0,1]\\kappa \\in [0, 1]κ[0,1](延髓反射中枢响应效率),表示个体将外部干扰信号转化为实际节律重置效果的效率。κ=1\\kappa = 1κ=1为理想个体(信号无衰减),κ=0\\kappa = 0κ=0为完全无响应。实际Nreset\\mathcal{N}_{\\text{reset}}Nreset为名义值乘以κ\\kappaκ

    Nresetactual=κ⋅Nresetnominal\\mathcal{N}_{\\text{reset}}^{\\text{actual}} = \\kappa \\cdot \\mathcal{N}_{\\text{reset}}^{\\text{nominal}}Nresetactual=κNresetnominal

    当外部信号拉满(x1=x2=x3=1x_1 = x_2 = x_3 = 1<sp

    Harness Engineering 从理论到实战:基于 Spring AI Alibaba 的完整实现指南

    master阅读(37)

    Harness Engineering 从理论到实战:基于 Spring AI Alibaba 的完整实现指南

    一份涵盖七层架构、完整代码和深度解析的实践手册


    Harness Engineering 实战:用 Spring AI Alibaba 构建可控的 AI 智能体:

    https://blog.csdn.net/BADAO_LIUMANG_QIZHI/article/details/162488035

    基于上述流程。

    目录

  • 什么是 Harness Engineering?
  • Harness Engineering 的核心架构
  • 环境准备与项目搭建
  • 完整代码实现
  • 七层架构深度解析
  • 运行与测试
  • 常见问题与解决方案
  • 总结与进阶方向

  • 一、什么是 Harness Engineering?

    Harness Engineering(驾驭工程)是 2026 年 AI 工程化领域兴起的一种新范式。它的核心思想是:将重心从优化 AI 模型本身,转移到为 AI 智能体(Agent)构建一个可靠、可控、可维护的运行环境。

    核心公式

    Agent = Model + Harness

    如果把大语言模型比作一匹力量强大但难以预测的野马,那么 Harness(马具) 就是套在它身上用来引导和控制的整套装备。

    Harness Engineering 的核心组件

    组件类型作用类比
    Rules(规则) 软约束 声明式规范,告诉 Agent “什么不能做” 交通规则
    Skills(技能) 半硬约束 步骤化操作手册,告诉 Agent “具体怎么做” 操作说明书
    Gate(门禁) 硬约束 强制校验输出,“不通过就拦截” 安检闸机
    State(状态) 上下文管理 记录 Agent 的进度和上下文 记事本
    Instructions(指令) 行为指导 告诉 Agent 做什么、按什么顺序做 工作清单
    Verification(验证) 质量保证 只有通过测试才算任务完成 质量检查员

    注:

    博客:

    https://blog.csdn.net/badao_liumang_qizhi

    二、Harness Engineering 的核心架构

    七层架构总览

    Harness Engineering 的七层架构是一个从内到外、层层递进的完整体系,其核心目标是将 AI 的不可控性通过工程化手段转化为确定性。

    ┌─────────────────────────────────────────────────────────────┐
    │ 7. 评估与反馈层 (Evaluation & Feedback) │ ← 持续优化闭环
    ├─────────────────────────────────────────────────────────────┤
    │ 6. 多 Agent 架构层 (Multi-Agent Architecture) │ ← 团队协作
    ├─────────────────────────────────────────────────────────────┤
    │ 5. 约束与防护层 (Constraints & Guardrails) │ ← 安全护栏
    ├─────────────────────────────────────────────────────────────┤
    │ 4. 上下文工程层 (Context Engineering) │ ← 记忆与知识
    ├─────────────────────────────────────────────────────────────┤
    │ 3. 项目搭建层 (Project Setup) │ ← 标准化地基
    ├─────────────────────────────────────────────────────────────┤
    │ 2. 工具编排层 (Tool Orchestration) │ ← 外部交互
    ├─────────────────────────────────────────────────────────────┤
    │ 1. 执行主循环层 (Execution Loop) │ ← 大脑与调度
    └─────────────────────────────────────────────────────────────┘

    各层职责速览

    层级核心职责简单理解
    1. 执行主循环 Agent 运行的"大脑"与调度核心 负责思考、决策、执行和反思的循环
    2. 工具编排 Agent 与外部世界交互的"手脚" 安全、可控地调用外部 API 或工具
    3. 项目搭建 标准化项目的"地基" 统一项目结构、依赖和规范
    4. 上下文工程 Agent 的"记忆"与"知识库" 管理、优化注入给模型的信息
    5. 约束与防护 Agent 的"安全护栏" 硬性校验输入输出,防止失控
    6. 多 Agent 架构 团队协作的"管理模式" 让多个专业 Agent 协同工作
    7. 评估与反馈 持续优化的"学习闭环" 通过测试和反馈不断改进系统

    三、环境准备与项目搭建

    3.1 Windows 本地环境要求

    软件版本要求下载地址
    JDK 17 或更高 Oracle JDK
    Maven 3.8+ https://maven.apache.org/download.cgi
    IntelliJ IDEA 2023.3+ https://www.jetbrains.com/idea/
    curl 任意版本 Windows 10/11 内置
    DashScope API Key 阿里云百炼平台获取 https://bailian.console.aliyun.com/

    验证安装:在命令行执行 java -version 和 mvn -version 确认环境变量配置正确。

    3.2 项目结构

    spring-ai-harness-demo/
    ├── pom.xml
    ├── src/main/
    │ ├── java/com/badao/ai/
    │ │ ├── SpringAiHarnessDemoApplication.java # 启动类
    │ │ ├── config/
    │ │ │ ├── HarnessAgentConfig.java # 单Agent配置
    │ │ │ └── MultiAgentConfig.java # 多Agent配置
    │ │ ├── controller/
    │ │ │ └── HarnessController.java # REST API
    │ │ ├── service/
    │ │ │ ├── HarnessAgentService.java # 核心业务服务
    │ │ │ └── MultiAgentService.java # 多Agent协同服务
    │ │ ├── model/
    │ │ │ ├── ContactInfo.java # 联系人POJO
    │ │ │ └── ProductReview.java # 商品评价POJO
    │ │ ├── harness/
    │ │ │ ├── rules/
    │ │ │ │ └── ExtractionRules.md # 规则文件(软约束)
    │ │ │ ├── skills/
    │ │ │ │ └── ReviewAnalysisSkill.java # 评价分析技能
    │ │ │ ├── gates/
    │ │ │ │ └── OutputValidator.java # 输出门禁(硬约束)
    │ │ │ └── tools/
    │ │ │ ├── WeatherTool.java # 天气查询工具
    │ │ │ └── CalculatorTool.java # 计算器工具
    │ │ └── evaluation/
    │ │ ├── QualityScorer.java # 质量评分器
    │ │ └── FeedbackLogger.java # 反馈日志
    │ └── resources/
    │ ├── application.yml
    │ └── harness/rules/
    │ └── ExtractionRules.md


    四、完整代码实现

    4.1 Maven 项目配置(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
    http://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>
    </parent>

    <groupId>com.badao.ai</groupId>
    <artifactId>spring-ai-harness-demo</artifactId>
    <version>1.0.0</version>

    <properties>
    <java.version>17</java.version>
    <spring-ai-alibaba.version>1.1.2.0</spring-ai-alibaba.version>
    <jackson.version>2.17.2</jackson.version>
    </properties>

    <dependencyManagement>
    <dependencies>
    <dependency>
    <groupId>com.fasterxml.jackson</groupId>
    <artifactId>jackson-bom</artifactId>
    <version>${jackson.version}</version>
    <type>pom</type>
    <scope>import</scope>
    </dependency>
    </dependencies>
    </dependencyManagement>

    <dependencies>
    <!– Spring Boot Web –>
    <dependency>
    <groupId>org.springframework.boot</groupId>
    <artifactId>spring-boot-starter-web</artifactId>
    </dependency>

    <!– Spring AI Alibaba Agent Framework –>
    <dependency>
    <groupId>com.alibaba.cloud.ai</groupId>
    <artifactId>spring-ai-alibaba-agent-framework</artifactId>
    <version>${spring-ai-alibaba.version}</version>
    </dependency>

    <!– DashScope 模型适配器 –>
    <dependency>
    <groupId>com.alibaba.cloud.ai</groupId>
    <artifactId>spring-ai-alibaba-starter-dashscope</artifactId>
    <version>${spring-ai-alibaba.version}</version>
    </dependency>

    <!– Jackson –>
    <dependency>
    <groupId>com.fasterxml.jackson.core</groupId>
    <artifactId>jackson-databind</artifactId>
    </dependency>
    <dependency>
    <groupId>com.fasterxml.jackson.core</groupId>
    <artifactId>jackson-core</artifactId>
    </dependency>
    <dependency>
    <groupId>com.fasterxml.jackson.core</groupId>
    <artifactId>jackson-annotations</artifactId>
    </dependency>
    </dependencies>

    <repositories>
    <repository>
    <id>spring-milestones</id>
    <name>Spring Milestones</name>
    <url>https://repo.spring.io/milestone</url>
    </repository>
    </repositories>

    <build>
    <plugins>
    <plugin>
    <groupId>org.springframework.boot</groupId>
    <artifactId>spring-boot-maven-plugin</artifactId>
    </plugin>
    </plugins>
    </build>
    </project>

    注意:必须显式管理 Jackson 版本,避免版本冲突导致 NoSuchMethodError。

    4.2 配置文件(application.yml)

    server:
    port: 885

    spring:
    ai:
    dashscope:
    api-key: ${DASHSCOPE_API_KEY}
    chat:
    options:
    model: qwenmax

    logging:
    level:
    com.alibaba.cloud.ai: debug
    com.badao.ai: debug

    4.3 数据模型(POJO)

    ContactInfo.java(简单 POJO)

    package com.badao.ai.model;

    public class ContactInfo {
    private String name;
    private String email;
    private String phone;

    public ContactInfo() {}

    public ContactInfo(String name, String email, String phone) {
    this.name = name;
    this.email = email;
    this.phone = phone;
    }

    public String getName() { return name; }
    public void setName(String name) { this.name = name; }
    public String getEmail() { return email; }
    public void setEmail(String email) { this.email = email; }
    public String getPhone() { return phone; }
    public void setPhone(String phone) { this.phone = phone; }

    @Override
    public String toString() {
    return "ContactInfo{name='" + name + "', email='" + email + "', phone='" + phone + "'}";
    }
    }

    ProductReview.java(嵌套 POJO)

    package com.badao.ai.model;

    import java.util.Arrays;

    public class ProductReview {
    private int rating;
    private String sentiment; // "positive", "neutral", "negative"
    private String[] keyPoints;
    private ReviewDetails details;

    public ProductReview() {}

    public int getRating() { return rating; }
    public void setRating(int rating) { this.rating = rating; }
    public String getSentiment() { return sentiment; }
    public void setSentiment(String sentiment) { this.sentiment = sentiment; }
    public String[] getKeyPoints() { return keyPoints; }
    public void setKeyPoints(String[] keyPoints) { this.keyPoints = keyPoints; }
    public ReviewDetails getDetails() { return details; }
    public void setDetails(ReviewDetails details) { this.details = details; }

    public static class ReviewDetails {
    private String[] pros;
    private String[] cons;
    private String summary;

    public ReviewDetails() {}

    public String[] getPros() { return pros; }
    public void setPros(String[] pros) { this.pros = pros; }
    public String[] getCons() { return cons; }
    public void setCons(String[] cons) { this.cons = cons; }
    public String getSummary() { return summary; }
    public void setSummary(String summary) { this.summary = summary; }

    @Override
    public String toString() {
    return "ReviewDetails{pros=" + Arrays.toString(pros) +
    ", cons=" + Arrays.toString(cons) +
    ", summary='" + summary + "'}";
    }
    }

    @Override
    public String toString() {
    return "ProductReview{rating=" + rating +
    ", sentiment='" + sentiment + "'" +
    ", keyPoints=" + Arrays.toString(keyPoints) +
    ", details=" + details + "}";
    }
    }

    4.4 Harness 核心组件

    ① 规则层(软约束):ExtractionRules.md

    # 联系人提取规则

    ## 必填字段
    – name: 必须提取完整姓名(中文或英文)
    – email: 必须提取完整邮箱地址
    – phone: 必须提取完整电话号码(含区号)

    ## 格式要求
    – 所有字段值必须去除首尾空格
    – 电话号统一为字符串格式,保留原始格式

    ## 禁止行为
    – 不要编造任何字段值
    – 如果某个字段在原文中找不到,设置为空字符串 ""

    ② 技能层(半硬约束):ReviewAnalysisSkill.java

    package com.badao.ai.harness.skills;

    import com.badao.ai.model.ProductReview;
    import org.springframework.stereotype.Component;

    @Component
    public class ReviewAnalysisSkill {

    public String buildPrompt(String reviewText) {
    return """
    请分析以下商品评价,按标准格式输出 JSON。

    【重要约束】
    1. 如果评价内容不足以提取优缺点,请基于评价的字面含义,给出合理的推断或直接输出空数组。
    2. keyPoints 必须是从评价中提取的具体产品维度(如"音质""价格""外观"),而不是对评价本身的描述。
    3. 如果评价少于5个字,你可以推断一个合理的评分,并在 sentiment 中标注 "neutral"。

    分析步骤:
    – 评分:1-5 整数
    – 情感:positive / neutral / negative
    – 关键点:至少1个,从评价中抽取具体维度
    – 优点/缺点:尽量提取,如果没有则保留空数组
    – 总结:一句话概括整体评价

    评价文本:
    """ + reviewText;
    }

    public ProductReview postProcess(ProductReview review) {
    if (review.getRating() < 1) review.setRating(1);
    if (review.getRating() > 5) review.setRating(5);
    if (review.getKeyPoints() == null || review.getKeyPoints().length == 0) {
    review.setKeyPoints(new String[]{"无关键点"});
    }
    if (review.getDetails() == null) {
    ProductReview.ReviewDetails details = new ProductReview.ReviewDetails();
    details.setPros(new String[0]);
    details.setCons(new String[0]);
    details.setSummary("无总结");
    review.setDetails(details);
    }
    return review;
    }
    }

    ③ 门禁层(硬约束):OutputValidator.java

    package com.badao.ai.harness.gates;

    import com.badao.ai.model.ContactInfo;
    import com.badao.ai.model.ProductReview;
    import org.springframework.stereotype.Component;
    import java.util.ArrayList;
    import java.util.List;

    @Component
    public class OutputValidator {

    public List<String> validateContact(ContactInfo contact) {
    List<String> errors = new ArrayList<>();
    if (contact == null) { errors.add("联系人信息为空"); return errors; }
    if (contact.getName() == null || contact.getName().trim().isEmpty())
    errors.add("姓名不能为空");
    if (contact.getEmail() == null || !contact.getEmail().contains("@"))
    errors.add("邮箱格式无效");
    if (contact.getPhone() == null || contact.getPhone().trim().isEmpty())
    errors.add("电话不能为空");
    return errors;
    }

    public List<String> validateReview(ProductReview review) {
    List<String> errors = new ArrayList<>();
    if (review == null) { errors.add("评价信息为空"); return errors; }
    if (review.getRating() < 1 || review.getRating() > 5)
    errors.add("评分必须在 1-5 之间");
    String sentiment = review.getSentiment();
    if (sentiment == null || !(sentiment.equals("positive") ||
    sentiment.equals("neutral") || sentiment.equals("negative")))
    errors.add("情感倾向必须是 positive/neutral/negative 之一");
    if (review.getKeyPoints() == null || review.getKeyPoints().length == 0)
    errors.add("关键点不能为空");
    return errors;
    }

    public boolean isValid(List<String> errors) {
    return errors == null || errors.isEmpty();
    }
    }

    ④ 工具层:WeatherTool.java

    package com.badao.ai.harness.tools;

    import org.springframework.ai.tool.annotation.Tool;
    import org.springframework.ai.tool.annotation.ToolParam;
    import org.springframework.stereotype.Component;

    @Component
    public class WeatherTool {

    @Tool(description = "根据城市名称查询当前天气")
    public String getWeather(@ToolParam(description = "城市名称,如Beijing") String city) {
    if ("Beijing".equalsIgnoreCase(city)) {
    return "北京:晴,25°C,湿度40%";
    } else if ("Shanghai".equalsIgnoreCase(city)) {
    return "上海:多云,28°C,湿度65%";
    } else {
    return city + ":天气未知,请稍后再查";
    }
    }
    }

    ⑤ 工具层:CalculatorTool.java

    package com.badao.ai.harness.tools;

    import org.springframework.ai.tool.annotation.Tool;
    import org.springframework.ai.tool.annotation.ToolParam;
    import org.springframework.stereotype.Component;

    @Component
    public class CalculatorTool {

    @Tool(description = "执行基本的四则运算")
    public double calculate(
    @ToolParam(description = "第一个操作数") double a,
    @ToolParam(description = "运算符,支持 + – * /") String operator,
    @ToolParam(description = "第二个操作数") double b) {
    return switch (operator) {
    case "+" -> a + b;
    case "-" -> a b;
    case "*" -> a * b;
    case "/" -> a / b;
    default -> throw new IllegalArgumentException("不支持的运算符: " + operator);
    };
    }
    }

    4.5 Agent 配置类

    HarnessAgentConfig.java(单 Agent 配置)

    package com.badao.ai.config;

    import com.alibaba.cloud.ai.graph.agent.ReactAgent;
    import com.alibaba.cloud.ai.graph.checkpoint.savers.MemorySaver;
    import com.badao.ai.model.ContactInfo;
    import com.badao.ai.model.ProductReview;
    import com.badao.ai.harness.skills.ReviewAnalysisSkill;
    import org.springframework.ai.chat.model.ChatModel;
    import org.springframework.ai.converter.BeanOutputConverter;
    import org.springframework.context.annotation.Bean;
    import org.springframework.context.annotation.Configuration;
    import org.springframework.core.io.ClassPathResource;
    import java.io.IOException;
    import java.nio.charset.StandardCharsets;

    @Configuration
    public class HarnessAgentConfig {

    private final ReviewAnalysisSkill reviewAnalysisSkill;

    public HarnessAgentConfig(ReviewAnalysisSkill reviewAnalysisSkill) {
    this.reviewAnalysisSkill = reviewAnalysisSkill;
    }

    @Bean
    public ReactAgent contactAgent(ChatModel chatModel) throws IOException {
    String rules = loadRules("harness/rules/ExtractionRules.md");
    String systemPrompt = "你是一个联系人信息提取专家。请严格遵循以下规则:\\n" + rules;
    return ReactAgent.builder()
    .name("contact_extractor")
    .model(chatModel)
    .systemPrompt(systemPrompt)
    .outputType(ContactInfo.class)
    .saver(new MemorySaver())
    .build();
    }

    @Bean
    public ReactAgent reviewAgent(ChatModel chatModel) {
    BeanOutputConverter<ProductReview> converter =
    new BeanOutputConverter<>(ProductReview.class);
    String schema = converter.getFormat();
    return ReactAgent.builder()
    .name("review_analyzer")
    .model(chatModel)
    .systemPrompt("你是一个商品评价分析专家,必须按以下 JSON Schema 格式输出:\\n" + schema)
    .outputSchema(schema)
    .saver(new MemorySaver())
    .build();
    }

    private String loadRules(String path) throws IOException {
    ClassPathResource resource = new ClassPathResource(path);
    return new String(resource.getInputStream().readAllBytes(), StandardCharsets.UTF_8);
    }
    }

    MultiAgentConfig.java(多 Agent 协同配置)

    package com.badao.ai.config;

    import com.alibaba.cloud.ai.graph.agent.ReactAgent;
    import com.alibaba.cloud.ai.graph.checkpoint.savers.MemorySaver;
    import com.badao.ai.harness.tools.CalculatorTool;
    import com.badao.ai.harness.tools.WeatherTool;
    import org.springframework.ai.chat.model.ChatModel;
    import org.springframework.context.annotation.Bean;
    import org.springframework.context.annotation.Configuration;

    @Configuration
    public class MultiAgentConfig {

    @Bean
    public ReactAgent plannerAgent(ChatModel chatModel) {
    return ReactAgent.builder()
    .name("planner")
    .model(chatModel)
    .systemPrompt("""
    你是一个任务规划专家。用户会给出一个复杂需求,你需要将其拆解为1~3个明确的子任务,
    并以JSON数组格式输出,每个子任务包含:description(描述)和 assigned_to(执行角色,只能是"executor")。
    示例输出:[{"description":"查询北京天气","assigned_to":"executor"}]
    """
    )
    .saver(new MemorySaver())
    .build();
    }

    @Bean
    public ReactAgent executorAgent(ChatModel chatModel,
    WeatherTool weatherTool,
    CalculatorTool calculatorTool) {
    return ReactAgent.builder()
    .name("executor")
    .model(chatModel)
    .systemPrompt("你是一个执行专家,负责具体执行用户分配的任务。你可以使用工具完成工作。")
    .methodTools(weatherTool, calculatorTool) // 自动扫描 @Tool 注解
    .saver(new MemorySaver())
    .build();
    }
    }

    4.6 Service 层

    HarnessAgentService.java(核心业务服务)

    package com.badao.ai.service;

    import com.alibaba.cloud.ai.graph.RunnableConfig;
    import com.alibaba.cloud.ai.graph.agent.ReactAgent;
    import com.alibaba.cloud.ai.graph.exception.GraphRunnerException;
    import com.badao.ai.model.ContactInfo;
    import com.badao.ai.model.ProductReview;
    import com.badao.ai.harness.gates.OutputValidator;
    import com.badao.ai.harness.skills.ReviewAnalysisSkill;
    import com.badao.ai.evaluation.QualityScorer;
    import com.badao.ai.evaluation.FeedbackLogger;
    import com.fasterxml.jackson.databind.ObjectMapper;
    import org.springframework.ai.chat.messages.AssistantMessage;
    import org.springframework.stereotype.Service;
    import java.util.List;

    @Service
    public class HarnessAgentService {

    private final ReactAgent contactAgent;
    private final ReactAgent reviewAgent;
    private final ReviewAnalysisSkill reviewAnalysisSkill;
    private final OutputValidator outputValidator;
    private final ObjectMapper objectMapper;
    private final QualityScorer qualityScorer;
    private final FeedbackLogger feedbackLogger;

    public HarnessAgentService(ReactAgent contactAgent, ReactAgent reviewAgent,
    ReviewAnalysisSkill reviewAnalysisSkill,
    OutputValidator outputValidator,
    ObjectMapper objectMapper,
    QualityScorer qualityScorer,
    FeedbackLogger feedbackLogger) {
    this.contactAgent = contactAgent;
    this.reviewAgent = reviewAgent;
    this.reviewAnalysisSkill = reviewAnalysisSkill;
    this.outputValidator = outputValidator;
    this.objectMapper = objectMapper;
    this.qualityScorer = qualityScorer;
    this.feedbackLogger = feedbackLogger;
    }

    // ========== 联系人提取 ==========
    public ContactInfo extractContact(String text, String sessionId) {
    RunnableConfig config = RunnableConfig.builder().threadId(sessionId).build();
    AssistantMessage response;
    try {
    response = contactAgent.call(text, config);
    } catch (GraphRunnerException e) {
    throw new RuntimeException("Agent 执行失败: " + e.getMessage(), e);
    }
    String json = response.getText();
    ContactInfo contact;
    try {
    contact = objectMapper.readValue(json, ContactInfo.class);
    } catch (Exception e) {
    throw new RuntimeException("解析结构化输出失败,原始 JSON: " + json, e);
    }
    // Gate 门禁校验
    List<String> errors = outputValidator.validateContact(contact);
    if (!outputValidator.isValid(errors)) {
    throw new RuntimeException("输出校验失败: " + String.join("; ", errors));
    }
    return contact;
    }

    public String extractContactRaw(String text, String sessionId) {
    RunnableConfig config = RunnableConfig.builder().threadId(sessionId).build();
    try {
    AssistantMessage response = contactAgent.call(text, config);
    return response.getText();
    } catch (GraphRunnerException e) {
    throw new RuntimeException("Agent 执行失败: " + e.getMessage(), e);
    }
    }

    // ========== 商品评价分析 ==========
    public ProductReview analyzeReview(String reviewText, String sessionId) {
    RunnableConfig config = RunnableConfig.builder().threadId(sessionId).build();
    String prompt = reviewAnalysisSkill.buildPrompt(reviewText);
    AssistantMessage response;
    try {
    response = reviewAgent.call(prompt, config);
    } catch (GraphRunnerException e) {
    throw new RuntimeException("Agent 执行失败: " + e.getMessage(), e);
    }
    String json = response.getText();
    ProductReview review;
    try {
    review = objectMapper.readValue(json, ProductReview.class);
    } catch (Exception e) {
    throw new RuntimeException("解析结构化输出失败,原始 JSON: " + json, e);
    }
    // Skill 后处理
    review = reviewAnalysisSkill.postProcess(review);
    // Gate 门禁校验
    List<String> errors = outputValidator.validateReview(review);
    if (!outputValidator.isValid(errors)) {
    throw new RuntimeException("输出校验失败: " + String.join("; ", errors));
    }
    return review;
    }

    // ========== 带评估的评价分析 ==========
    public ProductReview analyzeReviewWithEval(String reviewText, String sessionId) {
    ProductReview review = analyzeReview(reviewText, sessionId);
    // 质量评估
    int score = qualityScorer.scoreReview(review);
    if (score < 3) {
    feedbackLogger.logFailedCase(reviewText, review, "质量评分过低: " + score);
    feedbackLogger.triggerHumanReview(reviewText, review);
    throw new RuntimeException("质量门禁未通过(得分" + score + "/4),已转人工复核");
    }
    return review;
    }
    }

    MultiAgentService.java(多 Agent 协同服务)

    package com.badao.ai.service;

    import com.alibaba.cloud.ai.graph.RunnableConfig;
    import com.alibaba.cloud.ai.graph.agent.ReactAgent;
    import com.alibaba.cloud.ai.graph.exception.GraphRunnerException;
    import com.fasterxml.jackson.databind.JsonNode;
    import com.fasterxml.jackson.databind.ObjectMapper;
    import org.springframework.ai.chat.messages.AssistantMessage;
    import org.springframework.stereotype.Service;
    import java.util.ArrayList;
    import java.util.List;

    @Service
    public class MultiAgentService {

    private final ReactAgent plannerAgent;
    private final ReactAgent executorAgent;
    private final ObjectMapper objectMapper;

    public MultiAgentService(ReactAgent plannerAgent, ReactAgent executorAgent,
    ObjectMapper objectMapper) {
    this.plannerAgent = plannerAgent;
    this.executorAgent = executorAgent;
    this.objectMapper = objectMapper;
    }

    public String executeComplexTask(String userRequest, String sessionId) {
    RunnableConfig config = RunnableConfig.builder().threadId(sessionId).build();

    // 1. 规划阶段
    AssistantMessage planResponse;
    try {
    planResponse = plannerAgent.call("请拆解以下任务:" + userRequest, config);
    } catch (GraphRunnerException e) {
    throw new RuntimeException("规划失败: " + e.getMessage(), e);
    }
    String planJson = planResponse.getText();

    // 解析规划结果
    List<JsonNode> tasks;
    try {
    tasks = objectMapper.readValue(planJson, objectMapper.getTypeFactory()
    .constructCollectionType(List.class, JsonNode.class));
    } catch (Exception e) {
    throw new RuntimeException("规划结果解析失败: " + planJson, e);
    }

    // 2. 执行阶段
    List<String> results = new ArrayList<>();
    for (JsonNode task : tasks) {
    String description = task.get("description").asText();
    String assignedTo = task.get("assigned_to").asText();
    if (!"executor".equals(assignedTo)) {
    results.add("跳过不支持的角色: " + assignedTo);
    continue;
    }
    AssistantMessage execResponse;
    try {
    execResponse = executorAgent.call(description, config);
    } catch (GraphRunnerException e) {
    results.add("执行子任务失败: " + description + ",错误: " + e.getMessage());
    continue;
    }
    results.add(execResponse.getText());
    }

    return String.join("\\n", results);
    }
    }

    4.7 评估与反馈组件

    QualityScorer.java

    package com.badao.ai.evaluation;

    import com.badao.ai.model.ProductReview;
    import org.springframework.stereotype.Component;

    @Component
    public class QualityScorer {

    public int scoreReview(ProductReview review) {
    int score = 0;
    if (review.getRating() >= 1 && review.getRating() <= 5) score += 1;
    if (review.getSentiment() != null) score += 1;
    if (review.getKeyPoints() != null && review.getKeyPoints().length >= 2) score += 1;
    if (review.getDetails() != null && review.getDetails().getSummary() != null) score += 1;
    return score; // 满分4分
    }
    }

    FeedbackLogger.java

    package com.badao.ai.evaluation;

    import org.slf4j.Logger;
    import org.slf4j.LoggerFactory;
    import org.springframework.stereotype.Component;

    @Component
    public class FeedbackLogger {

    private static final Logger logger = LoggerFactory.getLogger(FeedbackLogger.class);

    public void logFailedCase(Object input, Object output, String reason) {
    logger.warn("失败案例 – 输入: {}, 输出: {}, 原因: {}", input, output, reason);
    }

    public void triggerHumanReview(Object input, Object output) {
    logger.info("人工复核触发 – 输入: {}, 输出: {}", input, output);
    }
    }

    4.8 Controller 层

    package com.badao.ai.controller;

    import com.badao.ai.model.ContactInfo;
    import com.badao.ai.model.ProductReview;
    import com.badao.ai.service.HarnessAgentService;
    import com.badao.ai.service.MultiAgentService;
    import org.springframework.web.bind.annotation.*;
    import java.util.Map;

    @RestController
    @RequestMapping("/api/harness")
    public class HarnessController {

    private final HarnessAgentService harnessAgentService;
    private final MultiAgentService multiAgentService;

    public HarnessController(HarnessAgentService harnessAgentService,
    MultiAgentService multiAgentService) {
    this.harnessAgentService = harnessAgentService;
    this.multiAgentService = multiAgentService;
    }

    // ===== 联系人提取 =====
    @PostMapping("/contact")
    public Map<String, Object> extractContact(
    @RequestParam String text,
    @RequestParam(required = false, defaultValue = "default") String sessionId) {
    ContactInfo contact = harnessAgentService.extractContact(text, sessionId);
    return Map.of("success", true, "data", contact, "sessionId", sessionId);
    }

    @PostMapping("/contact/raw")
    public Map<String, Object> extractContactRaw(
    @RequestParam String text,
    @RequestParam(required = false, defaultValue = "default") String sessionId) {
    String json = harnessAgentService.extractContactRaw(text, sessionId);
    return Map.of("success", true, "json", json, "sessionId", sessionId);
    }

    // ===== 商品评价分析 =====
    @PostMapping("/review")
    public Map<String, Object> analyzeReview(
    @RequestParam String reviewText,
    @RequestParam(required = false, defaultValue = "default") String sessionId) {
    ProductReview review = harnessAgentService.analyzeReview(reviewText, sessionId);
    return Map.of("success", true, "data", review, "sessionId", sessionId);
    }

    @PostMapping("/review/eval")
    public Map<String, Object> analyzeReviewWithEval(
    @RequestParam String reviewText,
    @RequestParam(required = false, defaultValue = "default") String sessionId) {
    ProductReview review = harnessAgentService.analyzeReviewWithEval(reviewText, sessionId);
    return Map.of("success", true, "data", review, "sessionId", sessionId);
    }

    // ===== 多 Agent 协同 =====
    @PostMapping("/complex")
    public Map<String, Object> complexTask(
    @RequestParam String request,
    @RequestParam(required = false, defaultValue = "default") String sessionId) {
    String result = multiAgentService.executeComplexTask(request, sessionId);
    return Map.of("success", true, "result", result, "sessionId", sessionId);
    }
    }

    4.9 启动类

    package com.badao.ai;

    import org.springframework.boot.SpringApplication;
    import org.springframework.boot.autoconfigure.SpringBootApplication;

    @SpringBootApplication
    public class SpringAiHarnessDemoApplication {
    public static void main(String[] args) {
    SpringApplication.run(SpringAiHarnessDemoApplication.class, args);
    }
    }


    五、七层架构深度解析

    5.1 执行主循环层

    知识点:这是 Agent 运行的"大脑",负责"规划-执行-反思"的 ReAct 循环。

    代码体现:

    // ReactAgent 内置了 ReAct 循环
    ReactAgent.builder()
    .model(chatModel) // 绑定模型
    .saver(new MemorySaver()) // 状态管理,支持错误恢复
    .build();

    最佳实践:

    • 避免无规划的扁平串行循环,必须加入前置规划和后置反思
    • 所有 LLM 调用应内置分级重试与降级策略
    • 使用 MemorySaver 实现状态持久化

    5.2 工具编排层

    知识点:Agent 与外部世界交互的"手脚",需要建立边界管控、权限校验和结果处理体系。

    代码体现:

    // 使用 @Tool 注解定义工具
    @Tool(description = "根据城市名称查询当前天气")
    public String getWeather(@ToolParam(description = "城市名称") String city) { ... }

    // 在 Agent 中注册工具
    .methodTools(weatherTool, calculatorTool)

    最佳实践:

    • 为每个工具定义清晰的 description,帮助模型决策
    • 工具结果应进行结构化处理和异常捕获
    • 建立工具黑白名单机制,防止危险操作

    5.3 项目搭建层

    知识点:标准化项目的"地基",统一项目结构、依赖和规范。

    代码体现:

    <!– pom.xml 统一版本管理 –>
    <parent>
    <groupId>org.springframework.boot</groupId>
    <artifactId>spring-boot-starter-parent</artifactId>
    <version>3.2.5</version>
    </parent>

    <properties>
    <java.version>17</java.version>
    <spring-ai-alibaba.version>1.1.2.0</spring-ai-alibaba.version>
    </properties>

    # application.yml 环境配置
    spring:
    ai:
    dashscope:
    api-key: ${DASHSCOPE_API_KEY} # 环境变量注入

    最佳实践:

    • 使用父 POM 统一管理依赖版本
    • 敏感信息通过环境变量注入,实现环境隔离
    • 建立模板仓库,包含基础依赖、安全策略和监控配置

    5.4 上下文工程层

    知识点:Agent 的"记忆"与"知识库",管理、优化注入给模型的信息。

    代码体现:

    // 1. 加载规则文件作为静态上下文
    String rules = loadRules("harness/rules/ExtractionRules.md");
    String systemPrompt = "你是一个联系人信息提取专家。\\n" + rules;

    // 2. Skill 构建结构化提示词
    public String buildPrompt(String reviewText) {
    return """
    请分析以下商品评价,按标准格式输出 JSON。
    分析步骤:
    1. 提取评分(1-5星整数)
    2. 判断情感倾向:positive / neutral / negative

    """
    + reviewText;
    }

    最佳实践:

    • 使用分层缓存:基础层(项目配置)+ 会话层(当前任务)+ 临时层(用户输入)
    • 构建知识图谱,让 Agent 理解项目全貌
    • 对上下文进行版本控制,便于追溯

    5.5 约束与防护层

    知识点:Agent 的"安全护栏",硬性校验输入输出,防止失控。

    代码体现:

    // Gate 门禁校验
    public List<String> validateContact(ContactInfo contact) {
    List<String> errors = new ArrayList<>();
    if (contact.getName() == null || contact.getName().trim().isEmpty())
    errors.add("姓名不能为空");
    if (contact.getEmail() == null || !contact.getEmail().contains("@"))
    errors.add("邮箱格式无效");
    return errors;
    }

    // 在 Service 中执行熔断
    List<String> errors = outputValidator.validateContact(contact);
    if (!outputValidator.isValid(errors)) {
    throw new RuntimeException("输出校验失败: " + String.join("; ", errors));
    }

    最佳实践:

    • 输入验证:正则表达式过滤非法请求
    • 输出过滤:检查敏感信息泄露或危险操作
    • 熔断机制:异常行为自动触发回滚或降级

    5.6 多 Agent 架构层

    知识点:团队协作的"管理模式",让多个专业 Agent 协同工作。

    代码体现:

    // 主从模式
    @Bean
    public ReactAgent plannerAgent(ChatModel chatModel) {
    return ReactAgent.builder()
    .name("planner")
    .systemPrompt("你是一个任务规划专家…")
    .build();
    }

    @Bean
    public ReactAgent executorAgent(ChatModel chatModel) {
    return ReactAgent.builder()
    .name("executor")
    .systemPrompt("你是一个执行专家…")
    .methodTools(weatherTool, calculatorTool)
    .build();
    }

    最佳实践:

    • 主从模式:主 Agent 分配任务,子 Agent 执行
    • 标准化通信协议,定义清晰的任务描述格式
    • 建立冲突解决机制,处理子 Agent 结果不一致的情况

    5.7 评估与反馈层

    知识点:持续优化的"学习闭环",通过测试和反馈不断改进系统。

    代码体现:

    // 质量评分
    public int scoreReview(ProductReview review) {
    int score = 0;
    if (review.getRating() >= 1 && review.getRating() <= 5) score += 1;
    if (review.getSentiment() != null) score += 1;
    if (review.getKeyPoints() != null && review.getKeyPoints().length >= 2) score += 1;
    if (review.getDetails() != null && review.getDetails().getSummary() != null) score += 1;
    return score;
    }

    // 质量门禁
    if (score < 3) {
    feedbackLogger.logFailedCase(input, output, "质量评分过低");
    feedbackLogger.triggerHumanReview(input, output);
    throw new RuntimeException("质量门禁未通过");
    }

    最佳实践:

    • 建立多维度的评估指标(准确性、完整性、安全性)
    • 设置质量门禁阈值,未达标触发人工复核
    • 收集用户修改行为,作为后续微调的数据

    六、运行与测试

    6.1 启动应用

    在项目根目录执行:

    mvn clean package
    java -jar target/spring-ai-harness-demo-1.0.0.jar

    或在 IDEA 中直接运行 SpringAiHarnessDemoApplication。

    6.2 测试接口

    ① 联系人提取

    curl -X POST "http://localhost:885/api/harness/contact?text=从以下信息提取联系方式:王五,wangwu@outlook.com,+86 139-9999-8888&sessionId=test01"

    预期返回:

    {
    "success": true,
    "data": {
    "name": "王五",
    "email": "wangwu@outlook.com",
    "phone": "+86 139-9999-8888"
    },
    "sessionId": "test01"
    }

    ② 商品评价分析

    curl -X POST "http://localhost:885/api/harness/review?reviewText=这款耳机音质不错,降噪效果好,但佩戴舒适度一般,价格略高。&sessionId=test02"

    预期返回:

    {
    "success": true,
    "data": {
    "rating": 4,
    "sentiment": "positive",
    "keyPoints": ["音质不错", "降噪效果好", "佩戴舒适度一般", "价格略高"],
    "details": {
    "pros": ["音质不错", "降噪效果好"],
    "cons": ["佩戴舒适度一般", "价格略高"],
    "summary": "整体满意,舒适度和价格有改进空间"
    }
    },
    "sessionId": "test02"
    }

    ③ 带质量评估的评价分析

    curl -X POST "http://localhost:885/api/harness/review/eval?reviewText=还可以吧&sessionId=eval01"

    预期返回(通过质量门禁时):

    {
    "success": true,
    "data": {
    "rating": 3,
    "sentiment": "neutral",
    "keyPoints": ["态度中立", "缺乏细节"],
    "details": {
    "pros": [],
    "cons": [],
    "summary": "评价者仅表示"还可以吧",态度中立,没有明确指出优劣。"
    }
    },
    "sessionId": "eval01"
    }

    ④ 多 Agent 协同

    curl -X POST "http://localhost:885/api/harness/complex?request=查询北京天气并计算25+17&sessionId=multi01"

    预期返回:

    {
    "success": true,
    "result": "北京:晴,25°C,湿度40%\\n42.0",
    "sessionId": "multi01"
    }


    七、常见问题与解决方案

    问题 1:Cannot resolve method 'tools(List<E>)'

    原因:ReactAgent.Builder 的 tools 方法不直接接受 Spring Bean 实例。

    解决方案:使用 methodTools 方法,它会自动扫描 @Tool 注解。

    // ❌ 错误
    .tools(weatherToolCallback, calculatorToolCallback)

    // ✅ 正确
    .methodTools(weatherTool, calculatorTool)

    问题 2:Cannot resolve method 'builder(String, WeatherTool)'

    原因:FunctionToolCallback.builder("getWeather", weatherTool) 的 API 在当前版本不支持。

    解决方案:改用 methodTools,或使用 lambda 方式构建。

    问题 3:解析结构化输出失败

    原因:模型返回的 JSON 格式与 POJO 不完全匹配。

    解决方案:

  • 在 systemPrompt 中明确要求 JSON 格式
  • 使用 outputType 或 outputSchema 强制格式
  • 在 postProcess 中增加兜底处理
  • 问题 4:质量门禁总是失败

    原因:评估阈值设置过高或评分逻辑过严。

    解决方案:

  • 根据实际测试结果调整阈值
  • 增加评分维度,使评分更精确
  • 区分"必须通过"和"建议通过"的指标

  • 八、总结与进阶方向

    8.1 总结

    通过本示例,我们完成了:

    目标完成情况
    理解 Harness Engineering 概念 ✅ 核心组件 + 七层架构
    Windows 本地环境搭建 ✅ JDK 17 + Maven + DashScope
    完整代码实现 ✅ 6 个模块,20+ 个类
    七层架构落地 ✅ 每层都有对应的代码实现
    运行与测试 ✅ 4 个接口均可正常调用

    8.2 Harness Engineering 在本示例中的映射

    Harness 组件实现位置作用
    Rules(规则) ExtractionRules.md 软约束,告诉 Agent 必须提取哪些字段
    Skills(技能) ReviewAnalysisSkill.java 半硬约束,封装标准化提示词和后处理
    Gate(门禁) OutputValidator.java 硬约束,强制校验输出结构
    State(状态) MemorySaver + threadId 管理会话上下文
    Instructions(指令) systemPrompt + outputType 明确输出格式和行为边界
    Verification(验证) 异常处理 + Gate 校验 只有通过才算任务完成
    Tools(工具) WeatherTool + CalculatorTool 扩展 Agent 能力边界
    Evaluation(评估) QualityScorer 量化输出质量
    Feedback(反馈) FeedbackLogger 收集失败案例,触发人工复核

    8.3 进阶方向

  • 动态规则加载:将 Rules 从本地文件迁移到配置中心,实现热更新
  • 更复杂的工具集成:接入数据库查询、Web 搜索、文件读写等
  • 多 Agent 深度协作:实现流水线模式或对等协商模式
  • 模型微调:将失败案例加入训练集,提升模型表现
  • 可观测性:集成 Prometheus + Grafana 监控 Agent 运行状态
  • 大规模部署:使用 Kubernetes 部署多个 Harness 实例

  • 参考资源

    • Spring AI Alibaba 官方文档
    • Learn Harness Engineering 开源课程
    • DashScope 百炼平台

    鸿蒙端侧 NLP 实战:分词器 + TF-IDF + 朴素贝叶斯分类 + 文本相似度

    master阅读(41)

    在这里插入图片描述

    技术栈:HarmonyOS NEXT(API 12+)|ArkTS 原生开发
    核心亮点:纯算法实现,不依赖任何第三方 NLP 库;覆盖分词 → 关键词提取 → 文本分类 → 相似度计算全链路
    难度定位:中高级 · 完整可运行代码


    一、为什么要在端侧做 NLP?

    随着设备算力提升和隐私保护需求增长,将 NLP 能力下沉到端侧已成为趋势:

    • 隐私优先:文本数据不出设备,符合数据最小化原则
    • 低延迟:无需网络往返,毫秒级响应
    • 离线可用:弱网或无网环境依然工作
    • 成本节省:减少云端算力开销

    本文选取 方向 A:端侧自然语言处理,从零实现一套完整的文本处理流水线,代码可直接嵌入 ArkTS 项目使用。


    二、整体架构

    用户文本输入


    ┌─────────────┐
    │ 文本预处理 │ → 全角转半角、大小写统一、去除特殊字符
    └──────┬──────┘

    ┌─────────────────────┐
    │ 正向最大匹配 (MMM) │ ← 基于词典的贪心分词
    │ 逆向最大匹配 (RMM) │
    │ 双向最大匹配 (BiMMM)│ ← 融合策略
    └──────┬──────────────┘

    ┌─────────────────────┐
    │ 停用词过滤 │
    └──────┬──────────────┘

    ┌─────────────┐
    │ TF-IDF │ → 关键词提取 + 向量化
    │ 关键词提取 │
    └──────┬──────┘

    ┌─────────────────────┐
    │ 朴素贝叶斯分类器 │ → 文本分类
    └──────┬──────────────┘

    ┌─────────────────────────────┐
    │ 余弦相似度 / Jaccard / │
    │ Jaro-Winkler 距离 │ → 文本相似度计算
    └─────────────────────────────┘


    三、完整 ArkTS 代码实现

    3.1 项目结构

    entry/src/main/ets/
    ├── utils/
    │ ├── TextPreprocessor.ets // 文本预处理工具
    │ ├── Dictionary.ets // 本地词典(内置 + 可扩展)
    │ ├── Segmenter.ets // 分词器(MMM / RMM / BiMMM)
    │ ├── TFIDFExtractor.ets // TF-IDF 关键词提取
    │ ├── NaiveBayesClassifier.ets // 朴素贝叶斯分类器
    │ └── TextSimilarity.ets // 文本相似度计算
    ├── model/
    │ └── NLPipeline.ets // NLP 全流程封装
    └── pages/
    └── NLPDemo.ets // 演示页面

    提示:以下代码均为单文件可直接运行,只需按结构放入项目对应目录即可。


    3.2 文本预处理工具 TextPreprocessor.ets

    // TextPreprocessor.ets
    // 文本预处理:标准化输入,为后续分词做准备

    /**
    * 字符类型枚举
    */

    enum CharType {
    CHINESE = 'CHINESE', // 中文
    ENGLISH = 'ENGLISH', // 英文字母
    DIGIT = 'DIGIT', // 数字
    PUNCTUATION = 'PUNCTUATION', // 标点符号
    SPACE = 'SPACE', // 空格/空白
    OTHER = 'OTHER' // 其他字符
    }

    /**
    * 中文停用词表(常用)
    */

    const STOP_WORDS: Set<string> = new Set([
    '的', '了', '在', '是', '我', '有', '和', '就', '不', '人',
    '都', '一', '一个', '上', '也', '很', '到', '说', '要', '去',
    '你', '会', '着', '没有', '看', '好', '自己', '这', '那', '个',
    '他', '她', '它', '们', '地', '得', '着', '而', '与', '及',
    '等', '或', '但', '如果', '因为', '所以', '虽然', '然后',
    '可以', '这个', '那个', '什么', '怎么', '为什么', '还', '又',
    '把', '被', '让', '对', '以', '于', '中', '从', '由', '已'
    ]);

    /**
    * 文本预处理器
    */

    export class TextPreprocessor {
    /**
    * 全角转半角
    */

    static fullWidthToHalfWidth(text: string): string {
    let result = '';
    for (let i = 0; i < text.length; i++) {
    const charCode = text.charCodeAt(i);
    // 全角字符范围:0xFF01 – 0xFF5E,对应半角 0x21 – 0x7E
    if (charCode >= 0xFF01 && charCode <= 0xFF5E) {
    result += String.fromCharCode(charCode 0xFEE0);
    } else if (charCode === 0x3000) { // 全角空格
    result += ' ';
    } else {
    result += text[i];
    }
    }
    return result;
    }

    /**
    * 判断字符类型
    */

    static getCharType(char: string): CharType {
    const code = char.charCodeAt(0);
    if (this.isChineseChar(code)) {
    return CharType.CHINESE;
    }
    if ((code >= 0x41 && code <= 0x5A) || (code >= 0x61 && code <= 0x7A)) {
    return CharType.ENGLISH;
    }
    if (code >= 0x30 && code <= 0x39) {
    return CharType.DIGIT;
    }
    if (this.isPunctuation(char)) {
    return CharType.PUNCTUATION;
    }
    if (/\\s/.test(char)) {
    return CharType.SPACE;
    }
    return CharType.OTHER;
    }

    /**
    * 判断是否为中文 Unicode 范围
    */

    private static isChineseChar(code: number): boolean {
    // CJK 统一汉字扩展A-F + 基本汉字 + 扩展B
    return (code >= 0x4E00 && code <= 0x9FFF) ||
    (code >= 0x3400 && code <= 0x4DBF) ||
    (code >= 0x20000 && code <= 0x2A6DF);
    }

    /**
    * 判断是否为标点符号
    */

    private static isPunctuation(char: string): boolean {
    const code = char.charCodeAt(0);
    // ASCII 标点
    if (code >= 0x21 && code <= 0x2F) return true;
    if (code >= 0x3A && code <= 0x40) return true;
    if (code >= 0x5B && code <= 0x60) return true;
    if (code >= 0x7B && code <= 0x7E) return true;
    // 中文标点(常见)
    const chinesePunct = ',。!?;:""''《》【】()—…·~「」『』';
    return chinesePunct.includes(char);
    }

    /**
    * 清理特殊字符:保留中文、英文、数字,替换标点为空格
    */

    static cleanText(text: string): string {
    let result = '';
    for (let i = 0; i < text.length; i++) {
    const char = text[i];
    const type = this.getCharType(char);
    if (type === CharType.CHINESE || type === CharType.ENGLISH) {
    result += char;
    } else if (type === CharType.DIGIT) {
    result += char;
    } else if (type === CharType.SPACE) {
    result += ' ';
    } else {
    result += ' '; // 标点和特殊字符替换为空格
    }
    }
    return result;
    }

    /**
    * 合并连续空格
    */

    static normalizeSpaces(text: string): string {
    return text.replace(/\\s+/g, ' ').trim();
    }

    /**
    * 过滤停用词
    */

    static removeStopWords(tokens: string[]): string[] {
    return tokens.filter(token =>
    token.length > 1 && !STOP_WORDS.has(token)
    );
    }

    /**
    * 一站式预处理
    */

    static preprocess(text: string, removeStop: boolean = true): string[] {
    let result = text;
    // 1. 全角转半角
    result = this.fullWidthToHalfWidth(result);
    // 2. 清理特殊字符
    result = this.cleanText(result);
    // 3. 合并空格
    result = this.normalizeSpaces(result);
    // 4. 转为小写(英文部分)
    result = result.toLowerCase();
    // 5. 分词(简单按空格分,英文/数字保留为独立词)
    let tokens = result.split(' ').filter(t => t.length > 0);
    // 6. 停用词过滤
    if (removeStop) {
    tokens = this.removeStopWords(tokens);
    }
    return tokens;
    }
    }


    3.3 本地词典 Dictionary.ets

    // Dictionary.ets
    // 内置词典 + 可运行时添加词语(适应不同垂直领域)

    /**
    * 词典条目
    */

    interface DictEntry {
    word: string;
    frequency: number; // 词频(用于 MM 分词时的优先排序)
    }

    /**
    * 本地词典管理器
    * 词典按词长索引,加速最大匹配查找
    */

    export class Dictionary {
    // 按词长分组的词典(key: 词长, value: 词长为该值的词集合)
    private dictByLength: Map<number, Map<string, number>> = new Map();
    // 单词集合(用于快速存在性判断)
    private wordSet: Set<string> = new Set();
    // 最大词长(超过此长度的词直接跳过)
    private maxWordLength: number = 6;

    constructor() {
    this.loadBuiltinDictionary();
    }

    /**
    * 加载内置词典(通用中文词典片段,可按需扩展)
    */

    private loadBuiltinDictionary(): void {
    const builtinWords: Array<[string, number]> = [
    // 常用词(词频仅作参考,实际使用中可忽略)
    ['手机', 100], ['电脑', 95], ['应用', 90], ['程序', 88],
    ['数据', 95], ['用户', 98], ['系统', 99], ['网络', 90],
    ['信息', 92], ['服务', 88], ['功能', 85], ['设置', 80],
    ['文件', 88], ['图片', 85], ['音乐', 75], ['视频', 80],
    ['下载', 82], ['安装', 85], ['卸载', 70], ['打开', 90],
    ['关闭', 88], ['搜索', 92], ['连接', 80], ['密码', 85],
    ['账号', 88], ['登录', 90], ['注册', 85], ['退出', 80],
    ['消息', 85], ['通知', 82], ['提醒', 75], ['日历', 70],
    ['相机', 80], ['相册', 78], ['地图', 75], ['导航', 72],
    ['天气', 78], ['新闻', 80], ['购物', 75], ['支付', 85],
    ['银行', 70], ['游戏', 78], ['语音', 82], ['文字', 75],
    ['输入', 85], ['编辑', 80], ['删除', 82], ['保存', 88],
    ['发送', 85], ['接收', 80], ['阅读', 75], ['分享', 78],
    ['打印', 65], ['扫描', 68], ['复制', 80], ['粘贴', 82],
    ['无线', 72], ['蓝牙', 70], ['屏幕', 85], ['电池', 75],
    ['充电', 78], ['内存', 82], ['存储', 80], ['处理器', 75],
    ['人工智能', 85], ['机器学习', 80], ['深度学习', 82],
    ['自然语言', 88], ['语音识别', 85], ['图像识别', 82],
    ['智慧助手', 78], ['智能家居', 75], ['健康管理', 72],
    ['传感器', 68], ['加速度', 65], ['陀螺仪', 62],
    // 常用双字词
    ['鸿蒙', 95], ['Harmony', 90], ['华为', 92], ['应用市场', 85],
    ['操作系统', 88], ['开发者', 90], ['开发板', 82],
    ['智能', 92], ['端侧', 85], ['设备', 90], ['本地', 88],
    ['隐私', 85], ['加密', 80], ['安全', 88], ['认证', 78],
    ['推荐', 82], ['搜索', 90], ['分类', 85], ['标签', 80],
    // 高频短语(三字及以上)
    ['文本处理', 75], ['关键词', 80], ['文本分类', 78],
    ['分词器', 72], ['相似度', 75], ['向量空间', 68],
    // 更多通用词(随机补充以增加覆盖率)
    ['今天', 88], ['明天', 85], ['现在', 90], ['时间', 92],
    ['工作', 85], ['生活', 82], ['学习', 80], ['问题', 88],
    ['方法', 82], ['结果', 80], ['原因', 78], ['情况', 75],
    ['公司', 85], ['项目', 88], ['产品', 86], ['技术', 90],
    ['版本', 82], ['更新', 80], ['修复', 75], ['测试', 78],
    ];

    for (const [word, freq] of builtinWords) {
    this.addWord(word, freq);
    }
    }

    /**
    * 添加词语到词典
    * @param word 词语
    * @param frequency 词频(数值越大越优先被分出)
    */

    addWord(word: string, frequency: number = 10): void {
    const len = word.length;
    if (len === 0 || len > this.maxWordLength) {
    return;
    }
    this.wordSet.add(word);

    if (!this.dictByLength.has(len)) {
    this.dictByLength.set(len, new Map());
    }
    this.dictByLength.get(len)!.set(word, frequency);
    }

    /**
    * 批量添加词语
    */

    addWords(words: Array<[string, number]>): void {
    for (const [word, freq] of words) {
    this.addWord(word, freq);
    }
    }

    /**
    * 检查词语是否在词典中
    */

    contains(word: string): boolean {
    return this.wordSet.has(word);
    }

    /**
    * 获取词语词频
    */

    getFrequency(word: string): number {
    const len = word.length;
    if (!this.dictByLength.has(len)) return 0;
    return this.dictByLength.get(len)!.get(word) ?? 0;
    }

    /**
    * 获取最大词长
    */

    getMaxWordLength(): number {
    return this.maxWordLength;
    }

    /**
    * 设置最大词长
    */

    setMaxWordLength(len: number): void {
    if (len >= 2 && len <= 20) {
    this.maxWordLength = len;
    }
    }

    /**
    * 获取指定长度的所有词(用于分词匹配)
    */

    getWordsOfLength(len: number): Array<[string, number]> {
    if (!this.dictByLength.has(len)) return [];
    const map = this.dictByLength.get(len)!;
    const result: Array<[string, number]> = [];
    map.forEach((freq, word) => {
    result.push([word, freq]);
    });
    return result;
    }

    /**
    * 获取当前词典词数
    */

    getWordCount(): number {
    return this.wordSet.size;
    }
    }


    3.4 分词器 Segmenter.ets(核心算法)

    // Segmenter.ets
    // 基于最大匹配(Maximum Matching)的中文分词器
    // 支持:正向最大匹配 (MMM)、逆向最大匹配 (RMM)、双向最大匹配 (BiMMM)

    import { Dictionary } from './Dictionary';

    /**
    * 分词模式枚举
    */

    export enum SegmentMode {
    /** 正向最大匹配 (Forward Maximum Matching) */
    FORWARD = 'FORWARD',
    /** 逆向最大匹配 (Reverse Maximum Matching) */
    BACKWARD = 'BACKWARD',
    /** 双向最大匹配 (Bidirectional Maximum Matching) */
    BIDIRECTIONAL = 'BIDIRECTIONAL'
    }

    /**
    * 分词结果
    */

    export interface SegmentResult {
    words: string[]; // 分词结果
    mode: SegmentMode; // 采用的分词模式
    matchCount: number; // 成功匹配词数
    oovCount: number; // 未登录词(OOV)数量
    }

    /**
    * 最大匹配分词器
    */

    export class MaxMatchSegmenter {
    private dictionary: Dictionary;

    constructor(dictionary?: Dictionary) {
    this.dictionary = dictionary ?? new Dictionary();
    }

    /**
    * ========== 正向最大匹配 (Forward Maximum Matching, FMM) ==========
    * 算法:从句子开头起,每次尽量匹配最长词
    * 步骤:
    * 1. 从未处理文本的开头取 N 个字(N = maxWordLength)
    * 2. 在词典中查找:
    * – 找到 → 输出该词,移动窗口
    * – 未找到 → 去除最后一个字,缩短 N-1
    * – N = 1 仍未找到 → 单字成词,输出,移动一格
    * 3. 重复直到文本结束
    */

    forwardMatch(text: string): SegmentResult {
    const words: string[] = [];
    let i = 0;
    const len = text.length;

    while (i < len) {
    let matched = false;
    // 从最大词长开始尝试,逐步缩短
    for (let n = this.dictionary.getMaxWordLength(); n >= 1; n) {
    if (i + n > len) continue;
    const substr = text.substring(i, i + n);
    if (this.dictionary.contains(substr)) {
    words.push(substr);
    i += n;
    matched = true;
    break;
    }
    }
    // 未找到任何匹配,单字成词
    if (!matched) {
    words.push(text[i]);
    i++;
    }
    }

    return {
    words,
    mode: SegmentMode.FORWARD,
    matchCount: words.filter(w => this.dictionary.contains(w)).length,
    oovCount: words.filter(w => !this.dictionary.contains(w)).length
    };
    }

    /**
    * ========== 逆向最大匹配 (Reverse Maximum Matching, RMM) ==========
    * 算法:从句子结尾起,每次尽量匹配最长词,然后逆序输出
    * 优点:对汉语"长词优先"有更好的处理,尤其适合动词短语
    */

    backwardMatch(text: string): SegmentResult {
    const words: string[] = [];
    let i = text.length;

    while (i > 0) {
    let matched = false;
    // 从最大词长开始尝试
    for (let n = this.dictionary.getMaxWordLength(); n >= 1; n) {
    if (i n < 0) continue;
    const substr = text.substring(i n, i);
    if (this.dictionary.contains(substr)) {
    words.unshift(substr); // 头部插入(保证顺序)
    i -= n;
    matched = true;
    break;
    }
    }
    // 未找到,单字
    if (!matched) {
    words.unshift(text[i 1]);
    i;
    }
    }

    return {
    words,
    mode: SegmentMode.BACKWARD,
    matchCount: words.filter(w => this.dictionary.contains(w)).length,
    oovCount: words.filter(w => !this.dictionary.contains(w)).length
    };
    }

    /**
    * ========== 双向最大匹配 (Bidirectional Maximum Matching, BiMMM) ==========
    * 融合策略:综合 FMM 和 RMM 的结果,选择最优解
    * 策略优先级:
    * 1. 词数少者优(期望更长的词被正确切分)
    * 2. 词数相同 → 总词长(单字总数)少者优
    * 3. 词数相同且词长相同 → 优先选择 FMM 结果
    */

    bidirectionalMatch(text: string): SegmentResult {
    const fmmResult = this.forwardMatch(text);
    const rmmResult = this.backwardMatch(text);

    const fmm = fmmResult.words;
    const rmm = rmmResult.words;

    // 策略比较
    if (fmm.length !== rmm.length) {
    // 词数少者优
    return fmm.length < rmm.length ? fmmResult : rmmResult;
    }

    // 词数相同,比较总词长(OOV 越少越优)
    const fmmOov = fmm.filter(w => !this.dictionary.contains(w)).length;
    const rmmOov = rmm.filter(w => !this.dictionary.contains(w)).length;
    if (fmmOov !== rmmOov) {
    return fmmOov < rmmOov ? fmmResult : rmmResult;
    }

    // 完全相同,随机选一个(通常很少见)
    return fmmResult;
    }

    /**
    * 统一分词接口
    */

    segment(text: string, mode: SegmentMode = SegmentMode.BIDIRECTIONAL): SegmentResult {
    switch (mode) {
    case SegmentMode.FORWARD:
    return this.forwardMatch(text);
    case SegmentMode.BACKWARD:
    return this.backwardMatch(text);
    case SegmentMode.BIDIRECTIONAL:
    return this.bidirectionalMatch(text);
    default:
    return this.bidirectionalMatch(text);
    }
    }

    /**
    * 批量分词
    */

    segmentBatch(texts: string[], mode: SegmentMode = SegmentMode.BIDIRECTIONAL): SegmentResult[] {
    return texts.map(text => this.segment(text, mode));
    }

    /**
    * 打印分词过程(用于调试)
    */

    debugSegment(text: string): void {
    const fmm = this.forwardMatch(text);
    const rmm = this.backwardMatch(text);
    const bi = this.bidirectionalMatch(text);

    console.info(`【分词调试】原文: "${text}"`);
    console.info(`FMM: ${fmm.words.join('/')} (词数=${fmm.words.length}, OOV=${fmm.oovCount})`);
    console.info(`RMM: ${rmm.words.join('/')} (词数=${rmm.words.length}, OOV=${rmm.oovCount})`);
    console.info(`BiMMM: ${bi.words.join('/')} (采用 ${bi.mode})`);
    }
    }

    /**
    * 工厂函数:快速创建分词器
    */

    export function createSegmenter(customWords?: Array<[string, number]>): MaxMatchSegmenter {
    const dict = new Dictionary();
    if (customWords) {
    dict.addWords(customWords);
    }
    return new MaxMatchSegmenter(dict);
    }


    3.5 TF-IDF 关键词提取 TFIDFExtractor.ets

    // TFIDFExtractor.ets
    // TF-IDF(Term Frequency – Inverse Document Frequency)关键词提取
    // 完全自实现,无需任何第三方数学库

    /**
    * 词项信息
    */

    interface TermInfo {
    term: string; // 词语
    tf: number; // 词频 TF
    idf: number; // 逆文档频率 IDF
    tfidf: number; // TF-IDF 综合得分
    df: number; // 文档频率(出现该词的文档数)
    }

    /**
    * 文档词汇表
    */

    interface DocumentVocabulary {
    [term: string]: number; // term -> 文档内出现次数
    }

    /**
    * TF-IDF 提取器
    */

    export class TFIDFExtractor {
    // 文档集合(分词后的文档列表)
    private documents: string[][] = [];
    // 全局词汇表(term -> 包含该词的文档数)
    private globalDF: Map<string, number> = new Map();
    // 全局 IDF 缓存
    private globalIDF: Map<string, number> = new Map();
    // 总文档数
    private totalDocs: number = 0;
    // 是否已训练
    private trained: boolean = false;

    /**
    * 添加训练文档
    * @param tokens 分词后的词语数组
    */

    addDocument(tokens: string[]): void {
    // 统计该文档内每个词的频率
    const localFreq: Map<string, number> = new Map();
    for (const token of tokens) {
    localFreq.set(token, (localFreq.get(token) ?? 0) + 1);
    }

    // 更新全局文档频率
    localFreq.forEach((_, term) => {
    this.globalDF.set(term, (this.globalDF.get(term) ?? 0) + 1);
    });

    this.documents.push(tokens);
    this.totalDocs++;
    this.trained = false; // 需要重新计算 IDF
    }

    /**
    * 批量添加文档(接受原始文本,需先分词)
    */

    addDocumentBatch(rawTexts: string[], tokenizer: (text: string) => string[]): void {
    for (const text of rawTexts) {
    const tokens = tokenizer(text);
    this.addDocument(tokens);
    }
    }

    /**
    * 计算 IDF
    * IDF = log(总文档数 / 包含该词的文档数 + 1)
    * 加1是为了防止除零,且使 IDF 更平滑
    */

    private computeIDF(term: string): number {
    if (this.globalIDF.has(term)) {
    return this.globalIDF.get(term)!;
    }
    const df = this.globalDF.get(term) ?? 0;
    const idf = Math.log((this.totalDocs + 1) / (df + 1)) + 1;
    this.globalIDF.set(term, idf);
    return idf;
    }

    /**
    * ========== TF 计算方式 ==========
    * 支持多种 TF 归一化策略:
    * – RAW: 原始词频
    * – LOG: 1 + log(tf)
    * – AUGMENTED: 0.5 + 0.5 * tf / max_tf(防止长文档 bias)
    */

    enum TFScoreType {
    RAW = 'RAW',
    LOG = 'LOG',
    AUGMENTED = 'AUGMENTED'
    }

    /**
    * 计算单个词的 TF
    */

    private computeTF(term: string, localFreq: Map<string, number>, type: TFScoreType): number {
    const tf = localFreq.get(term) ?? 0;
    if (tf === 0) return 0;

    switch (type) {
    case TFScoreType.RAW:
    return tf;
    case TFScoreType.LOG:
    return 1 + Math.log(tf);
    case TFScoreType.AUGMENTED: {
    const maxFreq = Math.max(Array.from(localFreq.values()));
    return 0.5 + 0.5 * tf / maxFreq;
    }
    default:
    return tf;
    }
    }

    /**
    * ========== 核心方法:为单个文档计算 TF-IDF 向量 ==========
    * @param docTokens 目标文档的分词结果
    * @param topN 返回前 N 个关键词(默认 10)
    * @param tfType TF 归一化方式
    */

    extractKeywords(
    docTokens: string[],
    topN: number = 10,
    tfType: TFScoreType = TFScoreType.AUGMENTED
    ): TermInfo[] {
    if (!this.trained && this.totalDocs === 0) {
    // 无训练数据,直接用频率
    return this.extractKeywordsByFreq(docTokens, topN);
    }

    // 构建该文档的局部词频表
    const localFreq: Map<string, number> = new Map();
    for (const token of docTokens) {
    localFreq.set(token, (localFreq.get(token) ?? 0) + 1);
    }

    // 计算每个词的 TF-IDF
    const results: TermInfo[] = [];
    localFreq.forEach((freq, term) => {
    const tf = this.computeTF(term, localFreq, tfType);
    const idf = this.totalDocs > 0 ? this.computeIDF(term) : 1;
    const tfidf = tf * idf;
    results.push({
    term,
    tf,
    idf,
    tfidf,
    df: this.globalDF.get(term) ?? 0
    });
    });

    // 按 TF-IDF 降序排列
    results.sort((a, b) => b.tfidf a.tfidf);
    return results.slice(0, topN);
    }

    /**
    * 无训练数据时的降级方案:纯词频提取
    */

    private extractKeywordsByFreq(docTokens: string[], topN: number): TermInfo[] {
    const freq: Map<string, number> = new Map();
    for (const token of docTokens) {
    freq.set(token, (freq.get(token) ?? 0) + 1);
    }
    const results: TermInfo[] = [];
    freq.forEach((tf, term) => {
    results.push({ term, tf, idf: 1, tfidf: tf, df: 1 });
    });
    results.sort((a, b) => b.tfidf a.tfidf);
    return results.slice(0, topN);
    }

    /**
    * 将文档转换为 TF-IDF 向量(用于后续文本分类/相似度计算)
    * @param docTokens 分词后的词语数组
    * @returns 稀疏向量表示 [{term, tfidf}, …]
    */

    toTFIDFVector(docTokens: string[]): Array<{ term: string; tfidf: number }> {
    const keywords = this.extractKeywords(docTokens, docTokens.length);
    return keywords.map(k => ({ term: k.term, tfidf: k.tfidf }));
    }

    /**
    * 获取文档集合的全局词汇表
    */

    getVocabulary(): string[] {
    return Array.from(this.globalDF.keys());
    }

    /**
    * 获取全局词数
    */

    getVocabularySize(): number {
    return this.globalDF.size;
    }

    /**
    * 重置(清空训练数据)
    */

    reset(): void {
    this.documents = [];
    this.globalDF.clear();
    this.globalIDF.clear();
    this.totalDocs = 0;
    this.trained = false;
    }
    }


    3.6 朴素贝叶斯分类器 NaiveBayesClassifier.ets

    // NaiveBayesClassifier.ets
    // 朴素贝叶斯文本分类器(简化版 Multinomial NB)
    // P(类别|文档) ∝ P(类别) × ∏ P(词项|类别)

    /**
    * 类别信息
    */

    interface ClassInfo {
    name: string;
    prior: number; // P(类别) — 先验概率
    wordFreq: Map<string, number>; // 该类别中每个词的出现次数
    totalWords: number; // 该类别的总词数
    docCount: number; // 该类别的文档数
    }

    /**
    * 分类结果
    */

    export interface ClassificationResult {
    predictedClass: string; // 预测类别
    confidence: number; // 置信度(归一化后的最大概率)
    scores: Map<string, number>; // 所有类别的原始得分
    }

    /**
    * 朴素贝叶斯分类器(多项式模型)
    */

    export class NaiveBayesClassifier {
    // 类别数据表
    private classes: Map<string, ClassInfo> = new Map();
    // 全局词汇表
    private vocabulary: Set<string> = new Set();
    // 总文档数
    private totalDocs: number = 0;
    // 全局词数(用于拉普拉斯平滑)
    private globalWordCount: number = 0;

    /**
    * 注册类别
    */

    registerClass(className: string): void {
    if (!this.classes.has(className)) {
    this.classes.set(className, {
    name: className,
    prior: 0,
    wordFreq: new Map(),
    totalWords: 0,
    docCount: 0
    });
    }
    }

    /**
    * ========== 训练阶段 ==========
    * 使用拉普拉斯平滑(加一平滑)防止零概率
    * P(词项|类别) = (该类中词项出现次数 + 1) / (该类总词数 + 全局词数)
    */

    train(className: string, tokens: string[]): void {
    // 确保类别存在
    this.registerClass(className);

    const classInfo = this.classes.get(className)!;
    classInfo.docCount++;
    this.totalDocs++;

    // 统计该文档中各词频(避免同一文档中同一词多次计入)
    const uniqueTokens = new Set(tokens);
    uniqueTokens.forEach(token => {
    this.vocabulary.add(token);
    classInfo.wordFreq.set(token, (classInfo.wordFreq.get(token) ?? 0) + 1);
    classInfo.totalWords++;
    });
    }

    /**
    * 批量训练
    */

    trainBatch(samples: Array<{ className: string; tokens: string[] }>): void {
    for (const sample of samples) {
    this.train(sample.className, sample.tokens);
    }
    // 训练后计算先验概率
    this.computePriors();
    this.globalWordCount = Array.from(this.classes.values())
    .reduce((sum, c) => sum + c.totalWords, 0);
    }

    /**
    * 计算各类别的先验概率 P(类别)
    */

    private computePriors(): void {
    this.classes.forEach(classInfo => {
    classInfo.prior = classInfo.docCount / this.totalDocs;
    });
    }

    /**
    * ========== 预测阶段 ==========
    * 对数域计算(避免概率连乘导致下溢)
    * log P(类别|文档) ≈ log P(类别) + Σ log P(词项|类别)
    */

    predict(tokens: string[]): ClassificationResult {
    const scores = new Map<string, number>();

    this.classes.forEach((classInfo, className) => {
    // log P(类别) — 先验
    let logProb = Math.log(classInfo.prior + 1e-10);

    // Σ log P(词项|类别) — 似然
    const uniqueTokens = new Set(tokens);
    uniqueTokens.forEach(token => {
    const wordCount = classInfo.wordFreq.get(token) ?? 0;
    // 拉普拉斯平滑
    const prob = (wordCount + 1) / (classInfo.totalWords + this.globalWordCount + 1);
    logProb += Math.log(prob + 1e-10);
    });

    scores.set(className, logProb);
    });

    // 找最大概率
    let maxClass = '';
    let maxScore = Infinity;
    scores.forEach((score, className) => {
    if (score > maxScore) {
    maxScore = score;
    maxClass = className;
    }
    });

    // 归一化为置信度(softmax 近似)
    const scoresArr = Array.from(scores.values());
    const maxLog = Math.max(scoresArr);
    const expScores = scoresArr.map(s => Math.exp(s maxLog));
    const sumExp = expScores.reduce((a, b) => a + b, 0);
    const confidence = Math.exp(maxScore maxLog) / sumExp;

    return {
    predictedClass: maxClass,
    confidence,
    scores
    };
    }

    /**
    * 获取各类别的详细信息(调试用)
    */

    getClassDetails(className: string): ClassInfo | null {
    return this.classes.get(className) ?? null;
    }

    /**
    * 获取最常见的 topN 词(按条件概率排序)
    */

    getTopWords(className: string, topN: number = 20): Array<[string, number]> {
    const classInfo = this.classes.get(className);
    if (!classInfo) return [];

    const words: Array<[string, number]> = [];
    classInfo.wordFreq.forEach((count, word) => {
    words.push([word, count]);
    });
    words.sort((a, b) => b[1] a[1]);
    return words.slice(0, topN);
    }

    /**
    * 重置分类器
    */

    reset(): void {
    this.classes.clear();
    this.vocabulary.clear();
    this.totalDocs = 0;
    this.globalWordCount = 0;
    }

    /**
    * 打印训练摘要
    */

    printSummary(): void {
    console.info('=== 朴素贝叶斯分类器训练摘要 ===');
    console.info(`总文档数: ${this.totalDocs}`);
    console.info(`词汇表大小: ${this.vocabulary.size}`);
    this.classes.forEach((info, name) => {
    console.info(`类别 "${name}": 文档=${info.docCount}, 先验=${info.prior.toFixed(3)}, 总词数=${info.totalWords}`);
    });
    }
    }


    3.7 文本相似度计算 TextSimilarity.ets

    // TextSimilarity.ets
    // 文本相似度计算:余弦相似度、Jaccard 系数、Jaro-Winkler 距离

    /**
    * 相似度结果
    */

    export interface SimilarityResult {
    cosine: number; // 余弦相似度(0~1)
    jaccard: number; // Jaccard 系数(0~1)
    jaroWinkler: number; // Jaro-Winkler 相似度(0~1)
    }

    /**
    * 向量工具(内联实现,避免引入数学库)
    */

    class VectorUtils {
    /**
    * 点积
    */

    static dot(a: number[], b: number[]): number {
    let sum = 0;
    const len = Math.min(a.length, b.length);
    for (let i = 0; i < len; i++) {
    sum += a[i] * b[i];
    }
    return sum;
    }

    /**
    * 向量模长
    */

    static norm(v: number[]): number {
    let sum = 0;
    for (let i = 0; i < v.length; i++) {
    sum += v[i] * v[i];
    }
    return Math.sqrt(sum);
    }
    }

    /**
    * 稀疏向量表示
    */

    type SparseVector = Map<string, number>;

    /**
    * 文本相似度计算器
    */

    export class TextSimilarity {
    /**
    * ========== 1. 余弦相似度 (Cosine Similarity) ==========
    * 衡量两个向量在方向上的接近程度,取值 [0, 1]
    * Cosine(θ) = (A · B) / (||A|| × ||B||)
    * 适合:TF-IDF 向量、词向量场景
    */

    static cosineSimilarity(vecA: SparseVector, vecB: SparseVector): number {
    // 构建稠密向量(使用全局词汇表)
    const allKeys = new Set([vecA.keys(), vecB.keys()]);
    if (allKeys.size === 0) return 0;

    const denseA: number[] = [];
    const denseB: number[] = [];
    allKeys.forEach(key => {
    denseA.push(vecA.get(key) ?? 0);
    denseB.push(vecB.get(key) ?? 0);
    });

    const dot = VectorUtils.dot(denseA, denseB);
    const normA = VectorUtils.norm(denseA);
    const normB = VectorUtils.norm(denseB);

    if (normA === 0 || normB === 0) return 0;
    return dot / (normA * normB);
    }

    /**
    * 简化版余弦相似度(用于分词后的词集合)
    * 将词集合转为 TF 向量,再计算余弦
    */

    static cosineByTokens(tokensA: string[], tokensB: string[]): number {
    const freqA = new Map<string, number>();
    const freqB = new Map<string, number>();

    for (const t of tokensA) freqA.set(t, <span class=\"token p

    统计学学习教程,从入门到精通,方差分析(15)

    master阅读(39)

    方差分析


    一、方差分析引论

    1.1 方差分析的基本概念

    方差分析(Analysis of Variance,简称 ANOVA)是由英国统计学家 R.A. Fisher 于20世纪20年代提出的一种统计分析方法。其核心思想是:将数据的总变异分解为不同来源的变异,通过比较各来源的变异大小来判断因素对试验指标是否有显著影响。

    1.2 为什么需要方差分析

    问题的提出: 假设有 kkk 个总体(或 kkk 种处理),我们想要检验它们的均值是否相等。

    H0:μ1=μ2=⋯=μkH_0: \\mu_1 = \\mu_2 = \\cdots = \\mu_kH0:μ1=μ2==μk

    为什么不能用两两 ttt 检验替代?

    若对 kkk 个总体进行两两 ttt 检验,共需进行 (k2)=k(k−1)2\\binom{k}{2} = \\frac{k(k-1)}{2}(2k)=2k(k1) 次比较。例如 k=5k=5k=5 时需要 10 次比较。

    每次检验的显著性水平为 α\\alphaα,则至少犯一次第一类错误的概率为:

    P(至少犯一次第I类错误)=1−(1−α)(k2)P(\\text{至少犯一次第I类错误}) = 1 – (1-\\alpha)^{\\binom{k}{2}}P(至少犯一次第I类错误)=1(1α)(2k)

    α=0.05\\alpha = 0.05α=0.05k=5k = 5k=5 时:

    1−(1−0.05)10=1−0.9510≈1−0.5987=0.40131 – (1-0.05)^{10} = 1 – 0.95^{10} \\approx 1 – 0.5987 = 0.40131(10.05)10=10.951010.5987=0.4013

    即犯第一类错误的概率高达约 40%,远超预设水平。因此需要方差分析这一整体检验方法。

    1.3 方差分析中的基本术语

    术语含义
    试验指标 试验中要考察的结果(响应变量)
    因素 影响试验指标的条件(用 A,B,…A, B, \\ldotsA,B, 表示)
    水平 因素所处的不同状态或等级(用 A1,A2,…A_1, A_2, \\ldotsA1,A2, 表示)
    处理 因素水平的组合
    单因素方差分析 只有一个因素变化,其余因素固定
    双因素方差分析 有两个因素同时变化

    1.4 方差分析的基本假定

    对每一个总体(水平),要求:

  • 正态性: 各总体的观测值服从正态分布,即 Xij∼N(μi,σ2)X_{ij} \\sim N(\\mu_i, \\sigma^2)XijN(μi,σ2)
  • 方差齐性(等方差性): 各总体的方差相等,即 σ12=σ22=⋯=σk2=σ2\\sigma_1^2 = \\sigma_2^2 = \\cdots = \\sigma_k^2 = \\sigma^2σ12=σ22==σk2=σ2
  • 独立性: 各观测值相互独立

  • 二、单因素方差分析

    2.1 问题与模型

    设因素 AAAkkk 个水平 A1,A2,…,AkA_1, A_2, \\ldots, A_kA1,A2,,Ak,在水平 AiA_iAi 下进行 nin_ini 次独立试验,得到观测值 XijX_{ij}Xijj=1,2,…,nij = 1, 2, \\ldots, n_ij=1,2,,ni)。

    线性统计模型:

    Xij=μi+εij,i=1,2,…,k;j=1,2,…,niX_{ij} = \\mu_i + \\varepsilon_{ij}, \\quad i = 1, 2, \\ldots, k; \\quad j = 1, 2, \\ldots, n_iXij=μi+εij,i=1,2,,k;j=1,2,,ni

    其中 εij∼N(0,σ2)\\varepsilon_{ij} \\sim N(0, \\sigma^2)εijN(0,σ2) 且相互独立。

    引入效应参数:

    令总均值 μ=1n∑i=1kniμi\\mu = \\frac{1}{n}\\sum_{i=1}^{k} n_i \\mu_iμ=n1i=1kniμi(其中 n=∑i=1knin = \\sum_{i=1}^{k} n_in=i=1kni 为总观测次数),定义水平 AiA_iAi 的效应为:

    αi=μi−μ,i=1,2,…,k\\alpha_i = \\mu_i – \\mu, \\quad i = 1, 2, \\ldots, kαi=μiμ,i=1,2,,k

    则模型改写为:

    Xij=μ+αi+εij,∑i=1kniαi=0\\boxed{X_{ij} = \\mu + \\alpha_i + \\varepsilon_{ij}, \\quad \\sum_{i=1}^{k} n_i \\alpha_i = 0}Xij=μ+αi+εij,i=1kniαi=0

    待检验的假设为:

    H0:μ1=μ2=⋯=μk⇔H0:α1=α2=⋯=αk=0H_0: \\mu_1 = \\mu_2 = \\cdots = \\mu_k \\quad \\Leftrightarrow \\quad H_0: \\alpha_1 = \\alpha_2 = \\cdots = \\alpha_k = 0H0:μ1=μ2==μkH0:α1=α2==αk=0

    H1:μ1,μ2,…,μk 不全相等H_1: \\mu_1, \\mu_2, \\ldots, \\mu_k \\text{ 不全相等}H1:μ1,μ2,,μk 不全相等

    2.2 平方和分解(核心推导)

    2.2.1 定义基本统计量
    • iii 组样本均值:Xˉi⋅=1ni∑j=1niXij\\bar{X}_{i\\cdot} = \\frac{1}{n_i}\\sum_{j=1}^{n_i} X_{ij}Xˉi=ni1j=1niXij
    • 总样本均值:Xˉ⋅⋅=1n∑i=1k∑j=1niXij\\bar{X}_{\\cdot\\cdot} = \\frac{1}{n}\\sum_{i=1}^{k}\\sum_{j=1}^{n_i} X_{ij}Xˉ⋅⋅=n1i=1kj=1niXij

    重要恒等式验证:

    ∑i=1kniXˉi⋅=∑i=1kni⋅1ni∑j=1niXij=∑i=1k∑j=1niXij=nXˉ⋅⋅\\sum_{i=1}^{k} n_i \\bar{X}_{i\\cdot} = \\sum_{i=1}^{k} n_i \\cdot \\frac{1}{n_i}\\sum_{j=1}^{n_i} X_{ij} = \\sum_{i=1}^{k}\\sum_{j=1}^{n_i} X_{ij} = n\\bar{X}_{\\cdot\\cdot}i=1kniXˉi=i=1knini1j=1niXij=i=1kj=1niXij=nXˉ⋅⋅

    2.2.2 总离差平方和的分解

    定义总离差平方和:

    ST=∑i=1k∑j=1ni(Xij−Xˉ⋅⋅)2S_T = \\sum_{i=1}^{k}\\sum_{j=1}^{n_i}(X_{ij} – \\bar{X}_{\\cdot\\cdot})^2ST=i=1kj=1ni(XijXˉ⋅⋅)2

    分解过程: 对 Xij−Xˉ⋅⋅X_{ij} – \\bar{X}_{\\cdot\\cdot}XijXˉ⋅⋅ 进行"加一项减一项":

    Xij−Xˉ⋅⋅=(Xij−Xˉi⋅)+(Xˉi⋅−Xˉ⋅⋅)X_{ij} – \\bar{X}_{\\cdot\\cdot} = (X_{ij} – \\bar{X}_{i\\cdot}) + (\\bar{X}_{i\\cdot} – \\bar{X}_{\\cdot\\cdot})XijXˉ⋅⋅=(XijXˉi)+(XˉiXˉ⋅⋅)

    两边平方后求和:

    ST=∑i=1k∑j=1ni[(Xij−Xˉi⋅)+(Xˉi⋅−Xˉ⋅⋅)]2S_T = \\sum_{i=1}^{k}\\sum_{j=1}^{n_i}\\left[(X_{ij} – \\bar{X}_{i\\cdot}) + (\\bar{X}_{i\\cdot} – \\bar{X}_{\\cdot\\cdot})\\right]^2ST=i=1kj=1ni[(XijXˉi)+(XˉiXˉ⋅⋅)]2

    展开平方项:

    ST=∑i=1k∑j=1ni(Xij−Xˉi⋅)2⏟SE+∑i=1k∑j=1ni(Xˉi⋅−Xˉ⋅⋅)2⏟SA+2∑i=1k∑j=1ni(Xij−Xˉi⋅)(Xˉi⋅−Xˉ⋅⋅)S_T = \\underbrace{\\sum_{i=1}^{k}\\sum_{j=1}^{n_i}(X_{ij} – \\bar{X}_{i\\cdot})^2}_{S_E} + \\underbrace{\\sum_{i=1}^{k}\\sum_{j=1}^{n_i}(\\bar{X}_{i\\cdot} – \\bar{X}_{\\cdot\\cdot})^2}_{S_A} + 2\\sum_{i=1}^{k}\\sum_{j=1}^{n_i}(X_{ij} – \\bar{X}_{i\\cdot})(\\bar{X}_{i\\cdot} – \\bar{X}_{\\cdot\\cdot})ST=SEi=1kj=1ni(XijXˉi)2+SAi=1kj=1ni(XˉiXˉ⋅⋅)2+2i=1kj=1ni(XijXˉi)(XˉiXˉ⋅⋅)

    证明交叉项为零:

    ∑i=1k∑j=1ni(Xij−Xˉi⋅)(Xˉi⋅−Xˉ⋅⋅)=∑i=1k(Xˉi⋅−Xˉ⋅⋅)∑j=1ni(Xij−Xˉi⋅)\\sum_{i=1}^{k}\\sum_{j=1}^{n_i}(X_{ij} – \\bar{X}_{i\\cdot})(\\bar{X}_{i\\cdot} – \\bar{X}_{\\cdot\\cdot}) = \\sum_{i=1}^{k}(\\bar{X}_{i\\cdot} – \\bar{X}_{\\cdot\\cdot})\\sum_{j=1}^{n_i}(X_{ij} – \\bar{X}_{i\\cdot})i=1kj=1ni(XijXˉi)(XˉiXˉ⋅⋅)=i=1k(XˉiXˉ⋅⋅)j=1ni(XijXˉi)

    由于 ∑j=1ni(Xij−Xˉi⋅)=∑j=1niXij−niXˉi⋅=niXˉi⋅−niXˉi⋅=0\\sum_{j=1}^{n_i}(X_{ij} – \\bar{X}_{i\\cdot}) = \\sum_{j=1}^{n_i}X_{ij} – n_i\\bar{X}_{i\\cdot} = n_i\\bar{X}_{i\\cdot} – n_i\\bar{X}_{i\\cdot} = 0j=1ni(XijXˉi)=j=1niXijniXˉi=niXˉiniXˉi=0

    因此交叉项为零,得到:

    ST=SE+SA\\boxed{S_T = S_E + S_A}ST=SE+SA

    其中:

    SA=∑i=1k∑j=1ni(Xˉi⋅−Xˉ⋅⋅)2=∑i=1kni(Xˉi⋅−Xˉ⋅⋅)2S_A = \\sum_{i=1}^{k}\\sum_{j=1}^{n_i}(\\bar{X}_{i\\cdot} – \\bar{X}_{\\cdot\\cdot})^2 = \\sum_{i=1}^{k}n_i(\\bar{X}_{i\\cdot} – \\bar{X}_{\\cdot\\cdot})^2SA=i=1kj=1ni(XˉiXˉ⋅⋅)2=i=1kni(XˉiXˉ⋅⋅)2

    SE=∑i=1k∑j=1ni(Xij−Xˉi⋅)2S_E = \\sum_{i=1}^{k}\\sum_{j=1}^{n_i}(X_{ij} – \\bar{X}_{i\\cdot})^2SE=i=1kj=1ni(XijXˉi)2

    各项含义:

    平方和名称含义
    STS_TST 总离差平方和 数据总的变异程度
    SAS_ASA 因素 AAA 的离差平方和(组间平方和) 由因素 AAA 不同水平引起的变异
    SES_ESE 误差平方和(组内平方和) 由随机误差引起的变异
    2.2.3 自由度的分解

    ST 的自由度:fT=n−1S_T \\text{ 的自由度:} \\quad f_T = n – 1ST 的自由度:fT=n1

    SA 的自由度:fA=k−1S_A \\text{ 的自由度:} \\quad f_A = k – 1SA 的自由度:fA=k1

    SE 的自由度:fE=n−kS_E \\text{ 的自由度:} \\quad f_E = n – kSE 的自由度:fE=nk

    验证: fA+fE=(k−1)+(n−k)=n−1=fTf_A + f_E = (k-1) + (n-k) = n – 1 = f_TfA+fE=(k1)+(nk)=n1=fT

    2.3 各平方和期望值的推导

    2.3.1 E(SE)E(S_E)E(SE) 的推导

    SE=∑i=1k∑j=1ni(Xij−Xˉi⋅)2S_E = \\sum_{i=1}^{k}\\sum_{j=1}^{n_i}(X_{ij} – \\bar{X}_{i\\cdot})^2SE=i=1kj=1<span class

    行列式杂题第二弹

    master阅读(42)

    文章目录

    • 前言
    • 题目梗概
          • 1000题
    • 参考解析
      • 1.1
      • 1.2
      • 1.3
      • 2.1
      • 2.2
      • 2.3
      • 2.4
    • 后话

    前言

    中档题

    题目梗概

  • 行列式1.1 行列式多项式x³项系数求解 设四阶行列式

    f

    (

    x

    )

    =

    x

    x

    1

    2

    x

    1

    x

    2

    1

    2

    1

    x

    1

    2

    1

    1

    x

    f(x) = \\begin{vmatrix} x & x & 1 & 2x \\\\ 1 & x & 2 & -1 \\\\ 2 & 1 & x & 1 \\\\ 2 & -1 & 1 & x \\end{vmatrix}

    f(x)=

    x122xx1112x12x11x

    f

    (

    x

    )

    f(x)

    f(x)

    x

    3

    x^3

    x3项的系数。

  • 行列式1.2 四阶三对角行列式行变换化上三角 计算四阶行列式

    D

    4

    =

    1

    a

    0

    0

    1

    2

    a

    b

    0

    0

    2

    3

    b

    c

    0

    0

    3

    4

    c

    D_4 = \\begin{vmatrix} 1 & a & 0 & 0 \\\\ -1 & 2-a & b & 0 \\\\ 0 & -2 & 3-b & c \\\\ 0 & 0 & -3 & 4-c \\end{vmatrix}

    D4=

    1100a2a200b3b300c4c

  • 行列式1.3 矩阵线性变换求行列式 已知

    α

    1

    ,

    α

    2

    ,

    α

    3

    \\boldsymbol{\\alpha}_1,\\boldsymbol{\\alpha}_2,\\boldsymbol{\\alpha}_3

    α1,α2,α3是三维线性无关列向量,

    A

    \\boldsymbol{A}

    A是3阶矩阵,满足

    {

    A

    α

    1

    =

    α

    2

    2

    α

    3

    A

    α

    2

    =

    α

    1

    α

    2

    +

    2

    α

    3

    A

    α

    3

    =

    2

    α

    1

    +

    α

    2

    \\begin{cases} \\boldsymbol{A}\\boldsymbol{\\alpha}_1 = \\boldsymbol{\\alpha}_2 – 2\\boldsymbol{\\alpha}_3 \\\\ \\boldsymbol{A}\\boldsymbol{\\alpha}_2 = \\boldsymbol{\\alpha}_1 – \\boldsymbol{\\alpha}_2 + 2\\boldsymbol{\\alpha}_3 \\\\ \\boldsymbol{A}\\boldsymbol{\\alpha}_3 = 2\\boldsymbol{\\alpha}_1 + \\boldsymbol{\\alpha}_2 \\end{cases}

    Aα1=α22α3Aα2=α1α2+2α3Aα3=2α1+α2

    A

    |\\boldsymbol{A}|

    A

  • 1000题
  • 行列式2.1 三阶行列式函数零点个数判定 设

    a

    i

    ,

    b

    i

    ,

    c

    i

    a_i,b_i,c_i

    ai,bi,ci为常数,

    f

    (

    x

    )

    f(x)

    f(x)不恒为0,讨论函数

    f

    (

    x

    )

    =

    a

    1

    +

    x

    b

    1

    +

    x

    c

    1

    +

    x

    a

    2

    +

    x

    b

    2

    +

    x

    c

    2

    +

    x

    a

    3

    +

    x

    b

    3

    +

    x

    c

    3

    +

    x

    f(x) = \\begin{vmatrix} a_1+x & b_1+x & c_1+x \\\\ a_2+x & b_2+x & c_2+x \\\\ a_3+x & b_3+x & c_3+x \\end{vmatrix}

    f(x)=

    a1+xa2+xa3+xb1+xb2+xb3+xc1+xc2+xc3+x

    的零点个数。

  • 行列式2.2 n阶三对角行列式递推计算 计算n阶三对角行列式

    D

    n

    =

    2

    1

    0

    0

    1

    2

    1

    0

    0

    1

    2

    0

    0

    0

    0

    2

    D_n = \\begin{vmatrix} 2 & 1 & 0 & \\cdots & 0 \\\\ 1 & 2 & 1 & \\cdots & 0 \\\\ 0 & 1 & 2 & \\cdots & 0 \\\\ \\vdots & \\vdots & \\vdots & \\ddots & \\vdots \\\\ 0 & 0 & 0 & \\cdots & 2 \\end{vmatrix}

    Dn=

    2100121001200002

  • 行列式2.3 两个同结构行列式的差 设

    D

    1

    =

    2

    1

    0

    1

    1

    2

    5

    3

    3

    0

    a

    b

    1

    3

    5

    0

    ,

    D

    2

    =

    2

    1

    0

    1

    1

    2

    5

    3

    3

    0

    a

    b

    1

    1

    1

    0

    D_1 = \\begin{vmatrix} 2 & 1 & 0 & -1 \\\\ -1 & 2 & -5 & 3 \\\\ 3 & 0 & a & b \\\\ 1 & -3 & 5 & 0 \\end{vmatrix}, \\quad D_2 = \\begin{vmatrix} 2 & 1 & 0 & -1 \\\\ -1 & 2 & -5 & 3 \\\\ 3 & 0 & a & b \\\\ 1 & -1 & 1 & 0 \\end{vmatrix}

    D1=

    2131120305a513b0

    ,D2=

    2131120105a113b0

    D

    1

    D

    2

    D_1 – D_2

    D1D2

  • 行列式2.4 向量组线性变换求矩阵行列式(双解法) 设

    A

    \\boldsymbol{A}

    A是3阶矩阵,

    α

    1

    ,

    α

    2

    ,

    α

    3

    \\boldsymbol{\\alpha}_1,\\boldsymbol{\\alpha}_2,\\boldsymbol{\\alpha}_3

    α1,α2,α3是三维线性无关列向量,且满足

    {

    A

    α

    1

    =

    α

    1

    +

    2

    α

    2

    +

    α

    3

    A

    α

    2

    =

    2

    α

    1

    +

    α

    2

    +

    α

    3

    A

    α

    3

    =

    α

    1

    +

    α

    2

    +

    2

    α

    3

    \\begin{cases} \\boldsymbol{A}\\boldsymbol{\\alpha}_1 = \\boldsymbol{\\alpha}_1 + 2\\boldsymbol{\\alpha}_2 + \\boldsymbol{\\alpha}_3 \\\\ \\boldsymbol{A}\\boldsymbol{\\alpha}_2 = 2\\boldsymbol{\\alpha}_1 + \\boldsymbol{\\alpha}_2 + \\boldsymbol{\\alpha}_3 \\\\ \\boldsymbol{A}\\boldsymbol{\\alpha}_3 = \\boldsymbol{\\alpha}_1 + \\boldsymbol{\\alpha}_2 + 2\\boldsymbol{\\alpha}_3 \\end{cases}

    Aα1=α1+2α2+α3Aα2=2α1+α2+α3Aα3=α1+α2+2α3

    A

    |\\boldsymbol{A}|

    A

  • 参考解析

    1.1

    考察点:行列式按行展开、多项式次数与系数分析

    题目:设四阶行列式

    f

    (

    x

    )

    =

    x

    x

    1

    2

    x

    1

    x

    2

    1

    2

    1

    x

    1

    2

    1

    1

    x

    f(x) = \\begin{vmatrix} x & x & 1 & 2x \\\\ 1 & x & 2 & -1 \\\\ 2 & 1 & x & 1 \\\\ 2 & -1 & 1 & x \\end{vmatrix}

    f(x)=

    x122xx1112x12x11x

    f

    (

    x

    )

    f(x)

    f(x)

    x

    3

    x^3

    x3项的系数。

    解: 将

    f

    (

    x

    )

    f(x)

    f(x)按第1行展开,得四项:

    f

    (

    x

    )

    =

    x

    A

    11

    +

    x

    A

    12

    +

    1

    A

    13

    +

    2

    x

    A

    14

    f(x) = x\\cdot A_{11} + x\\cdot A_{12} + 1\\cdot A_{13} + 2x\\cdot A_{14}

    f(x)=xA11+xA12+1A13+2xA14 其中

    A

    i

    j

    =

    (

    1

    )

    i

    +

    j

    M

    i

    j

    A_{ij}=(-1)^{i+j}M_{ij}

    Aij=(1)i+jMij为代数余子式,逐项分析

    x

    3

    x^3

    x3项的贡献:

  • 第1项

    x

    A

    11

    x\\cdot A_{11}

    xA11

    M

    11

    M_{11}

    M11为去掉第1行第1列的三阶行列式:

    M

    11

    =

    x

    2

    1

    1

    x

    1

    1

    1

    x

    =

    x

    3

    4

    x

    3

    M_{11} = \\begin{vmatrix}x & 2 & -1 \\\\ 1 & x & 1 \\\\ -1 & 1 & x\\end{vmatrix} = x^3 – 4x – 3

    M11=

    x112x111x

    =x34x3 最高次为

    x

    3

    x^3

    x3,乘以

    x

    x

    x后最高次为

    x

    4

    x^4

    x4,无

    x

    3

    x^3

    x3项贡献。

  • 第2项

    x

    A

    12

    x\\cdot A_{12}

    xA12

    M

    12

    M_{12}

    M12为去掉第1行第2列的三阶行列式:

    M

    12

    =

    1

    2

    1

    2

    x

    1

    2

    1

    x

    =

    x

    2

    2

    x

    +

    1

    M_{12} = \\begin{vmatrix}1 & 2 & -1 \\\\ 2 & x & 1 \\\\ 2 & 1 & x\\end{vmatrix} = x^2 – 2x + 1

    M12=

    1222x111x

    =x22x+1

    A

    12

    =

    M

    12

    A_{12} = -M_{12}

    A12=M12,乘以

    x

    x

    x后得

    x

    3

    +

    2

    x

    2

    x

    -x^3 + 2x^2 – x

    x3+2x2x

    x

    3

    x^3

    x3项系数为

    1

    -1

    1

  • 第3项

    1

    A

    13

    1\\cdot A_{13}

    1A13

    M

    13

    M_{13}

    M13为去掉第1行第3列的三阶行列式,最高次为

    x

    2

    x^2

    x2,乘以常数1后无

    x

    3

    x^3

    x3项贡献。

  • 第4项

    2

    x

    A

    14

    2x\\cdot A_{14}

    2xA14

    M

    14

    M_{14}

    M14为去掉第1行第4列的三阶行列式:

    M

    14

    =

    1

    x

    2

    2

    1

    x

    2

    1

    1

    =

    2

    x

    2

    x

    7

    M_{14} = \\begin{vmatrix}1 & x & 2 \\\\ 2 & 1 & x \\\\ 2 & -1 & 1\\end{vmatrix} = 2x^2 – x – 7

    M14=

    122x112x1

    =2x2x7

    A

    14

    =

    M

    14

    A_{14} = -M_{14}

    A14=M14,乘以

    2

    x

    2x

    2x后得

    4

    x

    3

    +

    2

    x

    2

    +

    14

    x

    -4x^3 + 2x^2 + 14x

    4x3+2x2+14x

    x

    3

    x^3

    x3项系数为

    4

    -4

    4

  • 合并所有

    x

    3

    x^3

    x3项系数:

    1

    +

    (

    4

    )

    =

    5

    -1 + (-4) = -5

    1+(4)=5

    x

    3

    项的系数为

    5

    \\boxed{x^3项的系数为 -5}

    x3项的系数为5

    1.2

    考察点:行列式初等行变换、上三角行列式计算

    题目:计算四阶行列式

    D

    4

    =

    1

    a

    0

    0

    1

    2

    a

    b

    0

    0

    2

    3

    b

    c

    0

    0

    3

    4

    c

    D_4 = \\begin{vmatrix} 1 & a & 0 & 0 \\\\ -1 & 2-a & b & 0 \\\\ 0 & -2 & 3-b & c \\\\ 0 & 0 & -3 & 4-c \\end{vmatrix}

    D4=

    1100a2a200b3b300c4c

    解: 依次做行变换消去下三角元素:

    • 第2行加上第1行:

      r

      2

      =

      r

      2

      +

      r

      1

      r_2 = r_2 + r_1

      r2=r2+r1

    • 第3行加上新的第2行:

      r

      3

      =

      r

      3

      +

      r

      2

      r_3 = r_3 + r_2

      r3=r3+r2

    • 第4行加上新的第3行:

      r

      4

      =

      r

      4

      +

      r

      3

      r_4 = r_4 + r_3

      r4=r4+r3

    化简后得到上三角行列式:

    D

    4

    =

    1

    a

    0

    0

    0

    2

    b

    0

    0

    0

    3

    c

    0

    0

    0

    4

    D_4 = \\begin{vmatrix} 1 & a & 0 & 0 \\\\ 0 & 2 & b & 0 \\\\ 0 & 0 & 3 & c \\\\ 0 & 0 & 0 & 4 \\end{vmatrix}

    D4=

    1000a2000b3000c4

    主对角线元素相乘得结果:

    D

    4

    =

    1

    ×

    2

    ×

    3

    ×

    4

    =

    24

    D_4 = 1 \\times 2 \\times 3 \\times 4 = 24

    D4=1×2×3×4=24

    D

    4

    =

    24

    \\boxed{D_4 = 24}

    D4=24

    1.3

    考察点:相似矩阵行列式、向量组线性无关性、矩阵乘法

    题目:已知

    α

    1

    ,

    α

    2

    ,

    α

    3

    \\boldsymbol{\\alpha}_1,\\boldsymbol{\\alpha}_2,\\boldsymbol{\\alpha}_3

    α1,α2,α3是三维线性无关列向量,

    A

    \\boldsymbol{A}

    A是3阶矩阵,满足

    {

    A

    α

    1

    =

    α

    2

    2

    α

    3

    A

    α

    2

    =

    α

    1

    α

    2

    +

    2

    α

    3

    A

    α

    3

    =

    2

    α

    1

    +

    α

    2

    \\begin{cases} \\boldsymbol{A}\\boldsymbol{\\alpha}_1 = \\boldsymbol{\\alpha}_2 – 2\\boldsymbol{\\alpha}_3 \\\\ \\boldsymbol{A}\\boldsymbol{\\alpha}_2 = \\boldsymbol{\\alpha}_1 – \\boldsymbol{\\alpha}_2 + 2\\boldsymbol{\\alpha}_3 \\\\ \\boldsymbol{A}\\boldsymbol{\\alpha}_3 = 2\\boldsymbol{\\alpha}_1 + \\boldsymbol{\\alpha}_2 \\end{cases}

    Aα1=α22α3Aα2=α1α2+2α3Aα3=2α1+α2<span class=\"vlist\" st

    vllm源码剖析13-vLLM 分布式推理-张量并行详解

    master阅读(34)

    文章目录

      • 一 AllReduce 原理
        • 1.1 Ring-AllReduce 算子原理
          • 1.1.1 Reduce-Scatter
          • 1.1.2 All-Gather
        • 1.2 Ring AllReduce 的通信成本(重)
      • 二 Transformer 模型的张量并行
        • 2.1 线性层权重不同切分方式
        • 2.2 MLP 层的张量并行
          • 2.2.1 拆分原理
          • 2.2.2 拆分方式
          • 2.2.3 MLA 层的通讯量分析(重)
        • 2.3 MHA 层的张量并行
          • 2.3.1 拆分原理
          • 2.3.2 MHA 层的通讯量分析
          • 2.3.3 attention 中的张量并行实例
        • 2.4 Embedding 层的张量并行
          • 2.4.1 输入嵌入层
          • 2.4.2 输出嵌入层
      • 三 vLLM 中的张量并行
        • 3.1 vLLM 中张量并行如何使用
        • 3.2 vLLM 中的分布式资源管理
          • 模拟分布式并行分分组算法
        • 3.3 vLLM 的并行线性层
          • ColumnParallelLinear 类源码剖析(列并行)
          • RowParallelLinear 类源码剖析
          • QKVParallelLinear 源码剖析
        • 3.4 vLLM 分布式推理的流程总结
          • 阶段一:初始化与模型加载
          • 阶段二:分布式前向传播 (Forward Pass)
          • 阶段三:Logits 聚合与 Token 采样

    张量并行(TP,Tensor Parallelism)的核心思路,是把模型中的大矩阵计算拆到多张 GPU 上完成。这样可以降低单张 GPU 的显存压力;当计算量足够大、通信开销可以被摊薄时,也可能提升 forward 吞吐。
    在 vLLM 的 Megatron 风格 TP 中,常见的切分位置主要有三类:

    • Embedding 层:例如 VocabParallelEmbedding,按词表维度切分 embedding 权重。每张卡只负责一部分 token,forward 后通过 All-Reduce 汇总完整 embedding。
    • 线性层:例如 MLP 里的 MergedColumnParallelLinear、RowParallelLinear,以及 Attention 里的 QKVParallelLinear 和 o_proj。Column Parallel 通常切输出维,Row Parallel 通常切输入维。
    • Attention:Q/K/V 通常按 head 维分到不同 rank 上,各 rank 计算本地 head 的 attention,最后的 o_proj 通常通过 Row Parallel 和 All-Reduce 合并结果。

    后面我们就按 VocabParallelEmbedding、Column Parallel、Row Parallel 和 LM Head 这几类层,分别看它们在 forward 中什么时候触发 All-Reduce、All-Gather 或 gather。 下面将结合具体并行层,说明这些集合通信在 forward 中的触发位置,以及它们如何配合 TP 的切分方式完成计算。

    一 AllReduce 原理

    All-Reduce 的目标,是将所有进程上的数据通过特定操作(如求和、取最大值)聚合后,把结果同步到每一个进程。常见底层实现包括 Ring AllReduce 和 Tree AllReduce。其最终目标,就是让每块 GPU 上的数据都变成汇总/归约后的同一个结果。 在 vLLM 的 TP 中,All-Reduce 典型出现在 Row Parallel 线性层和 Embedding 层:各 rank 先算出自己的部分结果,再通过 All-Reduce 求和,得到完整输出。具体我们会在本节的之后内容中讲到

    在这里插入图片描述

    1.1 Ring-AllReduce 算子原理

    Ring-AllReduce 的实现其实分为两个过程 Reduce-Scatter 和 All-Gather。

  • 分块传输 假设有

    N

    N

    N 个进程参与通信,通常也对应

    N

    N

    N 个 GPU。每个进程都有一份待归约的数据。Ring-AllReduce 会先把这份数据切成

    N

    N

    N 个 chunk,并让这些进程在逻辑上组成一个环。

  • Reduce-Scatter 阶段 在每一轮中,每个进程都会把一个 chunk 发送给环上的下一个进程,同时从上一个进程接收一个 chunk。收到 chunk 后,进程会把它和本地对应位置的 chunk 做逐元素归约(通常是求和)。经过

    N

    1

    N-1

    N1 轮后,每个进程都会得到一块已完成全局归约的结果分片。此时每个进程只保存完整结果的一部分。

  • All-Gather 阶段 接下来,各进程继续沿环传播这些已归约完成的结果分片。每一轮发送当前持有的某个结果分片,同时接收其他进程传来的结果分片。 再经过

    N

    1

    N-1

    N1 轮后,所有进程都能收集到全部归约分片,并把这些分片拼接成完整的 AllReduce 结果。

  • 1.1.1 Reduce-Scatter

    以 4 个 GPU 设备为例,可将它们组织成一个逻辑环,使每个 GPU 只与相邻 GPU 交换数据。待归约的数据被切成 4 个 chunk,每轮通信时,各 GPU 同时发送一个 chunk、接收一个 chunk,并将收到的 chunk 与本地对应位置的 chunk 做逐元素相加。 第一次通信和累加完成后,某些位置的 chunk 已包含两个 GPU 的部分结果;这些更新后的 chunk 会在后续轮次中继续传递和累加。经过 N-1 轮(此处 N=4,即 3 轮)后,每个 GPU 会得到一个已完成全局累加的结果分片。 这里每次通信的数据量是一个 chunk。若原始数据量为 K,进程数为 N,则每个 chunk 的大小约为

    K

    /

    N

    K/N

    K/N。Reduce-Scatter 阶段结束时,每个设备只持有完整结果的一部分;后续 All-Gather 阶段会将这些分片继续传播,使每个设备最终都获得完整结果。

    在这里插入图片描述 第二次累加完成后的示意图如下,同样,被更新的数据块,会作为下一次传递和累加的起点,继续参与下一轮的通信和计算。

    在这里插入图片描述 第三次累加完成后的示意图如下: 在这里插入图片描述 经过 3 次环形传递和规约后,每块 GPU 上都有一块数据拥有了对应位置完整的累加聚合(下图中红色块)。此时,Reduce-Scatter 通信阶段结束,进入 All-Gather 通信阶段。目标是将红色块继续沿环传播,并填充到其余 GPU 对应的位置上,使所有 GPU 最终都拥有全部数据。

    1.1.2 All-Gather

    All-Gather 通信操作依然遵循相邻 GPU 对应位置进行通讯的原则,但这一步不再做相加,而是将已归约好的分块拷贝到下一跳对应的位置上。All-Gather 以红色块作为起点,第一轮传递和填充完成后的示意图如下所示:

    在这里插入图片描述 同样的经过 3 轮更新,使得每块 GPU 上都汇总到了完整的数据,变成如下形式: 在这里插入图片描述

    1.2 Ring AllReduce 的通信成本(重)

    假设有

    N

    N

    N 个设备,原始数据总大小为

    K

    K

    K,在一次 AllReduce 过程中,进行了

    N

    1

    N-1

    N1 次 Scatter-Reduce 操作和

    N

    1

    N-1

    N1 次 Allgather 操作,又因为每一次操作所需要传递的数据大小为

    K

    /

    N

    K/N

    K/N,所以整个 AllReduce 过程所传输的数据大小为

    2

    (

    N

    1

    )

    K

    /

    N

    2(N-1) * K/N

    2(N1)K/N。随着 N 的增大,Ring AllReduce 通信算子的通信量可以近似为 2K。

    在张量并行加速时,Ring AllReduce 的吞吐通常受到环上最慢链路和实现开销的共同影响。每次传输的数据块只有

    K

    /

    N

    K/N

    K/N,所以当

    N

    N

    N 增大时,单轮数据块变小。对于大张量,Ring AllReduce 容易形成流水线并行,通信效率通常较高;但对于小张量,固定调度开销和协议开销的占比会明显上升,实际有效带宽利用率会下降。

    另外,Ring AllReduce 需要完成

    2

    ×

    (

    N

    1

    )

    2 \\times (N-1)

    2×(N1) 次通信(两个阶段各

    N

    1

    N-1

    N1 次)。例如 8 GPU 需要 14 轮通信,即 Reduce-Scatter 和 All-Gather 各 7 轮。小张量本身的计算时间通常很短,但多轮通信带来的延迟累积会成为瓶颈。因此,Ring AllReduce 一般不适合小张量频繁同步、且对延迟要求很高的场景。

    vLLM 在 CUDA all-reduce 路径中并不简单地固定使用一种传统 Ring-AllReduce。源码中会根据当前并行组、环境开关、world size、硬件拓扑、P2P 能力、张量大小、dtype、连续性等多种条件,尝试不同的通信实现,选择更合适的 all-reduce 路径,以降低延迟或提高带宽利用率。 以下伪代码用来抽象说明低延迟 AllReduce 的核心思想:各 GPU 通过预先建立的通信资源访问其他 rank 的数据,并完成逐元素求和,使每个 GPU 最终得到与 AllReduce(sum) 等价的全局规约结果。

    def low_latency_all_reduce(x):
    """
    逻辑效果等价于 AllReduce(sum)。
    每个 GPU 最终都会得到所有 rank 输入张量的逐元素求和结果。
    """

    # 所有 GPU 事先完成必要的通信资源初始化,
    # 例如注册 buffer、创建 workspace,或建立 symmetric memory 句柄。
    peer_buffers = get_all_peer_buffers()

    # 同步和就绪检查省略。
    out = zeros_like(x)

    for buf in peer_buffers:
    # 抽象表示:读取各个 rank 的输入数据,
    # 并在 GPU 上完成逐元素累加。
    out += load_from(buf)

    return out

    从逻辑效果上看,每个 GPU 都会执行一次等价的规约过程,并得到与 AllReduce(sum) 相同的输出。和 Ring AllReduce 相比,这类低延迟路径不强调多轮分块传递,而是利用预先建立的通信资源、共享缓冲区等资源,尽量减少通信轮次和 kernel 调度次数,从而降低小张量频繁同步时的总延迟

    二 Transformer 模型的张量并行

    decoder-only 架构的 LLM 中标准的 transformer 层如图 2 所示,其由一个自注意力(self-attention)模块和一个两层的多层感知机 (MLP)组成,可在这两个模块中分别引入模型并行(也叫张量并行)技术。 在这里插入图片描述 基于 transformer 网络 pytorch 代码的基础上,通常只需添加几个同步操作代码(synchronization primitives),就可实现一个简单的模型并行方案。下文我将会描述 Megatron-LM 的张量并行算法原理,以及在 transformer 模型中的应用。

    2.1 线性层权重不同切分方式

    张量并行的底层逻辑,可以先从矩阵乘法的拆分计算来理解。下面分别介绍线性层中的列并行和行并行,以及它们对应的矩阵形式。 其中

    X

    X

    X

    是输入,

    是输入,

    是输入,A

    是权重矩阵。按这个数学形式,

    是权重矩阵。按这个数学形式,

    是权重矩阵。按这个数学形式,A$的第一维对应输入维,第二维对应输出维。后面说的“列并行”和“行并行”,都基于这个矩阵形式来理解 在这里插入图片描述

    线性层是 Transformer 中最主要的 GEMM 来源之一,既出现在 MLP 中,也出现在 Attention 的 Q/K/V projection 和 output projection 中。开启张量并行后,vLLM 会让不同 TP rank 加载各自负责的权重分片,并在 forward 时分别执行本地 GEMM。 按不同方式切分权重后(TP=2)的线性层推理 forward 操作的可视化对比图如下图所示:

    在这里插入图片描述

    • 权重

      A

      A

      A 按列维度切分时,对应的是 ColumnParallelLinear。每个 GPU 负责计算自己的局部输出

      Y

      i

      =

      X

      A

      i

      Y_i = X A_i

      Yi=XAi。默认情况下,各 GPU 只保留自己的局部结果;若 gather_output=True,则在 forward 结束时执行 AllGather,将各分片拼成完整输出

      Y

      Y

      Y

    • 权重

      A

      A

      A 按行维度切分时,对应的是 RowParallelLinear。输入

      X

      X

      X 也需要沿最后一维切分到不同 GPU 上;若输入本身已是并行的,可直接使用,否则会在 forward 中先切分输入。每个 GPU 计算自己的局部结果后需要执行 AllReduce,将各 GPU 的部分结果相加,得到完整输出

      Y

      Y

      Y

    2.2 MLP 层的张量并行

    MLP 的张量并行实现相对 Self-Attention 简单。

    2.2.1 拆分原理

    MLP 模块的第一个操作是通用矩阵乘法 (GEMM),随后是一个 GeLU 非线性激活函数,计算公式如下所示。

    Y

    =

    G

    e

    L

    U

    (

    X

    A

    )

    Y = GeLU(XA)

    Y=GeLU(XA) 后续最新的 Llama 及 qwen 系列的 llm 用的激活函数是

    S

    i

    L

    U

    SiLU

    SiLU。 为了实现 GEMM 的并行计算。

  • 第一个方法是将权重矩阵

    A

    A

    A 按行拆分,同时将输入矩阵

    X

    X

    X 按列拆分,如下图所示:

    X

    =

    [

    X

    1

    ,

    X

    2

    ]

    ,

    A

    =

    [

    A

    1

    A

    2

    ]

    X = [X_1, X_2],\\quad A = \\begin{bmatrix} A_1 \\\\ A_2 \\end{bmatrix}

    X=[X1,X2],A=[A1A2] 这种切分方式会得到

    Y

    =

    GeLU

    (

    X

    1

    A

    1

    +

    X

    2

    A

    2

    )

    Y = \\text{GeLU}(X_1A_1 + X_2A_2)

    Y=GeLU(X1A1+X2A2),因为 GeLU 是一个非线性函数,所以

    GeLU

    (

    X

    1

    A

    1

    +

    X

    2

    A

    2

    )

    GeLU

    (

    X

    1

    A

    1

    )

    +

    GeLU

    (

    X

    2

    A

    2

    )

    \\text{GeLU}(X_1A_1 + X_2A_2) \\neq \\text{GeLU}(X_1A_1) + \\text{GeLU}(X_2A_2)

    GeLU(X1A1+X2A2)=GeLU(X1A1)+GeLU(X2A2)。因此这种方法在 GeLU 函数之前需要一个同步点。所谓同步点是指,在执行 GeLU 函数之前,需要等待各个设备的并行计算操作(即

    X

    1

    A

    1

    X_1A_1

    X1A1

    X

    2

    A

    2

    X_2A_2

    X2A2)都完成后,同时将各个设备输出的中间结果正确地聚合(synchronize)。可通过 reduce + broadcast 操作在 GPU0 上计算完整

    X

    A

    XA

    XA,然后广播

    GeLU

    (

    X

    A

    )

    \\text{GeLU}(XA)

    GeLU(XA) 到其他 GPU。

  • 另一种方案是将权重矩阵

    A

    A

    A 沿列方向切分为

    A

    =

    [

    A

    1

    ,

    A

    2

    ]

    A = [A_1, A_2]

    A=[A1,A2]。这种切分方式使得每个 GPU 设备能独立完成部分 GEMM 运算,并应用 GeLU 激活函数。列切分权重方法的优势在于避免了前向传播中一次全局同步通信(行切分方案需先同步合并分块结果,再统一应用激活函数)。

    [

    Y

    1

    ,

    Y

    2

    ]

    =

    [

    GeLU

    (

    X

    A

    1

    )

    ,

    GeLU

    (

    X

    A

    2

    )

    ]

    [Y_1, Y_2] = [\\text{GeLU}(X A_1), \\text{GeLU}(X A_2)]

    [Y1,Y2]=[GeLU(XA1),GeLU(XA2)]

  • 在这里插入图片描述 将第一个线性层的权重按照列并行方式切分后,第二个线性层的权重矩阵 B 自然沿着行方向拆分,使其能够直接处理来自 GeLU 层的输出而无需任何通信,如图 3a 所示。

    2.2.2 拆分方式
  • MLP 模块中两个线性层权重先列后行的切分方法,其张量并行的前向传播过程拆解如下:
  • 第一个 GEMM(如

    X

    A

    1

    XA_1

    XA1

    X

    A

    2

    XA_2

    XA2):每个 GPU 独立计算,无需通信。

  • 第二个 GEMM(如

    Y

    1

    B

    1

    Y_1B_1

    Y1B1

    Y

    2

    B

    2

    Y_2B_2

    Y2B2)之后,需要一次 All-Reduce 操作合并结果,再将结果输入 Dropout 层 – 之所以需要 All-Reduce 归约操作,是因为第二个线性层的权重是按行切分得到

    Z

    1

    Z_1

    Z1

    Z

    2

    Z_2

    Z2,所以需要执行加法操作得到最终的 Z。 – 具体来说,第二个 GEMM(如

    Y

    1

    B

    1

    Y_1B_1

    Y1B1

    Y

    2

    B

    2

    Y_2B_2

    Y2B2) 之后,需要一次 All-Reduce 操作合并结果(对

    Y

    1

    B

    1

    Y_1B_1

    Y1B1

    Y

    2

    B

    2

    Y_2B_2

    Y2B2求和)。 – 激活函数(如 GeLU):本地计算,无需通信。 因为第二个线性层的权重是按行切分得到

    Z

    1

    Z_ 1

    Z1

    Z

    2

    Z_2

    Z2,所以需要执行加法操作得到最终的

    Z

    Z

    Z,即第二个 GEMM(如 $Y_1B_1$ 和 $Y_2B_2$之后,需要一次 All-Reduce 操作合并结果(对 $Y_1B_1$ 和 $Y_2B_2$求和)。以下的例子code/course12/simple_tp.py用于帮助大家理解这里的内容,其中X按列切分为

    X

    1

    X_1

    X1

    X

    2

    X_2

    X2,A作为权重按行切分为

    A

    1

    A_1

    A1

    A

    2

    A_2

    A2,两部分独立相乘之后得到

    Y

    1

    Y_1

    Y1

    Y

    2

    Y_2

    Y2,随后再将它们相加就能得到最终的结果。

  • import numpy as np

    def gelu(x):
    """GeLU 激活函数 (使用 tanh 近似)"""
    return 0.5 * x * (1 + np.tanh(np.sqrt(2 / np.pi) * (x + 0.044715 * np.power(x, 3))))

    # 1. 初始化参数
    np.random.seed(42)
    B = 4 # Batch size
    H = 8 # Hidden dimension (输入维度)
    D = 4 # Output dimension (输出维度)

    # 创建输入矩阵 X (B, H) 和 权重矩阵 A (H, D)
    X = np.random.randn(B, H)
    A = np.random.randn(H, D)

    # ==========================================
    # 方法 1: 单机基准计算 (不拆分)
    # ==========================================
    # 直接计算 Y = GeLU(X @ A)
    target_output = gelu(np.dot(X, A))
    print(f"基准输出形状: {target_output.shape}")

    # ==========================================
    # 方法 2: 模拟行拆分并行 (Row Parallelism)
    # ==========================================
    print("\\n— 开始模拟行拆分并行 —")

    # 假设我们有 2 个设备 (GPU),将 A 按行切分,将 X 按列切分
    # split_size = H // 2

    # [切分权重矩阵 A] -> A1, A2
    # A 的形状是 (H, D),按行切分 (axis=0)
    A1 = A[:H//2, :]
    A2 = A[H//2:, :]
    print(f"设备1 权重 A1 形状: {A1.shape}")
    print(f"设备2 权重 A2 形状: {A2.shape}")

    # [切分输入矩阵 X] -> X1, X2
    # X 的形状是 (B, H),为了配合 A 的行切分,X 需要按列切分 (axis=1)
    X1 = X[:, :H//2]
    X2 = X[:, H//2:]
    print(f"设备1 输入 X1 形状: {X1.shape}")
    print(f"设备2 输入 X2 形状: {X2.shape}")

    # [并行计算]
    # 每个设备独立计算自己的部分积: Xi * Ai
    # 结果形状均为 (B, D)
    Y1_partial = np.dot(X1, A1)
    Y2_partial = np.dot(X2, A2)

    # [同步点 / All-Reduce]
    # 在 GeLU 之前,必须将各设备的部分结果相加
    # 对应公式: XA = X1A1 + X2A2
    Y_combined = Y1_partial + Y2_partial

    # [应用非线性激活函数]
    # 聚合后才能执行 GeLU
    parallel_output = gelu(Y_combined)

    # ==========================================
    # 验证结果
    # ==========================================
    # 检查并行计算结果与基准结果是否一致
    is_close = np.allclose(target_output, parallel_output)
    print(f"\\n结果验证: {'成功' if is_close else '失败'}")
    print(f"两者误差 (Max Diff): {np.max(np.abs(target_output parallel_output))}")

    # 演示如果直接在局部做 GeLU 再相加是错误的 (数学原理验证)
    wrong_output = gelu(Y1_partial) + gelu(Y2_partial)
    print(f"\\n错误做法 (先 GeLU 后聚合) 误差: {np.max(np.abs(target_output wrong_output))}")
    print("结论: GeLU(X1A1 + X2A2) != GeLU(X1A1) + GeLU(X2A2)")

  • 另外,图 3a 中的 f 和 g 操作解释如下:
    • f 的前向推理(forward)计算:对应列并行线性层前后的 identity 操作。输入

      X

      X

      X 在各个 TP rank 上可用于本地计算,每块 GPU 使用自己持有的权重分片独立计算,例如得到

      X

      A

      1

      X A_1

      XA1

      X

      A

      2

      X A_2

      XA2,forward 阶段不需要额外通信。

    • g 的前向推理(forward)计算:对应行并行线性层后的归约操作。每块 GPU 完成本地 GEMM 后,分别得到局部输出

      Z

      1

      Z_1

      Z1

      Z

      2

      Z_2

      Z2;随后各 GPU 间执行一次 All-Reduce,对局部输出求和,得到最终的

      Z

      Z

      Z

    MLP 的张量并行过程的形状变换公式拆解如下:

    • b

      b

      b 表示 batch size;

    • s

      s

      s 表示 sequence length;

    • h

      h

      h 表示 hidden size;

    • i

      i

      i 表示 MLP intermediate size。 在 TP=2 时,每个 GPU 上的形状变化可以写成:

    • GPU0:

      [

      b

      ,

      s

      ,

      h

      ]

      ×

      [

      h

      ,

      2

      i

      /

      2

      ]

      [

      b

      ,

      s

      ,

      2

      i

      /

      2

      ]

      [b, s, h] \\times [h, 2i/2] \\rightarrow [b, s, 2i/2]

      [b,s,h]×[h,2i/2][b,s,2i/2] 经过 SiluAndMul 后:

      [

      b

      ,

      s

      ,

      2

      i

      /

      2

      ]

      [

      b

      ,

      s

      ,

      i

      /

      2

      ]

      [b, s, 2i/2] \\rightarrow [b, s, i/2]

      [b,s,2i/2][b,s,i/2] 再进入 down_proj:

      [

      b

      ,

      s

      ,

      i

      /

      2

      ]

      ×

      [

      i

      /

      2

      ,

      h

      ]

      Z

      1

      :

      [

      b

      ,

      s

      ,

      h

      ]

      [b, s, i/2] \\times [i/2, h] \\rightarrow Z_1: [b, s, h]

      [b,s,i/2]×[i/2,h]Z1:[b,s,h]

    • GPU1:

      [

      b

      ,

      s

      ,

      h

      ]

      ×

      [

      h

      ,

      2

      i

      /

      2

      ]

      [

      b

      ,

      s

      ,

      2

      i

      /

      2

      ]

      [b, s, h] \\times [h, 2i/2] \\rightarrow [b, s, 2i/2]

      [b,s,h]×[h,2i/2][b,s,2i/2] 经过 SiluAndMul 后:

      [

      b

      ,

      s

      ,

      2

      i

      /

      2

      ]

      [

      b

      ,

      s

      ,

      i

      /

      2

      ]

      [b, s, 2i/2] \\rightarrow [b, s, i/2]

      [b,s,2i/2][b,s,i/2] 再进入 down_proj:

      [

      b

      ,

      s

      ,

      i

      /

      2

      ]

      ×

      [

      i

      /

      2

      ,

      h

      ]

      Z

      2

      :

      [

      b

      ,

      s

      ,

      h

      ]

      [b, s, i/2] \\times [i/2, h] \\rightarrow Z_2: [b, s, h]

      [b,s,i/2]×[i/2,h]Z2:[b,s,h] 最后,RowParallelLinear 会对各 rank 的局部和做 All-Reduce 求和:

      Z

      =

      Z

      1

      +

      Z

      2

      Z = Z_1 + Z_2

      Z=Z1+Z2 形状为:

      [

      b

      ,

      s

      ,

      h

      ]

      +

      [

      b

      ,

      s

      ,

      h

      ]

      [

      b

      ,

      s

      ,

      h

      ]

      [b, s, h] + [b, s, h] \\rightarrow [b, s, h]

      [b,s,h]+[b,s,h][b,s,h]

  • 下述代码是

    f

    f

    f运算符的实现示例:

  • """
    f operator 实现:
    – 前向传递:直接返回输入(恒等运算)
    – 反向传递:对梯度执行 All-Reduce
    对应的 g operator 行为对称:
    – 前向传递:执行 All-Reduce
    – 反向传递:直接返回梯度(恒等运算)

    Implementation of f operator. g is similar to f with
    identity in the backward and all-reduce in the forward functions.
    """class f(torch.autograd.Function):
    @staticmethoddef forward(ctx, x):
    return x # 前向传递无通信 @staticmethoddef backward(ctx, grad_output):
    all_reduce(grad_output) # 反向传递触发 All-Reducereturn grad_output

    class g(torch.autograd.Function):
    @staticmethoddef forward(ctx, x):
    all_reduce(x) # 前向传递触发 All-Reduce @staticmethoddef backward(ctx, grad_output):
    return grad_output # 反向传递无通信

    总结:当列并行与行并行级联使用时,在常见的 ColumnParallelLinear(gather_output=False) 接 RowParallelLinear(input_is_parallel=True) 的组合下,前级输出本来就是按最后一维切分后的局部结果,正好可以作为后级行并行所需的局部输入。 需要注意的是,先列后行的并行级联方式只是表示中间无需通信,并不意味着整个模块没有集合通信。RowParallelLinear 默认 reduce_results=True,通常会在自身输出侧执行 AllReduce,将各 GPU 的部分结果求和成完整输出。如果它的前面或后面还要与其他并行方式连接,也可能根据张量布局继续使用 AllGather、AllReduce、ReduceScatter 等集合通信操作。 先列后行的并行级联的 MLP 前向传播的可视化连接如下图所示: 在这里插入图片描述 vLLM 框架中 MLP 模块针对线性层的张量并行,它的实际代码实现也是先列后行线性层: 在这里插入图片描述

    2.2.3 MLA 层的通讯量分析(重)

    总结:MLP 层在 forward(前向推理时) 做一次 All-Reduce 操作,在 backward(前向推理时) 做一次 All-Reduce 操作。而 All-Reduce 的过程分为两个阶段,Reduce-Scatter 和 All-Gather,每个阶段的通讯量是相等的。假设输入张量大小为 [b, s, h],数据类型为 fp16,则每次 All-Reduce 操作通讯量为 2bsh。 模型训练和推理阶段,MLP 层的张量并行通信量如下所示:

    • 模型训练时,包含前向传播和反相传播两个过程,即两次 All-Reduce 操作,所以 MLP 层的总通讯量为:

      4

      b

      s

      h

      4bsh

      4bsh

    • 模型推理时,只有前向传播过程,即一次 All-Reduce 操作,所以 MLP 层的总通讯量为:

      2

      b

      s

      h

      2bsh

      2bsh。 这是因为前文中说到,随着 N 的增大,Ring AllReduce 通信算子的通信量可以近似为 2K,其中 K 表示传输数据的数据量大小。

    2.3 MHA 层的张量并行

    2.3.1 拆分原理

    多头注意力模块的结构如下图所示,可以看出,在设计上,MHA 层对于每个头(head),就有都有独立的 q/k/v 三个线性变换层以及对应的 self-attention 计算结构,然后将每个 head 输出的结果做拼接 concat,最后将拼接得到结果做线性变换得到最终的注意力层输出张量。

    在这里插入图片描述下图展示了当 num_attention_heads = 2 时 attention 层的 Q/K/V 线性变换的并行计算方法。对每一块权重,我们都沿着列方向(k_dim)维度切割一刀。此时每个 head 上的

    W

    Q

    W^Q

    WQ

    W

    K

    W^K

    WK

    W

    V

    W^V

    WV的维度都变成 (d_model, k_dim//2)。每个 head 上能独立做矩阵计算,最后将计算结果 concat起来即可。整个流程如下图所示:

    在这里插入图片描述

    可以发现,从多头注意力结构看,其计算机制真的是天然适合模型的张量并行计算,即每个头上都可以在每个设备上独立计算,即可以把每个头(也可以是 n 个头)的参数放到一块 GPU 上,最后将子结果 concat 后得到最终的张量。具体来说,多头注意力结构的张量并行计算过程拆解如下,实际模型中,会存在一个或多个 head 占用一块 GPU 的情况,且我们尽量保证 heads 总数能被 GPU 个数整除。

  • W

    Q

    W^Q

    WQ

    W

    K

    W^K

    WK

    W

    V

    W^V

    WV 做列并行切分 在多头注意力中,Q/K/V 投影可以按输出维,也就是 head 维进行切分。每个 TP rank 负责一部分 attention heads,并在本地完成这些 heads 对应的 Q/K/V 线性变换。因此,每个 rank 可以独立完成本地 heads 的 Q/K/V GEMM 和后续 attention 计算。在 Q/K/V 投影输出到本地 attention 计算之间,通常不需要立即进行跨 rank 集合通信。

  • 对输出线性层 o_proj 做行并行切分 各 rank 完成本地 heads 的 attention 计算后,会得到本地 attention 输出。接下来,这些输出会进入输出线性层,即 o_proj。在 vLLM 中,o_proj 通常使用 RowParallelLinear,即按输入维做行并行切分。每个 rank 先基于自己的本地 heads 输出计算一个局部和,然后在 forward 末尾通过 All-Reduce 将各 rank 的局部和相加,恢复完整的 hidden states。 因此,Attention 的 TP 结构和 MLP 中先列并行、再行并行的模式非常相似:Q/K/V projection 类似列并行,中间的本地 head attention 不需要立即汇总完整张量,最后的 o_proj 类似行并行,并通过 All-Reduce 合并输出。 Attention 模块的张量并行加速过程如图 3(b) 所示: 在这里插入图片描述
  • vLLM 框架中 Attention 模块针对线性层的张量并行,实际代码实现也是先列后行线性层:

    在这里插入图片描述

    2.3.2 MHA 层的通讯量分析

    很明显上述的设计对 MLP 和自注意力层均采用了将两组 GEMM 运算融合的策略,从而消除了一个同步步骤,并获得了更好的扩展性。基于此技术方案,在一个标准 transformer 层中,前向传播只需执行两次 all-reduce 操作,反向传播也仅需两次 all-reduce(详见图 4)。 在这里插入图片描述 和 MLP 模块类似,模型训练和推理阶段,MHA 层的张量并行通信量如下所示:

    • 模型训练时,包含前向传播和反相传播两个过程,即两次 All-Reduce 操作,所以 MLP 层的总通讯量为:4bsh。
    • 模型推理时,只有前向传播过程,即一次 All-Reduce 操作,所以 MLP 层的总通讯量为:2bsh。

    简单理解,这里的两次All-Reduce分别来自于self-attention层和MLP层

    2.3.3 attention 中的张量并行实例

    在多头注意力中,Q、K、V 的输出维度可以看成由多个 attention heads 拼接而成。因此,当我们对 wq、wk、wv 的第二维进行切分时,本质上就是把不同的 heads 分配到不同的 tensor parallel rank 上计算。

    在这个例子中,tp_size = 2,所以我们将 wq、wk、wv 按列切成两份。第一张卡(tp_rank = 0)负责前一半 heads;第二张卡(tp_rank = 1)负责后一半 heads。与 Q/K/V 投影不同,输出投影 wo 采用行切分。因为 attention 输出已经按 head 维度分布在不同的 rank 上,所以 wo 需要按输入维度(即第一维)进行切分:

    wq_sub1 = wq_weight[:, :hidden_dim // 2]
    wk_sub1 = wk_weight[:, :hidden_dim // 2]
    wv_sub1 = wv_weight[:, :hidden_dim // 2]
    wo_sub1 = wo_weight[:hidden_dim // 2, :]

    随后,将输入复制到每个 tensor parallel rank 上,并分别与本地切分后的 Q/K/V 权重进行矩阵乘法。下面是第一张卡上的计算结果,分别记作 q1、k1、v1:

    q1 = np.matmul(inputs, wq_sub1) # [bsz, seq_len, hidden_dim//2]
    k1 = np.matmul(inputs, wk_sub1)
    v1 = np.matmul(inputs, wv_sub1)

    第二张卡的计算方式完全相同,只是使用的是 Q/K/V 权重的后一半切片。这样,每张卡都会得到自己负责的那部分 heads 对应的 q、k、v。接下来,每张卡使用本地的 q、k、v 独立计算 self-attention。由于标准多头注意力中不同 heads 之间是相互独立的,因此每个 rank 可以只计算自己负责的 heads,而不需要在 attention 计算阶段与其他 rank 通信。 得到局部 attention 输出之后,每张卡再将自己的 attention 输出与 wo 的对应行分块相乘。由于完整的输出投影可以拆成多个局部矩阵乘法结果之和,因此最后将各个 rank 的局部结果相加,就可以得到完整的 attention 输出:

    output_tp_parallel = attn_output1 @ wo_sub1 + attn_output2 @ wo_sub2

    就是说,整个 self-attention 的张量并行过程可以概括为: Q/K/V 投影:列并行

    • wq、wk、wv 按输出维度切分
    • 等价于按 attention heads 切分
    • 每个 rank 计算一部分 heads 的 q、k、v

    Attention 计算:

    • 每个 rank 独立计算本地 heads 的 attention
    • 标准 MHA 中不同 heads 之间不需要通信

    输出投影 wo:行并行

    • wo 按输入维度切分
    • 每个 rank 计算局部输出
    • 最后通过求和(即实际分布式实现中的 all-reduce)得到完整结果 因此,这个例子体现的正是 self-attention 中常见的“先列并行、后行并行”的张量并行模式:Q/K/V 投影阶段按列切分权重(即切分 heads);输出投影阶段按行切分 wo,再对各个 rank 的局部输出进行求和。 以下是完整的代码,详见 code/course12/attention_tp.py,过程基本符合Megatron-LM 的张量并行算法中所述。

    import numpy as np

    if __name__ == '__main__':
    # 参数设置
    bsz = 4
    seq_len = 16
    hidden_dim = 128
    num_heads = 8
    head_dim = hidden_dim // num_heads

    np.random.seed(42) # 设置随机种子,保证结果可复现

    # 随机初始化权重矩阵
    wq_weight = np.random.randn(hidden_dim, hidden_dim)
    wk_weight = np.random.randn(hidden_dim, hidden_dim)
    wv_weight = np.random.randn(hidden_dim, hidden_dim)
    wo_weight = np.random.randn(hidden_dim, hidden_dim)

    inputs = np.random.randn(bsz, seq_len, hidden_dim)

    print("===== 标准版本的多头注意力计算 =====")
    # 1. 标准版本 – 不分割计算
    q = np.matmul(inputs, wq_weight) # [bsz, seq_len, hidden_dim]
    k = np.matmul(inputs, wk_weight) # [bsz, seq_len, hidden_dim]
    v = np.matmul(inputs, wv_weight) # [bsz, seq_len, hidden_dim]

    # 重塑为多头形式
    q = q.reshape(bsz, seq_len, num_heads, head_dim) # [bsz, seq_len, num_heads, head_dim]
    k = k.reshape(bsz, seq_len, num_heads, head_dim) # [bsz, seq_len, num_heads, head_dim]
    v = v.reshape(bsz, seq_len, num_heads, head_dim) # [bsz, seq_len, num_heads, head_dim]

    # 调整维度顺序
    q = np.transpose(q, (0, 2, 1, 3)) # [bsz, num_heads, seq_len, head_dim]
    k = np.transpose(k, (0, 2, 1, 3)) # [bsz, num_heads, seq_len, head_dim]
    v = np.transpose(v, (0, 2, 1, 3)) # [bsz, num_heads, seq_len, head_dim]

    # 注意力分数计算
    scores = np.matmul(q, np.transpose(k, (0, 1, 3, 2))) / np.sqrt(head_dim) # [bsz, num_heads, seq_len, seq_len]

    # 应用softmax
    attn_probs = np.exp(scores np.max(scores, axis=1, keepdims=True))
    attn_probs = attn_probs / np.sum(attn_probs, axis=1, keepdims=True) # [bsz, num_heads, seq_len, seq_len]

    # 注意力输出计算
    attn_output = np.matmul(attn_probs, v) # [bsz, num_heads, seq_len, head_dim]

    # 恢复原始维度
    attn_output = np.transpose(attn_output, (0, 2, 1, 3)) # [bsz, seq_len, num_heads, head_dim]
    attn_output = attn_output.reshape(bsz, seq_len, hidden_dim) # [bsz, seq_len, hidden_dim]

    # 最终输出投影
    output = np.matmul(attn_output, wo_weight) # [bsz, seq_len, hidden_dim]

    print("===== 张量并行版本的多头注意力计算 =====")
    # 2. 张量并行版本 – 按头切分
    # 每个并行组处理一半的头
    heads_per_gpu = num_heads // 2

    # GPU 1处理前half_heads个头
    # 按列切分QKV权重 – 每个GPU负责一半头的权重
    wq_sub1 = wq_weight[:, :hidden_dim // 2] # 前half_heads个头的权重
    wk_sub1 = wk_weight[:, :hidden_dim // 2]
    wv_sub1 = wv_weight[:, :hidden_dim // 2]
    wo_sub1 = wo_weight[:hidden_dim // 2, :]
    # GPU 1上的计算
    q1 = np.matmul(inputs, wq_sub1) # [bsz, seq_len, hidden_dim//2]
    k1 = np.matmul(inputs, wk_sub1)
    v1 = np.matmul(inputs, wv_sub1)

    # 重塑为多头形式
    q1 = q1.reshape(bsz, seq_len, heads_per_gpu, head_dim)
    k1 = k1.reshape(bsz, seq_len, heads_per_gpu, head_dim)
    v1 = v1.reshape(bsz, seq_len, heads_per_gpu, head_dim)

    # 调整维度顺序
    q1 = np.transpose(q1, (0, 2, 1, 3)) # [bsz, heads_per_gpu, seq_len, head_dim]
    k1 = np.transpose(k1, (0, 2, 1, 3))
    v1 = np.transpose(v1, (0, 2, 1, 3))

    # 计算注意力分数
    scores1 = np.matmul(q1, np.transpose(k1, (0, 1, 3, 2))) / np.sqrt(head_dim)

    # 应用softmax
    attn_probs1 = np.exp(scores1 np.max(scores1, axis=1, keepdims=True))
    attn_probs1 = attn_probs1 / np.sum(attn_probs1, axis=1, keepdims=True)

    # 注意力输出计算
    attn_output1 = np.matmul(attn_probs1, v1) # [bsz, heads_per_gpu, seq_len, head_dim]

    # 恢复原始维度
    attn_output1 = np.transpose(attn_output1, (0, 2, 1, 3)) # [bsz, seq_len, heads_per_gpu, head_dim]
    attn_output1 = attn_output1.reshape(bsz, seq_len, hidden_dim // 2) # [bsz, seq_len, hidden_dim//2]

    # GPU 2处理后half_heads个头
    wq_sub2 = wq_weight[:, hidden_dim // 2:] # 后half_heads个头的权重
    wk_sub2 = wk_weight[:, hidden_dim // 2:]
    wv_sub2 = wv_weight[:, hidden_dim // 2:]
    wo_sub2 = wo_weight[hidden_dim // 2:, :]

    # GPU 2上的计算
    q2 = np.matmul(inputs, wq_sub2)
    k2 = np.matmul(inputs, wk_sub2)
    v2 = np.matmul(inputs, wv_sub2)

    # 重塑为多头形式
    q2 = q2.reshape(bsz, seq_len, heads_per_gpu, head_dim)
    k2 = k2.reshape(bsz, seq_len, heads_per_gpu, head_dim)
    v2 = v2.reshape(bsz, seq_len, heads_per_gpu, head_dim)

    # 调整维度顺序
    q2 = np.transpose(q2, (0, 2, 1, 3)) # [bsz, heads_per_gpu, seq_len, head_dim]
    k2 = np.transpose(k2, (0, 2, 1, 3))
    v2 = np.transpose(v2, (0, 2, 1, 3))

    # 计算注意力分数
    scores2 = np.matmul(q2, np.transpose(k2, (0, 1, 3, 2))) / np.sqrt(head_dim)

    # 应用softmax
    attn_probs2 = np.exp(scores2 np.max(scores2, axis=1, keepdims=True))
    attn_probs2 = attn_probs2 / np.sum(attn_probs2, axis=1, keepdims=True)

    # 注意力输出计算
    attn_output2 = np.matmul(attn_probs2, v2) # [bsz, heads_per_gpu, seq_len, head_dim]

    # 恢复原始维度
    attn_output2 = np.transpose(attn_output2, (0, 2, 1, 3)) # [bsz, seq_len, heads_per_gpu, head_dim]
    attn_output2 = attn_output2.reshape(bsz, seq_len, hidden_dim // 2) # [bsz, seq_len, hidden_dim//2]

    # 合并结果 (相当于在head维度上concatenate)
    output_tp_parallel = attn_output1 @ wo_sub1 + attn_output2 @ wo_sub2

    print(np.mean(np.abs(output output_tp_parallel)))

    2.4 Embedding 层的张量并行

    Transformer 语言模型的输出侧通常会通过 LM head 将 hidden states 投影到词表维度。这里需要区分两个概念:LM head 的权重矩阵通常是 [vocab_size, hidden_size],而计算得到的 logits 通常是 [batch_size 或 token 数, vocab_size]。由于当前 Transformer 语言模型的词汇表通常至少有数万个 token,因此对 LM head 这类大词表投影进行张量并行,通常可以减少单张 GPU 上的计算量和显存占用。 在 vLLM 中,输入 embedding 和输出 LM head 都采用按词表维度切分的方式。需要注意的是,输入 embedding 的通信方式和输出 LM head 不一样:

  • 输入 embedding 在各 rank 得到局部结果后,通过 all-reduce 把结果相加;
  • 输出 LM head 则是每个 rank 先计算自己负责的词表分片 logits,然后在 logits 处理阶段执行 gather 或 all-gather,得到完整词表维度的 logits。
  • 2.4.1 输入嵌入层

    Embedding 层开启 TP 时,会将输入嵌入层的权重矩阵

    E

    E

    E(尺寸为 [vocab_size, hidden_size])按词汇维度拆分。由于词汇维度对应矩阵的第 0 维,所以可以理解为“按行拆分”:每个 GPU 只保存一段 token id 范围对应的 embedding 行。例如 TP=2 时,可以把词表分成两个连续范围,rank 0 负责前半段 token,rank 1 负责后半段 token。 因为每个分块只包含嵌入表的一部分,所以每张卡只能直接查到属于自己词表范围内的 token。对于不属于当前 rank 的 token,vLLM 会先通过 mask 将其排除,并把对应位置的 embedding 输出置为 0。随后,各个 rank 对局部 embedding 结果执行一次 all-reduce。由于同一个 token 只会在负责它的 rank 上产生非零 embedding,all-reduce 求和后就能得到完整的输入嵌入结果。 简单理解这个过程:假设现在有两个 GPU,每个 GPU 负责一半词表。输入请求的 seq=2,input_ids=[31, 33]。假设 token 31 属于 GPU 0 的词表范围,token 33 属于 GPU 1 的词表范围。那么在 GPU 0 上,token 31 能查到真实 embedding,token 33 不属于本卡词表范围,会被置为 0:

    GPU 0:
    [
    [0.11, 0.22, 0.33, …, 0.12],
    [0.00, 0.00, 0.00, …, 0.00]
    ]

    在 GPU 1 上则相反,token 31 会被置为 0,token 33 能查到真实 embedding:

    GPU 1:
    [
    [0.00, 0.00, 0.00, …, 0.00],
    [0.31, 0.36, 0.13, …, 0.62]
    ]

    最后对两个 GPU 的结果做 all-reduce 求和,就得到完整的输入嵌入:

    [
    [0.11, 0.22, 0.33, …, 0.12],
    [0.31, 0.36, 0.13, …, 0.62]
    ]

    也就是说,输入 embedding 的 TP 流程可以概括为:先判断每个 token 是否属于当前 rank 的词表范围;属于本 rank 的 token 正常查表;不属于本 rank 的 token 先 mask 掉并把输出置 0;最后通过 all-reduce 把各个 rank 的局部结果相加,恢复完整的输入 embedding。

    import numpy as np

    # ==========================================
    # 0. 初始化
    # ==========================================
    np.random.seed(42)
    Vocab = 10
    Hidden = 4
    B, S = 2, 2

    # 完整的 Embedding 表 (模拟 Ground Truth)
    # 形状: (10, 4)
    E_full = np.random.randn(Vocab, Hidden)

    # 输入 Token IDs (Batch=2, Seq=2)
    # 包含落在两个 GPU 范围内的 ID
    # 1, 3 -> GPU 1 (0-4)
    # 6, 8 -> GPU 2 (5-9)
    Input_IDs = np.array([
    [1, 6],
    [3, 8]
    ])

    print("输入 IDs:\\n", Input_IDs)

    # ==========================================
    # 1. 单机基准 (Lookup)
    # ==========================================
    # Numpy 的高级索引模拟 Embedding Lookup
    Output_Ref = E_full[Input_IDs]
    print(f"\\n基准输出形状: {Output_Ref.shape}") # (2, 2, 4)

    # ==========================================
    # 2. 并行模拟 (Parallel Embedding)
    # ==========================================
    print("\\n— 开始并行模拟 —")

    # [切分权重] 按词汇维度切分 (Row Parallel in terms of Matrix,
    # 但通常称为 Vocab Parallel)
    # Split Size = 5
    V_per_gpu = Vocab // 2

    # GPU 1: 负责 ID 0-4
    E_gpu1 = E_full[:V_per_gpu, :] # (5, 4)
    range_start_1 = 0
    range_end_1 = V_per_gpu

    # GPU 2: 负责 ID 5-9
    E_gpu2 = E_full[V_per_gpu:, :] # (5, 4)
    range_start_2 = V_per_gpu
    range_end_2 = Vocab

    def parallel_embedding_forward(input_ids, local_weight, start_idx, end_idx):
    """
    每个 GPU 独立执行的 Forward 函数
    """
    # 1. 创建掩码: 找出哪些 ID 属于当前 GPU
    mask = (input_ids >= start_idx) & (input_ids < end_idx)

    # 2. 将全局 ID 映射为本地 ID (Offset)
    # 例如 GPU 2 负责 5-9,ID=6 对应的本地索引是 1
    local_ids = input_ids – start_idx

    # 3. 为了避免索引越界,将不属于自己的 ID 置为 0
    safe_ids = np.where(mask, local_ids, 0)

    # 4. 查表 (Lookup)
    local_output = local_weight[safe_ids]

    mask_expanded = mask[:, :, np.newaxis]

    # 5. 只有属于自己的 ID 保留 Lookup 结果,其他的变成 0.0
    final_local_output = local_output * mask_expanded

    return final_local_output

    # — GPU 1 计算 —
    Out_gpu1 = parallel_embedding_forward(Input_IDs, E_gpu1, range_start_1, range_end_1)
    print("\\nGPU 1 输出 (部分为0):")
    print(Out_gpu1[0]) # 看第一行: [Vector(ID=1), Vector(0.0)]

    # — GPU 2 计算 —
    Out_gpu2 = parallel_embedding_forward(Input_IDs, E_gpu2, range_start_2, range_end_2)
    print("\\nGPU 2 输出 (部分为0):")
    print(Out_gpu2[0]) # 看第一行: [Vector(0.0), Vector(ID=6)]

    # ==========================================
    # 3. 同步聚合 (All-Reduce)
    # ==========================================
    Output_Fused = Out_gpu1 + Out_gpu2

    print(f"\\n验证结果: {np.allclose(Output_Ref, Output_Fused)}")
    print("逻辑: Embedding(ID) = Embedding_GPU1(ID) + Embedding_GPU2(ID)")
    print(" 其中一个必然是 0向量,另一个是真实向量")

    2.4.2 输出嵌入层

    对于输出 LM head,先考虑一种朴素的完整 logits 汇总方案。假设 hidden states 为

    X

    X

    X,输出权重矩阵按词表维度切分为

    E

    1

    ,

    E

    2

    E_1, E_2

    E1,E2。若

    E

    i

    E_i

    Ei 的形状沿用 embedding 权重的表示,即 [vocab_size/tp_size, hidden_size],那么每个 rank 上的本地 logits 应写成:

    Y

    i

    =

    X

    E

    i

    T

    Y_i = X E_i^T

    Yi=XEiT 两个 rank 并行计算后,可以得到词表分片上的 logits:

    [

    Y

    1

    ,

    Y

    2

    ]

    =

    [

    X

    E

    1

    T

    ,

    X

    E

    2

    T

    ]

    [Y_1, Y_2] = [X E_1^T, X E_2^T]

    [Y1,Y2]=[XE1T,XE2T],接下来再通过 gather 或 all-gather 将这些本地 logits 按词表维度拼接起来,得到完整 logits:

    Y

    =

    gather/all-gather

    (

    [

    Y

    1

    ,

    Y

    2

    ]

    )

    Y = \\text{gather/all-gather}([Y_1, Y_2])

    Y=gather/all-gather([Y1,Y2])

    三 vLLM 中的张量并行

    3.1 vLLM 中张量并行如何使用

    使用 vLLM 启动模型服务时,如果模型太大,无法放入单个 GPU,但可以放入单个节点中的多个 GPU,就可以使用张量并行。只考虑 TP、没有叠加 PP、DP、DCP 等其他并行维度时,tensor_parallel_size 通常就等于你希望用于单个模型副本的 GPU 数量。例如,单节点有 4 个 GPU,可以将张量并行大小设置为 4。 对于多 GPU 的离线推理,可以在 LLM 类中设置 tensor_parallel_size 为所需的 GPU 数量。例如,要在 4 个 GPU 上运行推理::

    from vllm import LLM
    llm = LLM("Qwen/Qwen3-32B", tensor_parallel_size=4) # 或者改为模型权重的本地路径
    output = llm.generate("San Francisco is a")

    对于多 GPU 服务,也就是在线推理,可以在启动服务器时包含 –tensor-parallel-size。例如,在 4 个 GPU 上运行 API server:

    # 前提是安装成功了 vllm,在可通过下述命令启动多 GPU 服务
    vllm serve Qwen/Qwen3-32B \\
    –tensor-parallel-size 4

    在 vLLM 的多进程执行路径中,前端进程通过 ZMQ socket 与后台 EngineCore 通信。开启 DP 时,vLLM 会按 DP 模式管理多个 EngineCore;MoE 的 DP 场景会使用 DPEngineCoreProc,非 MoE 的 DP rank 则更接近多个相互独立的 EngineCore。 如果同时开启 DP 和 TP,可以把每个 DP rank 理解为一个模型副本,副本内部再按 tensor_parallel_size 切分到多个 GPU。只开启 DP+TP 时,1 个 DP 副本内部的 GPU 数量等于 TP 数;如果还叠加 PP(pipeline parallel)或 prefill context parallel,则还要继续乘上这些并行维度。

    3.2 vLLM 中的分布式资源管理

    vLLM 中与分布式资源管理相关的核心逻辑主要位于 vllm/distributed 目录。其中,distributed/parallel_state.py 负责分布式并行状态管理。它封装了 PyTorch ProcessGroup 等底层通信机制,为推理过程提供统一的并行组管理和通信接口。 工作流如下:

  • 调用 init_distributed_environment 初始化分布式环境,例如 NCCL 或 Gloo。
  • 调用 initialize_model_parallel 或 ensure_model_parallel_initialized 初始化模型并行组,例如 TP、PP、DP、EP 等。
  • 调用 destroy_model_parallel 和 destroy_distributed_environment 清理模型并行组和分布式环境。
  • 在 initialize_model_parallel() 中,vLLM 会根据当前 world_size、rank,以及 TP、PP、DP、prefill/decode context parallel、EP 等配置,构造全局 rank 布局,布局顺序是:ExternalDP x DP x PP x TP 随后,initialize_model_parallel() 通过 reshape、transpose 等操作生成不同并行维度的 group_ranks。例如,TP 组会把同一个 PP stage 内的 rank 放到一起;PP 组会把同一个 TP rank、跨不同 PP stage 的 rank 放到一起。 init_model_parallel_group() 更偏底层:它接收已经计算好的 group_ranks,并创建对应的 GroupCoordinator。GroupCoordinator 是对 PyTorch ProcessGroup 的封装,会为这些 rank 创建 device 通信组,并记录当前进程在组内的 rank_in_group、world_size 等信息。后续的 all_reduce、all_gather、send/recv 等操作,都会通过对应的 GroupCoordinator 在正确的 rank 范围内执行。 vLLM 使用模块级全局变量保存并行组状态。例如,_TP 记录张量并行组,_PP 记录流水线并行组,_DP 记录数据并行组;当前版本还包含 _DCP、_PCP、_EP、_EPLB 等组状态。其中,_EP 主要用于 MoE 场景,dense model 通常不会创建 EP group。后续代码可以通过 get_tp_group()、get_pp_group()、get_dp_group() 等函数获取对应通信组。 例如,假设有 8 张 GPU,并设置:

    • tensor_parallel_size = 2
    • pipeline_parallel_size = 4
    • data_parallel_size = 1

    此时world_size = 8,avLLM 会先将 global rank 组织成一个二维结构:

    PP stage 0: GPU0 GPU1
    PP stage 1: GPU2 GPU3
    PP stage 2: GPU4 GPU5
    PP stage 3: GPU6 GPU7
    TP0 TP1

    然后生成 TP 和 PP 两类通信组: TP groups:

    • [GPU0, GPU1]
    • [GPU2, GPU3]
    • [GPU4, GPU5]
    • [GPU6, GPU7]

    PP groups:

    • [GPU0, GPU2, GPU4, GPU6]
    • [GPU1, GPU3, GPU5, GPU7]

    其中,TP 组负责同一个 PP stage 内的张量并行通信。例如,[GPU2, GPU3] 共同负责 stage 1 的模型层,RowParallelLinear 的 AllReduce 只会在 GPU2 和 GPU3 之间发生。 PP 组负责同一个 TP rank 跨不同 PP stage 的通信。例如,[GPU1, GPU3, GPU5, GPU7] 表示 TP rank 1 这一路的流水线通道。前向执行时,stage 0 的 GPU1 会将对应激活发送给 stage 1 的 GPU3,随后再沿 GPU3 -> GPU5 -> GPU7 向后传递

    _TP = init_model_parallel_group(group_ranks,
    get_world_group().local_rank,
    backend,
    use_message_queue_broadcaster=True,
    group_name="tp")

    四种并行组作用如下:

    • TP(张量并行):每组负责模型参数的分区计算,通常用于分割权重矩阵。
    • PP(流水线并行):每组负责模型的不同层或阶段,适合深层模型分段计算。
    • DP(数据并行):每组处理不同的数据批次,实现批量训练或推理;每个 DP group 的成员在模型上参数完全相同,只负责数据并行。
    • EP(专家并行):用于 MoE(Mixture of Experts)等结构,每组分配部分专家网络。专家并行 = 数据并行 × 张量并行的组合,每个 group 内 rank 会共享一组专家层的负载。

    实例理解,TP + PP 并行组的概念。在 8 卡 GPU 上配置 流水线并行度(PP)为 4 和 张量并行度(TP)为 2 时,并行分组逻辑如下:

    • 全局并行设置:
      • world_size(总 GPU 数)= 8
      • pipeline_model_parallel_size(流水线并行组数)= 4
      • tensor_model_parallel_size(张量并行组大小)= 2
    • 并行分组详解:
    • 张量并行(TP)分组:模型参数在组内的 2 张 GPU 间进行切分。系统会形成 4 个 TP 组,每个组处理模型的一部分。
    • 流水线并行(PP)分组:模型层被分配到 4 个连续的流水线阶段。系统会形成 2 个 PP 组,每个组内的 4 张 GPU 分别负责一个阶段,共同完成一个完整的批次处理。

    # 共 4 个张量并行(TP)组,每组在 2 张卡上切分模型参数
    TP Group 0: [GPU0, GPU1]
    TP Group 1: [GPU2, GPU3]
    TP Group 2: [GPU4, GPU5]
    TP Group 3: [GPU6, GPU7] # 共 2 个流水线并行(PP)组,每组由 4 张卡构成一个完整流水线
    PP Group 0: [GPU0, GPU2, GPU4, GPU6] # 处理微批次 A
    PP Group 1: [GPU1, GPU3, GPU5, GPU7] # 处理微批次 B

    在这里插入图片描述 结合图表与分组信息,可以这样理解 TP + PP 并行组内的交互:

  • 组内协作(纵向):以一个 PP 组中的多个 pipeline stage 为例,比如 [GPU0, GPU2, GPU4, GPU6]。这些 GPU 对应同一 TP 位置上的不同 stage,一个微批次会按顺序经过这些 stage,数据从 GPU0 流到 GPU6,逐段完成对应层的前向或反向计算。
  • 组间并行(横向):同一个 TP 组内的 GPU,比如 GPU0 和 GPU1,分别位于不同的 PP 组中。它们持有同一层模型参数的不同切片,并在对应层计算后通过 TP 组内的集合通信同步局部结果,具体可能是 all-gather、all-reduce 或 reduce-scatter,取决于该层的切分方式。
  • 模拟分布式并行分分组算法

    可以通过下述示例代码(CPU模拟,不需要多机)来模拟分布式并行分分组算法,得到不同 gpu 设备和不同推理配置下的并行分组信息。rank网格简化为 [ExternalDP, DP, PP, TP]。

    import torch

    def compute_groups(world_size, tp, pp, dp):
    assert world_size % (tp * pp * dp) == 0
    E = world_size // (tp * pp * dp) # ExternalDP 大小(通常为 1)
    all_ranks = torch.arange(world_size).reshape(E, dp, pp, tp)

    # TP 组
    tp_groups = [x.tolist() for x in all_ranks.view(1, tp).unbind(0)]

    # PP 组
    pp_groups = [x.tolist() for x in all_ranks.transpose(2, 3).reshape(1, pp).unbind(0)]

    # DP 组(模型内部 DP)
    dp_groups = [x.tolist() for x in all_ranks.transpose(1, 3).reshape(1, dp).unbind(0)]

    # EP 组(同 PP stage 下合并 DP×TP)
    ep_groups = [x.tolist() for x in all_ranks.transpose(1, 2).reshape(1, dp * tp).unbind(0)]

    return tp_groups, pp_groups, dp_groups, ep_groups

    print("TP groups:", tp_groups)
    print("PP groups:", pp_groups)
    print("DP groups:", dp_groups)
    print("EP groups:", ep_groups)

    if __name__ == "__main__":
    """
    world_size = 16
    tp_size = 4
    pp_size = 1
    dp_size = 4
    """

    compute_groups(16, 4, 1, 4) # 卡数 16

    输出结果如下所示:

    TP groups: [[0, 1, 2, 3], [4, 5, 6, 7], [8, 9, 10, 11], [12, 13, 14, 15]]
    PP groups: [[0], [1], [2], [3], [4], [5], [6], [7], [8], [9], [10], [11], [12], [13], [14], [15]]
    DP groups: [[0, 4, 8, 12], [1, 5, 9, 13], [2, 6, 10, 14], [3, 7, 11, 15]]
    EP groups: [[0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15]]

    这里把这些 rank 按照不同类型拆分到 4 维上。这里的 ext_dp_size 指的是本系统之外的一层数据并行维度,在本节课中不做深究。

    • DP 维度对应传统意义上的数据并行组(组内 rank 拥有完整模型、处理不同子 batch)
    • PP 维度对应流水线并行的层切分(每个 stage 只持有一段层)
    • TP 维度对应张量并行的参数切分(同一层内部的大矩阵按维度拆到多个 rank 上计算,并 All-Reduce 合并) 在此基础上,引入 EP 时,通常是把同一个 PP stage 下若干个 (DP, TP) 的笛卡尔积子网格看作一个 EP group。也就是说,同一个 PP stage 内的不同 DP、TP rank 共同承担同一批 token 的 MoE 路由与专家计算,同时各自只持有部分专家,即让一个 EP group 里的所有 GPU 一起负责一整批 token 的 MoE 计算。 以上 rank 网格中我们忽略了 PCP(prefill_context_parallel_size)。这里的 prefill_context_parallel_size 指的是:在 prefill 阶段,将长上下文按 token 维度切分到多个 rank 上并行处理,因此实际的 rank 网格里除了 DP、PP、TP 之外,还会额外增加一个 PCP 维度。

    3.3 vLLM 的并行线性层

    vllm/model_executor/models/qwen3.py 代码中的 Qwen3Attention 模块的核心网络层组成如下所示:

    Qwen3Attention (继承自 nn.Module)
    ├── qkv_proj: QKVParallelLinear (用于合并计算 Q, K, V)
    │ └── (继承自) ColumnParallelLinear (列并行层)
    │ └── (继承自) LinearBase (并行层抽象基类)

    ├── o_proj: RowParallelLinear (用于 Attention 输出)
    │ └── (继承自) LinearBase (并行层抽象基类)

    ├── rotary_emb: RotaryEmbedding (旋转位置编码)

    ├── attn: Attention (底层 Attention 计算核心)
    │ └── (使用) Context (用于获取 prefill/decode 状态及相关参数)

    ├── q_norm: RMSNorm (对 Query 向量进行归一化)

    └── k_norm: RMSNorm (对 Key 向量进行归一化)

    其中 qkv 线性层使用 TP 是按列切分权重,o 线性层是按行切分权重。

    ColumnParallelLinear 类源码剖析(列并行)

    ColumnParallelLinear 和 RowParallelLinear 类的基类是 LinearBase,它为所有并行及量化的线性层提供了基础框架和通用属性,其核心支持了普通与量化两种线性层实现。 ColumnParallelLinear 类是一种实现列并行(Column Parallelism)的线性层。其中权重 A 按第二维(输出维度)列切分,被划分为多个子矩阵:

    A

    =

    [

    A

    1

    ,

    A

    2

    ,

    ,

    A

    p

    ]

    A = [A_1, A_2, \\dots, A_p]

    A=[A1,A2,,Ap],每个 GPU 只存储并计算其中一块

    A

    i

    A_i

    Ai,对应输出

    Y

    i

    =

    X

    A

    i

    Y_i = X A_i

    Yi=XAi。如果需要完整的输出,多个 GPU 上的

    Y

    i

    Y_i

    Yi<span class=

    鸿蒙分布式能力与超级终端实战:设备发现·数据同步·任务流转

    master阅读(50)

    在这里插入图片描述

    在 HarmonyOS NEXT 的设计哲学中,"设备"不再是一个个孤立的信息孤岛,而是可以随时被调用的分布式资源。超级终端(Super Device)的理念,使得手机、平板、智慧屏、手表、PC、车机甚至 IoT 设备能够协同工作,实现能力的跨设备延伸。这一切背后,分布式软总线(Distributed Soft Bus)扮演了"神经系统"的角色——它负责设备发现、认证、数据传输和任务调度,对开发者完全透明却又高度可控。

    本文选取分布式能力中最核心的三条链路——设备发现与信任组管理、跨设备数据同步、跨设备任务流转——配以完整可运行的 ArkTS 示例代码,带你从原理到实践彻底掌握 HarmonyOS NEXT 的超级终端开发范式。


    一、分布式软总线的架构概览

    1.1 软总线的分层模型

    分布式软总线位于 HarmonyOS 系统架构的底层,对上为Ability框架、ArkData数据管理、WantAgent任务调度等模块提供统一的分布式通信能力。其核心可以分为四层:

    • 设备管理层(DeviceManager):负责周边设备的发现、认证、上下线感知与信任组维护。
    • 协议适配层:自动适配 Wi-Fi、蓝牙、NFC、USB 等多种物理传输介质,开发者无需关心底层细节。
    • 分布式调度层(DistributedScheduler):负责跨设备任务的发起、分配与生命周期管理。
    • 数据通道层(DistributedData):提供统一的 KVStore 接口,实现跨设备键值对同步。

    理解这四层的职责边界,有助于在实际开发中选择正确的 API 入口。

    1.2 超级终端的设备角色

    超级终端中,每个设备都有两种动态角色:

    • 可信设备(Trusted Device):已通过设备认证并加入信任组的设备,可以直接进行数据同步和任务流转。
    • 协同设备(Collaborative Device):临时被拉起协同的设备,通常由WantAgent 拉起对方的特定Ability后短暂连接。

    设备角色并非固定,而是根据业务场景动态变化。例如,当手机与平板建立信任组后,两者互为可信设备;但如果手机通过扫码将笔记本拉起协同编辑文档,笔记本在文档协作场景中即为协同设备。


    二、设备发现与信任组管理:DeviceManager 实战

    2.1 核心原理

    设备发现的第一步是感知周围的同局域网或蓝牙范围内的 HarmonyOS 设备。DeviceManager 模块封装了这一过程,提供:

    • startTrustAgent():启动认证代理,弹出认证 UI,等待对方设备扫码确认。
    • authenticateDevice():主动认证一个已发现的设备。
    • getTrustedDeviceList():获取当前信任组中所有已认证设备。
    • checkDeviceAuthentication():检查特定设备是否已认证。

    值得注意的是,设备认证是双向的——设备 A 认证设备 B 后,设备 B 需要在自己的设备管理界面确认,认证才正式生效。这是软总线安全模型的设计核心。

    2.2 完整示例:设备发现与认证

    以下示例展示如何在一个 EntryAbility 中初始化 DeviceManager、监听设备列表变化,并在 UI 上展示已认证设备。

    // entry/src/main/ets/pages/DeviceDiscovery.ets

    import { deviceInfo } from '@kit.BasicServicesKit';
    import { distributedDeviceManager } from '@kit.DistributedServiceKit';
    import { BusinessError } from '@kit.BasicServicesKit';

    class DeviceDiscoveryViewModel {
    // 设备管理器单例,全应用只需一个实例
    private deviceManager: distributedDeviceManager.DeviceManager | null = null;
    // 设备列表状态,用于 UI 绑定
    deviceList: distributedDeviceManager.DeviceBasicInfo[] = [];

    // 设备列表变化回调
    private onDeviceListChange = (data: distributedDeviceManager.DeviceBasicInfo[]): void => {
    this.deviceList = data.filter(device => {
    // 仅保留可信设备(STATE_ACTIVE)
    return device.state === distributedDeviceManager.DeviceState.STATE_ACTIVE;
    });
    console.info(`[DeviceDiscovery] Active devices: ${this.deviceList.length}`);
    };

    // 设备上下线回调
    private onDeviceChange = (type: distributedDeviceManager.SubscribeType,
    data: distributedDeviceManager.DeviceBasicInfo): void => {
    if (type === distributedDeviceManager.SubscribeType.SUBECRIBE_TYPE_DEVICELIST_CHANGE) {
    console.info(`[DeviceDiscovery] Device list changed: ${JSON.stringify(data)}`);
    this.refreshDeviceList();
    }
    };

    async initialize(): Promise<void> {
    const context = getContext(this);

    try {
    // 创建设备管理器实例,传入包名
    this.deviceManager = distributedDeviceManager.createDeviceManager(
    context.applicationInfo.name
    );

    // 注册设备列表变化监听
    this.deviceManager.on('deviceListChange',
    distributedDeviceManager.SubscribeType.SUBECRIBE_TYPE_DEVICELIST_CHANGE,
    this.onDeviceChange
    );

    // 初始加载一次设备列表
    await this.refreshDeviceList();
    console.info('[DeviceDiscovery] DeviceManager initialized successfully');
    } catch (err) {
    const error = err as BusinessError;
    console.error(`[DeviceDiscovery] Init failed: ${error.code}${error.message}`);
    }
    }

    async refreshDeviceList(): Promise<void> {
    if (!this.deviceManager) return;

    try {
    const list = this.deviceManager.getTrustedDeviceListSync();
    this.deviceList = list.filter(d => d.state === distributedDeviceManager.DeviceState.STATE_ACTIVE);
    console.info(`[DeviceDiscovery] Found ${this.deviceList.length} trusted devices`);
    } catch (err) {
    console.error(`[DeviceDiscovery] Get device list failed: ${(err as BusinessError).message}`);
    }
    }

    async startTrustAgent(context: Context): Promise<void> {
    if (!this.deviceManager) return;

    try {
    // 启动认证代理,拉起系统认证界面
    this.deviceManager.startTrustAgent({
    onError: (code: number, message: string) => {
    console.error(`[DeviceDiscovery] TrustAgent error: ${code}${message}`);
    }
    });
    console.info('[DeviceDiscovery] TrustAgent started');
    } catch (err) {
    console.error(`[DeviceDiscovery] Start TrustAgent failed: ${(err as BusinessError).message}`);
    }
    }

    destroy(): void {
    if (this.deviceManager) {
    this.deviceManager.off('deviceListChange');
    this.deviceManager.release();
    this.deviceManager = null;
    }
    }

    // 获取本机设备 UUID
    getLocalDeviceUuid(): string {
    return deviceInfo.uuid;
    }
    }

    export { DeviceDiscoveryViewModel };

    上述 ViewModel 封装了设备发现的核心逻辑。接下来将其接入 ArkUI 页面:

    // entry/src/main/ets/pages/DeviceDiscoveryPage.ets
    import { DeviceDiscoveryViewModel } from '../viewmodel/DeviceDiscoveryViewModel';

    @Entry
    @Component
    struct DeviceDiscoveryPage {
    @State viewModel: DeviceDiscoveryViewModel = new DeviceDiscoveryViewModel();
    @State localUuid: string = '';
    @State isDiscovering: boolean = false;

    async aboutToAppear(): Promise<void> {
    await this.viewModel.initialize();
    this.localUuid = this.viewModel.getLocalDeviceUuid();
    }

    aboutToDisappear(): void {
    this.viewModel.destroy();
    }

    build() {
    Navigation() {
    Column({ space: 16 }) {
    // 本机信息卡片
    Row() {
    Column() {
    Text('本机 UUID')
    .fontSize(12)
    .fontColor('#999999')
    Text(this.localUuid.substring(0, 8) + '…')
    .fontSize(14)
    .fontFamily('monospace')
    }
    .alignItems(HorizontalAlign.Start)
    }
    .width('100%')
    .padding(16)
    .backgroundColor('#F5F5F5')
    .borderRadius(12)

    // 设备列表标题
    Row() {
    Text('可信设备')
    .fontSize(16)
    .fontWeight(FontWeight.Bold)

    Text(`${this.viewModel.deviceList.length} 台)`)
    .fontSize(14)
    .fontColor('#666666')
    }
    .width('100%')
    .padding({ left: 16, right: 16 })

    // 设备列表
    List() {
    ForEach(this.viewModel.deviceList, (device: distributedDeviceManager.DeviceBasicInfo) => {
    ListItem() {
    Row() {
    Column() {
    Text(device.deviceName)
    .fontSize(15)
    .fontWeight(FontWeight.Medium)
    Text(device.deviceId.substring(0, 16) + '…')
    .fontSize(11)
    .fontColor('#999999')
    .fontFamily('monospace')
    }
    .alignItems(HorizontalAlign.Start)
    .layoutWeight(1)

    Text(device.deviceType?.toString() ?? 'Unknown')
    .fontSize(12)
    .backgroundColor('#E8F5E9')
    .padding({ left: 8, right: 8, top: 4, bottom: 4 })
    .borderRadius(4)
    }
    .padding(16)
    }
    }, (device: distributedDeviceManager.DeviceBasicInfo) => device.deviceId)
    }
    .width('100%')
    .layoutWeight(1)
    .divider({ strokeWidth: 0.5, color: '#EEEEEE', startMargin: 16, endMargin: 16 })

    // 添加设备按钮
    Button('发现新设备', { type: ButtonType.Capsule })
    .width('80%')
    .height(48)
    .onClick(() => {
    this.viewModel.startTrustAgent(getContext(this));
    })
    }
    .width('100%')
    .height('100%')
    .padding(16)
    }
    .title('设备发现')
    .navDestination(this.PageMap)
    }

    PageMap: NavPathStack = new NavPathStack();
    }

    在这里插入图片描述

    上面这段代码展示了完整的设备发现流程:初始化 DeviceManager → 注册设备列表变化监听 → 刷新设备列表 → 展示设备信息 → 启动认证代理。

    需要特别说明的是,startTrustAgent() 启动后,系统会弹出认证二维码界面。另一台设备扫描后,双方设备均需在各自界面确认,设备才会进入 STATE_ACTIVE 状态。认证完成后,onDeviceListChange 回调会自动触发,UI 随之更新。


    三、跨设备数据同步:Distributed KVStore 实战

    3.1 为什么选择 KVStore 而不是普通 AppStorage

    HarmonyOS 提供了多种数据持久化方案:AppStorage、UserInfoRepo、Distributed KVStore。那么何时该用分布式 KVStore?

    简单来说:如果数据需要在多台设备间实时同步,选择 KVStore;如果数据仅存在于本地,选择 AppStorage。 KVStore 的底层实现基于分布式软总线,自动处理冲突合并(Last-Write-Wins 策略)、断点续传和网络切换恢复,对开发者屏蔽了全部传输层细节。

    KVStore 有三种模式:

    • DeviceSingle KVStore:单设备键值存储,不可跨设备同步。
    • DeviceDistributed KVStore:可跨设备同步的分布式 KVStore,同步范围为同一用户下所有可信设备。
    • SingleKVStore + 手动同步:通过 sync() 方法按需触发同步,灵活性最高。

    对于超级终端场景,推荐使用 DeviceDistributed KVStore,因为它天然支持设备组网后自动同步,无需手动调用 sync()。

    3.2 完整示例:分布式笔记同步

    假设我们要实现一个跨设备笔记应用,数据在手机和平板之间实时同步。以下是完整的分布式数据管理架构。

    // entry/src/main/ets/data/NoteModel.ets

    // 笔记数据模型
    interface Note {
    id: string;
    title: string;
    content: string;
    updatedAt: number; // 时间戳,用于冲突解决
    deviceId: string; // 最后修改的设备 ID
    }

    class DistributedNoteStore {
    private kvStore: distributedKVStore.DeviceKVStore | null = null;
    private storeId: string = 'note_distributed_store';
    private currentDeviceId: string = '';

    async initialize(context: Context): Promise<void> {
    // 获取本机设备 ID,用于追踪修改来源
    const options: distributedKVStore.StoreOptions = {
    createIfMissing: true,
    // 开启跨设备加密同步,数据在传输过程中全程加密
    encrypt: true,
    // 允许数据自动同步到其他可信设备
    autoSync: true,
    // 同步策略:优先以本机数据为准
    conflictStrategy: distributedKVStore.ConflictStrategy.VERSION
    };

    try {
    const mgr = distributedKVStore.createKVManager(context);
    this.kvStore = await mgr.getKVStore<distributedKVStore.DeviceKVStore>(
    this.storeId,
    options
    );

    // 监听数据变化(来自本设备或其他设备的变化均会触发)
    this.kvStore.on('dataChange', distributedKVStore.SubscribeType.SUBSCRIBE_TYPE_ALL,
    (data: distributedKVStore.KVStoreDataChange) => {
    console.info(`[NoteStore] Data changed. Inserted: ${data.insertEntries.length}, ` +
    `Updated: ${data.updateEntries.length}, Deleted: ${data.deleteEntries.length}`);
    this.notifyChange();
    }
    );

    console.info('[NoteStore] Initialized successfully');
    } catch (err) {
    console.error(`[NoteStore] Init failed: ${(err as Error).message}`);
    }
    }

    // 保存一条笔记,自动带上时间戳和设备标识
    async saveNote(note: Note): Promise<void> {
    if (!this.kvStore) {
    throw new Error('KVStore not initialized');
    }

    const noteWithMeta: Note = {
    note,
    updatedAt: Date.now(),
    deviceId: this.currentDeviceId
    };

    try {
    await this.kvStore.put(`note_${note.id}`, JSON.stringify(noteWithMeta));
    console.info(`[NoteStore] Saved note: ${note.id}`);
    } catch (err) {
    console.error(`[NoteStore] Save failed: ${(err as Error).message}`);
    }
    }

    // 读取指定笔记
    async getNote(noteId: string): Promise<Note | null> {
    if (!this.kvStore) return null;

    try {
    const raw = await this.kvStore.get(`note_${noteId}`);
    if (raw) {
    return JSON.parse(raw as string) as Note;
    }
    return null;
    } catch (err) {
    console.error(`[NoteStore] Get note failed: ${(err as Error).message}`);
    return null;
    }
    }

    // 获取所有笔记(按更新时间倒序)
    async getAllNotes(): Promise<Note[]> {
    if (!this.kvStore) return [];

    try {
    const entries = await this.kvStore.getEntries(''); // 空字符串匹配所有 key
    const notes: Note[] = [];

    for (const entry of entries) {
    if (entry.key.startsWith('note_')) {
    const note = JSON.parse(entry.value.value as string) as Note;
    notes.push(note);
    }
    }

    // 按更新时间倒序排列
    notes.sort((a, b) => b.updatedAt a.updatedAt);
    return notes;
    } catch (err) {
    console.error(`[NoteStore] GetAll failed: ${(err as Error).message}`);
    return [];
    }
    }

    // 删除笔记
    async deleteNote(noteId: string): Promise<void> {
    if (!this.kvStore) return;

    try {
    await this.kvStore.delete(`note_${noteId}`);
    console.info(`[NoteStore] Deleted note: ${noteId}`);
    } catch (err) {
    console.error(`[NoteStore] Delete failed: ${(err as Error).message}`);
    }
    }

    // 手动触发跨设备同步(仅在非自动模式或需要立即同步时使用)
    async forceSync(): Promise<void> {
    if (!this.kvStore) return;

    try {
    await this.kvStore.sync(
    distributedKVStore.SyncMode.PUSH_ONLY,
    3000 // 超时 3 秒
    );
    console.info('[NoteStore] Force sync triggered');
    } catch (err) {
    console.error(`[NoteStore] Sync failed: ${(err as Error).message}`);
    }
    }

    private changeCallback: (() => void) | null = null;

    onChange(callback: () => void): void {
    this.changeCallback = callback;
    }

    private notifyChange(): void {
    if (this.changeCallback) {
    this.changeCallback();
    }
    }

    destroy(): void {
    if (this.kvStore) {
    this.kvStore.off('dataChange');
    this.kvStore = null;
    }
    }
    }

    export { DistributedNoteStore, Note };

    接下来,将上述数据层接入笔记编辑器页面:

    // entry/src/main/ets/pages/NoteEditorPage.ets
    import { DistributedNoteStore, Note } from '../data/NoteModel';
    import { BusinessError } from '@kit.BasicServicesKit';

    @Entry
    @Component
    struct NoteEditorPage {
    @State note: Note = {
    id: '',
    title: '',
    content: '',
    updatedAt: 0,
    deviceId: ''
    };
    @State allNotes: Note[] = [];
    @State isEditing: boolean = false;
    @State syncStatus: string = 'idle';

    private store: DistributedNoteStore = new DistributedNoteStore();

    async aboutToAppear(): Promise<void> {
    await this.store.initialize(getContext(this));
    this.store.onChange(() => {
    // 数据变化时重新加载列表
    this.loadNotes();
    });
    await this.loadNotes();
    }

    async loadNotes(): Promise<void> {
    this.allNotes = await this.store.getAllNotes();
    }

    async saveCurrentNote(): Promise<void> {
    if (!this.note.title.trim()) return;

    // 生成唯一 ID(生产环境建议用 UUID 库)
    if (!this.note.id) {
    this.note.id = `note_${Date.now()}_${Math.random().toString(36).substring(2, 9)}`;
    }

    this.syncStatus = 'syncing';
    await this.store.saveNote(this.note);
    this.syncStatus = 'synced';

    setTimeout(() => { this.syncStatus = 'idle'; }, 2000);
    await this.loadNotes();
    }

    async deleteNote(id: string): Promise<void> {
    await this.store.deleteNote(id);
    await this.loadNotes();
    }

    newNote(): void {
    this.note = { id: '', title: '', content: '', updatedAt: 0, deviceId: '' };
    this.isEditing = true;
    }

    editNote(note: Note): void {
    this.note = { note };
    this.isEditing = true;
    }

    build() {
    NavDestination() {
    Column() {
    // 同步状态指示器
    Row() {
    if (this.syncStatus === 'syncing') {
    Text('↻ 同步中…')
    .fontSize(12)
    .fontColor('#1976D2')
    } else if (this.syncStatus === 'synced') {
    Text('✓ 已同步')
    .fontSize(12)
    .fontColor('#388E3C')
    }
    }
    .width('100%')
    .height(24)
    .padding({ left: 16 })

    if (this.isEditing) {
    // 编辑视图
    Column({ space: 12 }) {
    TextInput({ placeholder: '标题', text: this.note.title })
    .width('100%')
    .height(48)
    .fontSize(16)
    .onChange((v: string) => { this.note.title = v; })

    TextArea({ placeholder: '内容', text: this.note.content })
    .width('100%')
    .layoutWeight(1)
    .fontSize(15)
    .onChange((v: string) => { this.note.content = v; })

    Row({ space: 12 }) {
    Button('保存', { type: ButtonType.Capsule })
    .onClick(() => this.saveCurrentNote())

    Button('取消', { type: ButtonType.Capsule })
    .type(ButtonType.Normal)
    .backgroundColor('#EEEEEE')
    .fontColor('#333333')
    .onClick(() => { this.isEditing = false; })
    }
    .width('100%')
    .padding(16)
    }
    .width('100%')
    .height('100%')
    .padding(16)
    } else {
    // 笔记列表视图
    Column() {
    List() {
    ForEach(this.allNotes, (note: Note) => {
    ListItem() {
    Column({ space: 6 }) {
    Text(note.title)
    .fontSize(15)
    .fontWeight(FontWeight.Medium)
    .maxLines(1)
    .textOverflow({ overflow: TextOverflow.Ellipsis })

    Text(note.content)
    .fontSize(13)
    .fontColor('#666666')
    .maxLines(2)
    .textOverflow({ overflow: TextOverflow.Ellipsis })

    Row() {
    Text(new Date(note.updatedAt).toLocaleString())
    .fontSize(11)
    .fontColor('#AAAAAA')

    if (note.deviceId) {
    Text(` 来自: ${note.deviceId.substring(0, 6)}`)
    .fontSize(11)
    .fontColor('#90CAF9')
    }
    }
    }
    .width('100%')
    .padding(12)
    .alignItems(HorizontalAlign.Start)
    }
    .swipeAction({
    end: {
    label: '删除',
    color: '#F44336',
    action: () => this.deleteNote(note.id)
    }
    })
    .onClick(() => this.editNote(note))
    }, (note: Note) => note.id)
    }
    .width('100%')
    .layoutWeight(1)
    .divider({ strokeWidth: 0.5, color: '#EEEEEE', startMargin: 16, endMargin: 16 })
    }
    .width('100%')
    .layoutWeight(1)

    // 新建按钮
    Button({ type: ButtonType.Circle }) {
    Text('+')
    .fontSize(28)
    .fontWeight(FontWeight.Light)
    }
    .width(56)
    .height(56)
    .backgroundColor('#1976D2')
    .alignItems(VerticalAlign.Center)
    .justifyContent(FlexAlign.Center)
    .position({ x: '75%', y: '85%' })
    .onClick(() => this.newNote())
    }
    }
    .width('100%')
    .height('100%')
    }
    .title('分布式笔记')
    .onBackPressed(() => {
    this.isEditing = false;
    return true;
    })
    }
    }

    从上述两段代码可以看出,KVStore 的使用体验非常接近本地存储——put、get、delete、getEntries 这些操作与本地 API 完全一致,但背后会自动完成跨设备数据同步。当平板上编辑了一条笔记,手机端几乎可以实时感知到变化(通常在 1~3 秒内),无需任何额外代码。

    数据冲突的处理同样值得注意。KVStore 默认使用 Last-Write-Wins 策略:当同一 key 在多台设备被同时修改时,以 updatedAt 时间戳最大的版本为准。在笔记场景下,这种策略是合理的;但如果你的业务需要更精细的冲突保留(比如保留两个版本),可以在 Note 模型中维护一个 versions 数组,由应用层自行合并。


    四、跨设备任务流转:WantAgent 与分布式调度实战

    4.1 什么是任务流转

    任务流转(Task Continuity)是超级终端最直观的能力之一:用户在手机上开始编辑文档、查看地图或者播放音乐,可以随时将任务"迁移"到平板或车机上继续操作,全程无需重新打开应用或手动传输数据。

    从技术角度看,任务流转的实现依赖 WantAgent 和 DistributedScheduler 两个模块:

    • WantAgent:封装了一个"意图"(Want),包含目标设备、目标Ability、传参数据。相当于一个跨设备的函数调用请求。
    • DistributedScheduler:负责在设备间传递 WantAgent,并管理目标设备上 Ability 的生命周期。

    任务流转有三种典型模式:

  • 拉起(Start):在目标设备上启动指定的 Ability,并传递数据。
  • 续写(Continue):将本设备当前 Ability 的状态快照传到目标设备,目标设备以该状态继续运行。
  • 后台运行(BackgroundRunning):在后台设备上持续运行任务(如音乐播放、导航),不影响前台设备。
  • 4.2 完整示例:文档续写与任务迁移

    以下示例演示完整的任务流转场景:从设备 A 发起文档续写请求,设备 B 接收并以相同状态继续编辑。

    // entry/src/main/ets/data/TaskFlowManager.ets

    import { distributedMissionManager } from '@kit.MiscServicesKit';
    import { wantAgent, Want } from '@kit.AbilityKit';
    import { Caller, Callee } from '@kit.IPCKit';
    import { distributedDeviceManager } from '@kit.DistributedServiceKit';
    import { BusinessError } from '@kit.BasicServicesKit';

    // 任务流转的上下文数据结构
    interface DocumentTaskContext {
    taskId: string;
    documentId: string;
    documentTitle: string;
    cursorPosition: number;
    scrollOffset: number;
    deviceName: string;
    timestamp: number;
    }

    class TaskFlowManager {
    // 记录本端注册的 Callee,用于接收来自其他设备的续写请求
    private calleeStub: Callee | null = null;
    private deviceManager: distributedDeviceManager.DeviceManager | null = null;
    private localDeviceId: string = '';

    async initialize(context: Context): Promise<void> {
    // 获取本机设备 ID
    const deviceInfo = await distributedMissionManager.getLocalDeviceInfo();
    this.localDeviceId = deviceInfo.id;

    // 创建设备管理器(用于查询目标设备信息)
    this.deviceManager = distributedDeviceManager.createDeviceManager(
    context.applicationInfo.name
    );

    // 注册跨设备续写回调
    await this.registerContinueCallback(context);
    console.info(`[TaskFlow] Initialized, local device: ${this.localDeviceId}`);
    }

    // 注册续写回调——当其他设备请求续写到本设备时触发
    private async registerContinueCallback(context: Context): Promise<void> {
    try {
    const subscriberInfo: distributedMissionManager.SubscriberInfo = {
    subscriberName: 'DocumentContinueSubscriber',
    subscriberId: 'document_continue_001'
    };

    const callback: distributedMissionManager.ContinueCallback = {
    onContinue: this.handleContinueRequest.bind(this),
    onComplete: this.handleContinueComplete.bind(this),
    onTimeout: this.handleContinueTimeout.bind(this)
    };

    await distributedMissionManager.subscribeMissionCallback(subscriberInfo, callback);
    console.info('[TaskFlow] Continue callback registered');
    } catch (err) {
    const error = err as BusinessError;
    console.error(`[TaskFlow] Subscribe failed: ${error.code}${error.message}`);
    }
    }

    // 处理续写请求:将文档上下文序列化后启动文档编辑器
    private async handleContinueRequest(ctx: Context,
    params: Record<string, Object>): Promise<void> {
    console.info('[TaskFlow] Continue request received');

    // 从 params 中提取文档上下文
    const taskContext = params['taskContext'] as DocumentTaskContext;
    if (!taskContext) {
    console.error('[TaskFlow] taskContext is null');
    return;
    }

    // 将上下文数据通过 AppStorage 传递给 Ability
    AppStorage.setOrCreate('continueDocumentContext', taskContext);

    // 启动文档编辑 Ability,携带续写参数
    const want: Want = {
    deviceId: this.localDeviceId,
    bundleName: 'com.example.documentapp',
    abilityName: 'DocumentEditorAbility',
    parameters: {
    'taskContext': taskContext,
    'isContinued': true
    }
    };

    try {
    const controller = await wantAgent.getWantAgent(want);
    await wantAgent.startWantAgent(controller, {
    wantStartTime: 5000
    });
    console.info(`[TaskFlow] Launched DocumentEditorAbility with context: ${taskContext.documentTitle}`);
    } catch (err) {
    console.error(`[TaskFlow] Launch failed: ${(err as BusinessError).message}`);
    }
    }

    private handleContinueComplete(sourceDeviceId: string): void {
    console.info(`[TaskFlow] Continue completed from device: ${sourceDeviceId}`);
    }

    private handleContinueTimeout(sourceDeviceId: string): void {
    console.warn(`[TaskFlow] Continue timeout from device: ${sourceDeviceId}`);
    }

    // 发起任务流转:将当前文档上下文发送到目标设备
    async continueToDevice(context: Context, targetDeviceId: string,
    documentContext: DocumentTaskContext): Promise<void> {
    if (!targetDeviceId || targetDeviceId === this.localDeviceId) {
    console.warn('[TaskFlow] Invalid or local target device');
    return;
    }

    // 构造续写 Want
    const continueWant: Want = {
    deviceId: targetDeviceId,
    bundleName: 'com.example.documentapp',
    abilityName: 'DocumentEditorAbility',
    flags: 0x00000001, // FLAG_ABILITY_CONTINUATION
    parameters: {
    'taskContext': documentContext,
    'isContinued': true,
    'sourceDeviceId': this.localDeviceId
    }
    };

    // 获取目标设备信息(用于日志展示)
    if (this.deviceManager) {
    const device = this.deviceManager.getDeviceInfo(targetDeviceId);
    console.info(`[TaskFlow] Continuing to: ${device?.deviceName ?? targetDeviceId}`);
    }

    try {
    // 创建续写代理
    const agentInfo: wantAgent.wantAgentInfo = {
    wants: [continueWant],
    operationType: wantAgent.OperationType.CONTINUATION,
    requestCode: 0,
    wantAgentFlags: [wantAgent.WantAgentFlags.UPDATE_PRESENT_FLAG]
    };

    const agent = await wantAgent.getWantAgent(agentInfo);

    // 执行续写(系统会在目标设备上拉起对应 Ability)
    await wantAgent.startWantAgent(agent, {
    wantStartTime: 10000
    });

    console.info(`[TaskFlow] Task continued to ${targetDeviceId}`);
    } catch (err) {
    const error = err as BusinessError;
    console.error(`[TaskFlow] Continue failed: ${error.code}${error.message}`);
    throw error;
    }
    }

    // 查询可信设备列表,供用户选择流转目标
    getAvailableDevices(): distributedDeviceManager.DeviceBasicInfo[] {
    if (!this.deviceManager) return [];

    const allDevices = this.deviceManager.getTrustedDeviceListSync();
    // 过滤掉本设备
    return allDevices.filter(d =>
    d.state === distributedDeviceManager.DeviceState.STATE_ACTIVE &&
    d.deviceId !== this.localDeviceId
    );
    }

    // 获取当前任务快照(用于续写时传递完整状态)
    createDocumentSnapshot(documentId: string, title: string,
    cursorPos: number, scrollOffset: number): DocumentTaskContext {
    return {
    taskId: `task_${Date.now()}`,
    documentId,
    documentTitle: title,
    cursorPosition: cursorPos,
    scrollOffset: scrollOffset,
    deviceName: this.localDeviceId,
    timestamp: Date.now()
    };
    }

    destroy(): void {
    if (this.calleeStub) {
    this.calleeStub = null;
    }
    if (this.deviceManager) {
    this.deviceManager.release();
    this.deviceManager = null;
    }
    }
    }

    export { TaskFlowManager, DocumentTaskContext };

    下面是将任务流转能力集成到文档编辑器 UI 的完整页面代码:

    // entry/src/main/ets/pages/DocumentEditorPage.ets
    import { TaskFlowManager, DocumentTaskContext } from '../data/TaskFlowManager';
    import { distributedDeviceManager } from '@kit.DistributedServiceKit';

    @Entry
    @Component
    struct DocumentEditorPage {
    @State documentId: string = 'doc_001';
    @State documentTitle: string = '项目需求文档';
    @State documentContent: string = '';
    @State cursorPosition: number = 0;
    @State scrollOffset: number = 0;
    @State availableDevices: distributedDeviceManager.DeviceBasicInfo[] = [];
    @State showDevicePicker: boolean = false;
    @State flowStatus: string = '';

    private taskFlowManager: TaskFlowManager = new TaskFlowManager();
    private scroller: Scroller = new Scroller();

    async aboutToAppear(): Promise<void> {
    await this.taskFlowManager.initialize(getContext(this));

    // 检查是否有续写上下文传入(从其他设备流转过来)
    const continueCtx = AppStorage.get<DocumentTaskContext>('continueDocumentContext');
    if (continueCtx) {
    this.documentId = continueCtx.documentId;
    this.documentTitle = continueCtx.documentTitle;
    this.cursorPosition = continueCtx.cursorPosition;
    this.scrollOffset = continueCtx.scrollOffset;
    console.info(`[DocEditor] Opened with continue context: ${continueCtx.documentTitle}`);
    // 清除上下文,防止下次重复使用
    AppStorage.delete('continueDocumentContext');
    }

    // 加载可用流转设备
    this.availableDevices = this.taskFlowManager.getAvailableDevices();
    }

    aboutToDisappear(): void {
    this.taskFlowManager.destroy();
    }

    // 获取当前文档快照
    getCurrentSnapshot(): DocumentTaskContext {
    return this.taskFlowManager.createDocumentSnapshot(
    this.documentId,
    this.documentTitle,
    this.cursorPosition,
    this.scrollOffset
    );
    }

    // 执行跨设备流转
    async continueToDevice(targetDeviceId: string): Promise<void> {
    this.flowStatus = '流转中…';
    this.showDevicePicker = false;

    try {
    const snapshot = this.getCurrentSnapshot();
    await this.taskFlowManager.continueToDevice(getContext(this), targetDeviceId, snapshot);
    this.flowStatus = '已流转';

    // 3 秒后隐藏状态提示
    setTimeout(() => { this.flowStatus = ''; }, 3000);
    } catch {
    this.flowStatus = '流转失败';
    setTimeout(() => { this.flowStatus = ''; }, 3000);
    }
    }

    build() {
    Stack() {
    Column({ space: 16 }) {
    // 标题栏
    Row() {
    TextInput({ text: this.documentTitle, placeholder: '文档标题' })
    .fontSize(18)
    .fontWeight(FontWeight.Bold)
    .layoutWeight(1)
    .onChange((v: string) => { this.documentTitle = v; })

    // 流转按钮
    if (this.availableDevices.length > 0) {
    Row({ space: 4 }) {
    Image($r('sys.media.ohos_ic_public_arrow_right'))
    .width(16)
    .height(16)
    .fillColor('#1976D2')
    Text('流转')
    .fontSize(14)
    .fontColor('#1976D2')
    }
    .padding({ left: 12, right: 8, top: 6, bottom: 6 })
    .border({ width: 1, color: '#1976D2', radius: 16 })
    .onClick(() => { this.showDevicePicker = true; })
    }
    }
    .width('100%')
    .padding({ left: 16, right: 16, top: 12 })

    Divider().padding({ left: 16, right: 16 })

    // 文档内容编辑区
    Scroll(this.scroller) {
    TextArea({ text: this.documentContent, placeholder: '在此输入文档内容…' })
    .width('100%')
    .minHeight(400)
    .fontSize(15)
    .lineSpacing({ leading: 8, trailing: 8 })
    .onChange((v: string) => {
    this.documentContent = v;
    })
    .onTextSelectionChange((selection: { selectionStart: number, selectionEnd: number }) => {
    this.cursorPosition = selection.selectionStart;
    })
    }
    .scrollable(ScrollDirection.Vertical)
    .scrollBar(BarState.Auto)
    .layoutWeight(1)
    .padding(16)

    // 流转状态提示
    if (this.flowStatus) {
    Row() {
    if (this.flowStatus === '流转中…') {
    LoadingProgress()
    .width(16)
    .height(16)
    } else if (this.flowStatus === '已流转') {
    Text('✓')
    .fontSize(14)
    .fontColor('#388E3C')
    }
    Text(` ${this.flowStatus}`)
    .fontSize(13)
    .fontColor(this.flowStatus.includes('失败') ? '#D32F2F' : '#666666')
    }
    .width('100%')
    .padding(12)
    .backgroundColor('#FAFAFA')
    }
    }
    .width('100%')
    .height('100%')

    // 设备选择器浮层
    if (this.showDevicePicker) {
    Column() {
    // 半透明遮罩
    Column()
    .width('100%')
    .height('100%')
    .backgroundColor('rgba(0,0,0,0.4)')
    .onClick(() => { this.showDevicePicker = false; })

    // 选择面板
    Column() {
    Text('流转到')
    .fontSize(16)
    .fontWeight(FontWeight.Bold)
    .width('100%')
    .padding({ left: 20, bottom: 16 })

    ForEach(this.availableDevices, (device: distributedDeviceManager.DeviceBasicInfo) => {
    Row() {
    Column() {
    Text(device.deviceName)
    .fontSize(15)
    .fontWeight(FontWeight.Medium)
    Text(device.deviceType?.toString() ?? '')
    .fontSize(12)
    .fontColor('#666666')
    }
    .layoutWeight(1)
    .alignItems(HorizontalAlign.Start)

    Text('流转')
    .fontSize(13)
    .fontColor('#1976D2')
    }
    .width('100%')
    .padding(16)
    .onClick(() => this.continueToDevice(device.deviceId))
    })
    }
    .width('100%')
    .backgroundColor('#FFFFFF')
    .borderRadius({ topLeft: 16, topRight: 16 })
    .transition({ type: TransitionType.Insert, translate: { y: 300 } })
    }
    .width('100%')
    .height('100%')
    .position({ x: 0, y: 0 })
    }
    }
    .width('100%')
    .height('100%')
    }
    }

    4.3 后台任务流转:音乐续播示例

    除了文档续写,任务流转的另一个高频场景是音乐跨设备续播——用户在家用手机听音乐,出门时将播放任务无缝迁移到车机上。以下是后台运行模式的实现:

    // entry/src/main/ets/service/MusicContinuityService.ets

    import { wantAgent } from '@kit.AbilityKit';
    import { BackgroundTaskManager, BackgroundTaskTiming } from '@kit.BackgroundTasksKit';
    import { BusinessError } from '@kit.BasicServicesKit';

    class MusicContinuityService {
    private readonly TARGET_BUNDLE = 'com.example.carkitapp';
    private readonly TARGET_ABILITY = 'MusicPlayerAbility';

    // 将音乐播放任务流转到车机
    async transferMusicToCar(carDeviceId: string, trackInfo: MusicTrack): Promise<void> {
    const continueWant = {
    deviceId: carDeviceId,
    bundleName: this.TARGET_BUNDLE,
    abilityName: this.TARGET_ABILITY,
    flags: 0x00000001, // FLAG_ABILITY_CONTINUATION
    parameters: {
    'trackId': trackInfo.trackId,
    'trackTitle': trackInfo.title,
    'artist': trackInfo.artist,
    'albumArt': trackInfo.albumArtUri,
    'playbackPosition': trackInfo.currentPosition, // 当前播放位置
    'mode': 'continue'
    }
    };

    const agentInfo: wantAgent.wantAgentInfo = {
    wants: [continueWant],
    operationType: wantAgent.OperationType.CONTINUATION,
    requestCode: 1,
    wantAgentFlags: [wantAgent.WantAgentFlags.UPDATE_PRESENT_FLAG]
    };

    const agent = await wantAgent.getWantAgent(agentInfo);

    // 启动车机上的音乐播放器,并携带播放状态
    await wantAg

    © 2010-2026   171主机测评   网站地图

    请求次数:49 次,加载用时:2.821 秒,内存占用:40.48 MB