You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PySpark技术问询:如何提取DataFrame中距离矩阵的上三角?

提取PySpark DataFrame中距离矩阵的三角部分

这个需求完全可以用PySpark原生操作实现,全程保留分布式计算优势,不用依赖numpy这类本地库!下面是具体的实现思路和代码示例:

核心思路

我们需要给每行标记连续的行索引,然后针对每个列,根据行索引和列索引的相对关系来保留三角区域的值,其余区域置0:

  • 提取上三角(含对角线):当行索引 ≤ 列索引时保留原值,否则设为0
  • 提取下三角(含对角线):当行索引 ≥ 列索引时保留原值,否则设为0

完整代码示例

1. 创建测试用的距离矩阵DataFrame

先构建你提供的示例数据:

from pyspark.sql import SparkSession
from pyspark.sql import functions as F
from pyspark.sql.window import Window

# 初始化Spark会话
spark = SparkSession.builder.appName("triangular_matrix").getOrCreate()

# 构建示例距离矩阵
data = [
    (1, 2, 3, 4),
    (2, 1, 2, 3),
    (3, 2, 1, 2),
    (4, 3, 2, 1)
]
df = spark.createDataFrame(data, ["col0", "col1", "col2", "col3"])
df.show()

2. 添加连续行索引

用窗口函数生成从0开始的连续行索引,和列索引的起始对齐:

# 生成行索引(从0开始)
window_spec = Window.orderBy(F.monotonically_increasing_id())
df_with_row_idx = df.withColumn("row_idx", F.row_number().over(window_spec) - 1)
df_with_row_idx.show()

3. 转换为三角矩阵

遍历所有列,通过条件判断保留目标三角区域的值:

# 获取所有原始列名
original_cols = df.columns

# 对每个列进行三角区域转换
transformed_columns = []
for col_idx, col_name in enumerate(original_cols):
    # 这里的条件对应保留上三角(含对角线),如果要下三角就改成 F.col("row_idx") >= col_idx
    transformed_col = F.when(F.col("row_idx") <= col_idx, F.col(col_name)).otherwise(0).alias(col_name)
    transformed_columns.append(transformed_col)

# 生成结果DataFrame,移除行索引列
result_df = df_with_row_idx.select(*transformed_columns)
result_df.show()

运行后会得到你想要的结果:

+----+----+----+----+
|col0|col1|col2|col3|
+----+----+----+----+
|   1|   2|   3|   4|
|   0|   1|   2|   3|
|   0|   0|   1|   2|
|   0|   0|   0|   1|
+----+----+----+----+

为什么这个方案更好?

  • 完全分布式:全程没有把数据拉到Driver端,充分利用Spark的集群并发能力,适合处理大规模距离矩阵
  • 灵活调整:只需要修改when里的条件,就能快速切换上下三角的提取逻辑
  • 无额外依赖:只用PySpark原生函数,避免了numpy这类本地库带来的性能瓶颈

内容的提问来源于stack exchange,提问作者Barry Behrmann II

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.20 06:54:40