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

如何加速5V5游戏团队轨迹聚类的组合式距离计算?

5V5游戏团队轨迹聚类的性能优化方案

一、替换TypedDict为Numpy数组,降低距离查询开销

TypedDict的嵌套查询(area_distance_matrix[map_name][area1id][area2id][dist_type])是核心性能瓶颈之一,改用Numpy多维数组可大幅提升访问速度,同时完美兼容Numba优化:

  1. 映射字符串索引为整数:将map_name和dist_type分别映射为整数ID(如{"MapA":0, "MapB":1}、{"euclidean":0, "manhattan":1}),彻底避免字符串哈希查询的额外开销。
  2. 转换为三维Numpy数组:将预计算的tile距离矩阵转换为(num_tiles, num_tiles, num_dist_types)的float32数组,示例代码:
import numpy as np

# 假设原始area_distance_matrix是嵌套字典结构
map_id_map = {"MapA": 0, "MapB": 1}
dist_type_map = {"euclidean": 0, "manhattan": 1}
area_dist_np = {}

for map_name in area_distance_matrix:
    map_id = map_id_map[map_name]
    num_tiles = len(area_distance_matrix[map_name])
    dist_array = np.zeros((num_tiles, num_tiles, len(dist_type_map)), dtype=np.float32)
    for area1 in range(num_tiles):
        for area2 in range(num_tiles):
            for dist_str, dist_id in dist_type_map.items():
                dist_array[area1][area2][dist_id] = area_distance_matrix[map_name][area1][area2][dist_str]
    area_dist_np[map_id] = dist_array

后续查询直接用数组索引(如area_dist_np[map_id][area1][area2][dist_id]),速度远快于字典嵌套查询。

二、用匈牙利算法替代全排列枚举,消除排列生成开销

当前枚举5! = 120种玩家配对排列的方式,本质是求解二分图最小权匹配问题,改用匈牙利算法可将时间复杂度从O(n!)降至O(n³)(n=5时计算量相当,但完全避免排列生成的额外开销),且可通过Numba完全加速:

  1. 实现Numba加速的匈牙利算法:
from numba import njit

@njit(cache=True)
def hungarian_min_avg(matrix):
    n = matrix.shape[0]
    u = np.zeros(n + 1)
    v = np.zeros(n + 1)
    p = np.zeros(n + 1, dtype=np.int32)
    way = np.zeros(n + 1, dtype=np.int32)

    for i in range(1, n + 1):
        p[0] = i
        minv = np.full(n + 1, np.inf)
        used = np.zeros(n + 1, dtype=np.bool_)
        j0 = 0
        while True:
            used[j0] = True
            i0 = p[j0]
            delta = np.inf
            j1 = 0
            for j in range(1, n + 1):
                if not used[j]:
                    cur = matrix[i0 - 1][j - 1] - u[i0] - v[j]
                    if cur < minv[j]:
                        minv[j] = cur
                        way[j] = j0
                    if minv[j] < delta:
                        delta = minv[j]
                        j1 = j
            for j in range(n + 1):
                if used[j]:
                    u[p[j]] += delta
                    v[j] -= delta
                else:
                    minv[j] -= delta
            j0 = j1
            if p[j0] == 0:
                break
        while True:
            j1 = way[j0]
            p[j0] = p[j1]
            j0 = j1
            if j0 == 0:
                break
    total = 0.0
    for j in range(1, n + 1):
        if p[j] != 0:
            total += matrix[p[j] - 1][j - 1]
    return total / n
  1. 重构状态距离计算函数:
@njit(cache=True)
def position_state_distance(team1_tiles, team2_tiles, area_dist_np, dist_id=0):
    # team1_tiles、team2_tiles为长度5的int32 Numpy数组
    dist_matrix = np.zeros((5, 5), dtype=np.float32)
    for i in range(5):
        for j in range(5):
            dist_matrix[i][j] = area_dist_np[team1_tiles[i]][team2_tiles[j]][dist_id]
    return hungarian_min_avg(dist_matrix)

此版本完全避免itertools.permutations的使用,解决Numba与itertools的兼容性问题。

三、轨迹距离计算的向量化与并行优化

1. 平均状态距离的加速

将轨迹存储为(时间步长, 5)的int32 Numpy数组,用Numba加速遍历计算:

@njit(cache=True)
def area_trajectory_distance_avg(traj1, traj2, area_dist_np):
    total = 0.0
    T = traj1.shape[0]
    for t in range(T):
        total += position_state_distance(traj1[t], traj2[t], area_dist_np)
    return total / T

2. DTW距离的Numba加速

实现纯Numba版本的DTW,避免Python原生循环的开销:

@njit(cache=True)
def area_trajectory_distance_dtw(traj1, traj2, area_dist_np):
    len1 = traj1.shape[0]
    len2 = traj2.shape[0]
    dp = np.full((len1 + 1, len2 + 1), np.inf)
    dp[0][0] = 0.0
    for i in range(1, len1 + 1):
        for j in range(1, len2 + 1):
            cost = position_state_distance(traj1[i-1], traj2[j-1], area_dist_np)
            dp[i][j] = cost + min(dp[i-1][j], dp[i][j-1], dp[i-1][j-1])
    return dp[len1][len2] / max(len1, len2)

四、预计算轨迹距离矩阵的并行化

聚类前的距离矩阵预计算是CPU密集型任务,可通过多进程并行利用多核资源:

from joblib import Parallel, delayed

def precompute_dist_matrix(trajectories, area_dist_np, use_dtw=False):
    n = len(trajectories)
    dist_matrix = np.zeros((n, n), dtype=np.float32)
    # 仅计算上三角,再对称复制以减少重复计算
    def compute_pair(i, j):
        if use_dtw:
            d = area_trajectory_distance_dtw(trajectories[i], trajectories[j], area_dist_np)
        else:
            d = area_trajectory_distance_avg(trajectories[i], trajectories[j], area_dist_np)
        return i, j, d
    
    # 用所有CPU核心并行计算
    results = Parallel(n_jobs=-1)(delayed(compute_pair)(i, j) for i in range(n) for j in range(i+1, n))
    for i, j, d in results:
        dist_matrix[i][j] = d
        dist_matrix[j][i] = d
    return dist_matrix

五、其他细节优化

  • 数据类型压缩:用int32存储tile ID,float32存储距离值,减少内存占用与数据传输开销。
  • 缓存复用:用Numba的cache=True装饰器,避免重复编译JIT函数。
  • 轨迹长度对齐:若使用平均状态距离,提前过滤或对齐轨迹长度,避免额外的长度判断开销。

内容的提问来源于stack exchange,提问作者J.N.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 01:38:09