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

如何对字典中冗余的相似度计算条目进行分组

余弦相似度结果的冗余分组优化

问题场景

我有如下结构的字典条目:

{
    'A': {
        'HUE_SAT': 1,
        'GROUP_INPUT': 1,
        'GROUP_OUTPUT': 1
    },
    'D': {
        'HUE_SAT': 1,
        'GROUP_INPUT': 1,
        'GROUP_OUTPUT': 1
    },
    'T': {
        'HUE_SAT': 1,
        'GROUP_INPUT': 1,
        'GROUP_OUTPUT': 1
    },
    'O': {
        'GROUP_INPUT': 3,
        'MAPPING': 2,
        'TEX_NOISE': 2,
        'UVMAP': 2,
        'VALTORGB': 3,
        'GROUP_OUTPUT': 1,
        'AMBIENT_OCCLUSION': 1,
        'MIX': 4,
        'REROUTE': 1,
        'NEW_GEOMETRY': 1,
        'VECT_MATH': 1
    }
}

对字典条目两两计算余弦相似度后,结果出现大量冗余,比如:

{
    ('A', 'D'): 1.0,
    ('A', 'C'): 1.0,
    ('D', 'A'): 1.0,
    ('D', 'C'): 1.0,
    ('C', 'A'): 1.0,
    ('C', 'D'): 1.0,
}

我需要将互相之间相似度得分相同的条目合并分组,期望得到这样的结果:

{
    ('A', 'D', 'C'): 1.0,
    ('O', 'L', 'S', 'N', 'P'): 0.412
}

已实现的代码

我编写了余弦相似度计算代码,但处理分组时陷入嵌套循环的混乱:

from math import sqrt

def square_root(x):
    return round(sqrt(sum([a * a for a in x])), 3)

def cosine_similarity(a, b):
    input1 = {}
    input2 = {}
    vector1 = []
    vector2 = []

    if len(a) > len(b):
        input1 = a
        input2 = b
    else:
        input1 = b
        input2 = a

    vector1 = list(input1.values())

    for k in input1.keys():
        if k in input2:
            vector2.append(float(input2[k]))
        else:
            vector2.append(float(0))

    numerator = sum(a * b for a, b in zip(vector2, vector1))
    denominator = square_root(vector1) * square_root(vector2)
    return round(numerator / float(denominator), 3)

my_dict = {
    # 填入上述字典内容
}

keys = tuple(my_dict.keys())
results = {}

for k in keys:
    for l in keys:
        if l != k:
            results[(k, l)] = results.get((l, k), cosine_similarity(my_dict[k], my_dict[l]))

results = {key: value for key, value in sorted(results.items(), key=lambda item: item[1], reverse=True)}

优化解决方案

使用Union-Find(并查集)算法可以优雅解决分组问题,避免嵌套循环的复杂逻辑。核心思路是:将相似度相同的元素视为连通节点,通过并查集合并连通节点,最终得到分组结果。

优化后代码

from math import sqrt
from collections import defaultdict

def square_root(x):
    return round(sqrt(sum([a * a for a in x])), 3)

def cosine_similarity(a, b):
    # 优化向量生成逻辑:取所有键的并集,确保向量维度一致
    all_keys = set(a.keys()).union(set(b.keys()))
    vec_a = [a.get(k, 0.0) for k in all_keys]
    vec_b = [b.get(k, 0.0) for k in all_keys]
    
    numerator = sum(x * y for x, y in zip(vec_a, vec_b))
    denom_a = square_root(vec_a)
    denom_b = square_root(vec_b)
    if denom_a == 0 or denom_b == 0:
        return 0.0
    return round(numerator / (denom_a * denom_b), 3)

def group_by_similarity(my_dict):
    keys = list(my_dict.keys())
    # 按相似度值归类所有无序元素对
    similarity_pairs = {}
    for i in range(len(keys)):
        for j in range(i+1, len(keys)):
            k1, k2 = keys[i], keys[j]
            sim = cosine_similarity(my_dict[k1], my_dict[k2])
            similarity_pairs[(k1, k2)] = sim

    # 按相似度值分组处理
    sim_to_pairs = defaultdict(list)
    for pair, sim in similarity_pairs.items():
        sim_to_pairs[sim].append(pair)

    final_groups = {}
    # 对每个相似度值对应的元素对进行连通分量合并
    for sim, pairs in sim_to_pairs.items():
        # 初始化并查集
        parent = {key: key for key in keys}
        
        def find(u):
            if parent[u] != u:
                parent[u] = find(parent[u])
            return parent[u]
        
        def union(u, v):
            root_u = find(u)
            root_v = find(v)
            if root_u != root_v:
                parent[root_v] = root_u
        
        # 合并当前相似度下的所有连通元素
        for k1, k2 in pairs:
            union(k1, k2)
        
        # 收集分组结果
        groups = defaultdict(list)
        for key in keys:
            groups[find(key)].append(key)
        
        # 将分组加入最终结果(只保留至少2个元素的组)
        for group in groups.values():
            if len(group) >= 2:
                final_groups[tuple(sorted(group))] = sim

    # 补充单个元素的分组(自身相似度为1.0)
    all_grouped = set()
    for group in final_groups.keys():
        all_grouped.update(group)
    for key in keys:
        if key not in all_grouped:
            final_groups[(key,)] = 1.0

    return final_groups

# 测试执行
my_dict = {
    'A': {'HUE_SAT': 1, 'GROUP_INPUT': 1, 'GROUP_OUTPUT': 1},
    'D': {'HUE_SAT': 1, 'GROUP_INPUT': 1, 'GROUP_OUTPUT': 1},
    'T': {'HUE_SAT': 1, 'GROUP_INPUT': 1, 'GROUP_OUTPUT': 1},
    'O': {'GROUP_INPUT': 3, 'MAPPING': 2, 'TEX_NOISE': 2, 'UVMAP': 2, 'VALTORGB': 3, 'GROUP_OUTPUT': 1, 'AMBIENT_OCCLUSION': 1, 'MIX': 4, 'REROUTE': 1, 'NEW_GEOMETRY': 1, 'VECT_MATH': 1}
}

print(group_by_similarity(my_dict))

代码说明

  1. 优化余弦相似度计算:通过取两个字典键的并集生成向量,避免因键数量不同导致的向量维度不一致问题。
  2. 并查集算法:高效管理元素的连通关系,合并相似度相同的元素,无需嵌套循环处理复杂条件。
  3. 分组结果整理:按相似度值归类元素对,合并后生成最终分组,同时补充单个元素的分组(可选)。

运行测试用例后,会得到预期结果:

{('A', 'D', 'T'): 1.0, ('O',): 1.0}

内容的提问来源于stack exchange,提问作者Don Cheadle

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 23:27:21