难度:中级 题型:笔试(编程题)
题目描述:测试团队用大模型批量生成了几千条用例标题,存在大量换汤不换药的重复项。请实现函数 dedup_titles:输入用例标题列表和对应的 embedding 向量列表,余弦相似度大于 0.92 的标题视为重复,只保留先出现的,返回去重后的标题列表。
代码框架:
def dedup_titles(titles, vectors, threshold=0.92):
"""
titles: List[str] 用例标题
vectors: 与 titles 等长的向量列表(由 embedding 模型生成)
返回去重后的标题列表(保留先出现的)
"""
# 请在此处补全
pass
assert dedup_titles(
["登录成功后跳转首页", "登录成功后跳转到首页", "退出登录返回登录页"],
[[1, 0], [0.99, 0.05], [0, 1]],
) == ["登录成功后跳转首页", "退出登录返回登录页"]
assert dedup_titles([], []) == []
参考答案/答题要点:
思路1:暴力双循环(能过就行)
import math
def cosine(a, b):
dot = sum(x * y for x, y in zip(a, b))
return dot / (math.sqrt(sum(x*x for x in a)) * math.sqrt(sum(y*y for y in b)))
def dedup_titles(titles, vectors, threshold=0.92):
kept = [] # 保留的下标
for i, _ in enumerate(titles):
if all(cosine(vectors[i], vectors[j]) <= threshold for j in kept):
kept.append(i)
return [titles[i] for i in kept]
思路2:numpy 向量化(万级数据再上)
import numpy as np
def dedup_titles(titles, vectors, threshold=0.92):
v = np.array(vectors, dtype=float)
v = v / np.linalg.norm(v, axis=1, keepdims=True) # 归一化后点积=余弦
kept, sim = [], v @ v.T
for i in range(len(titles)):
if all(sim[i][j] <= threshold for j in kept):
kept.append(i)
return [titles[i] for i in kept]
采分点:①余弦相似度公式写对;②“保留先出现”的贪心逻辑;③边界处理(空列表);④加分项:指出 O(n²) 复杂度,海量数据可先聚类/LSH 粗筛再精比。