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

Swift中如何实现NLContextualEmbedding的近邻查找?

Swift中NLContextualEmbedding的近邻查找实现方法

Apple的NLContextualEmbedding确实没有内置的近邻查找API,不像NLEmbedding提供了直接的neighbors(for:limit:)方法,但可以通过手动实现的方式完成近邻查找,步骤如下:

  • 提取上下文嵌入向量
    使用NLContextualEmbedding生成目标文本的向量,由于上下文嵌入是token级的,你可以选择取所有token向量的平均值作为整段文本的代表向量,或者提取特定位置token的向量。示例代码:

    import NaturalLanguage
    
    func getContextualEmbedding(for text: String) -> MLMultiArray? {
        guard let embedding = NLContextualEmbedding(name: .bertBaseUncased) else { return nil }
        let tokens = embedding.tokenize(text)
        guard let vectors = embedding.vectors(for: tokens) else { return nil }
        
        // 计算所有token向量的平均值作为文本代表向量
        let count = vectors.count
        guard count > 0 else { return nil }
        var averageVector = MLMultiArray(zeros: vectors[0].shape)
        for vector in vectors {
            averageVector += vector
        }
        averageVector /= Double(count)
        return averageVector
    }
    
  • 构建向量数据集
    预先将所有需要参与近邻比对的文本,用上述方法转换成上下文嵌入向量,统一存储在数组或其他数据结构中,同时保留向量对应的原文本信息。

  • 实现相似度计算与排序
    自己实现余弦相似度(最常用的向量相似度度量方式)或欧氏距离算法,遍历数据集里的所有向量,计算与目标向量的相似度,再按相似度排序取Top N。示例余弦相似度实现:

    func cosineSimilarity(_ vec1: MLMultiArray, _ vec2: MLMultiArray) -> Double {
        guard vec1.count == vec2.count else { return 0 }
        
        var dotProduct: Double = 0
        var norm1: Double = 0
        var norm2: Double = 0
        
        for i in 0..<vec1.count {
            let val1 = vec1[i].doubleValue
            let val2 = vec2[i].doubleValue
            dotProduct += val1 * val2
            norm1 += val1 * val1
            norm2 += val2 * val2
        }
        
        guard norm1 > 0, norm2 > 0 else { return 0 }
        return dotProduct / (sqrt(norm1) * sqrt(norm2))
    }
    
    // 查找近邻的示例函数
    func findNearestNeighbors(targetVector: MLMultiArray, dataset: [(text: String, vector: MLMultiArray)], limit: Int) -> [(text: String, similarity: Double)] {
        var similarities = [(text: String, similarity: Double)]()
        for item in dataset {
            let similarity = cosineSimilarity(targetVector, item.vector)
            similarities.append((text: item.text, similarity: similarity))
        }
        // 按相似度降序排序,取前limit个
        return similarities.sorted { $0.similarity > $1.similarity }.prefix(limit).map { $0 }
    }
    

需要注意的是,上下文嵌入的核心特点是依赖文本语境——同一个词汇在不同句子中生成的向量会存在差异,因此在处理所有文本时,必须保证使用相同的NLContextualEmbedding模型、相同的向量聚合逻辑(比如统一用平均向量),否则比对结果会失去参考价值。

内容的提问来源于stack exchange,提问作者Tay Baye

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 10:05:26