使用向量嵌入与 Qdrant 搜索代码

作者:Qdrant 团队。

本 Notebook 演示如何用向量嵌入浏览代码库、查找相关代码片段:既可以使用自然语言语义查询,也可以按相似逻辑搜索代码。原文提供了实时部署的 Web 界面,可以搜索 Qdrant 代码库。

方法

需要两个模型:

  • 通用自然语言编码器 sentence-transformers/all-MiniLM-L6-v2,下文称 NLP 模型。
  • 用于代码相似性搜索的 jinaai/jina-embeddings-v2-base-code,支持英语和 30 种常用编程语言,序列长度为 8192,下文称代码模型。

NLP 模型需要先把代码转换为接近自然语言的格式。代码模型本身支持多种标准语言,无需预处理,可以直接使用代码。

安装依赖

使用以下软件包:

  • inflection:字符串变换,支持英文单复数以及驼峰转下划线。
  • fastembed:优先面向 CPU 的轻量向量嵌入库,也支持 GPU。
  • qdrant-client:连接 Qdrant 的官方 Python 库。
%pip install inflection qdrant-client fastembed

数据准备

将应用源码切成小块并不简单。函数、类方法、结构体、枚举等语言结构通常是合适的分块:足够大,包含有意义的信息,又足够小,适合嵌入模型有限的上下文窗口。文档字符串、注释及其他元数据也可以丰富内容。

文本搜索主要基于函数签名,而代码搜索可能返回循环等更小片段。因此,如果 NLP 模型找到某个函数签名,代码模型找到其实现的一部分,就需要合并结果。

将代码库按结构切分,并结合位置与上下文信息
将代码库按结构切分,并结合位置与上下文信息

解析代码库

本例使用 Rust 编写的 Qdrant,但方法适用于其他语言。可以使用语言服务器协议 LSP 工具构建代码图,再提取片段。作者使用 rust-analyzer,将解析结果导出为代码智能数据标准 LSIF,再据此导航代码并提取分块。其他语言也有许多相应实现。

最终将片段导出为 JSON,不仅包含代码,还包含项目中的位置上下文。Google Cloud Storage 中提供了已经解析好的 structures.jsonl,下载后作为搜索数据:

!wget https://storage.googleapis.com/tutorial-attachments/code-search/structures.jsonl

读取文件,将每行解析为字典:

import json

structures = []
with open("structures.jsonl", "r") as fp:
    for i, row in enumerate(fp):
        entry = json.loads(row)
        structures.append(entry)

查看一个条目:

structures[0]
{'name': 'InvertedIndexRam',
 'signature': '# [doc = " Inverted flatten index from dimension id to posting list"] # [derive (Debug , Clone , PartialEq)] pub struct InvertedIndexRam { # [doc = " Posting lists for each dimension flattened (dimension id -> posting list)"] # [doc = " Gaps are filled with empty posting lists"] pub postings : Vec < PostingList > , # [doc = " Number of unique indexed vectors"] # [doc = " pre-computed on build and upsert to avoid having to traverse the posting lists."] pub vector_count : usize , }',
 'code_type': 'Struct',
 'docstring': '= " Inverted flatten index from dimension id to posting list"',
 'line': 15,
 'line_from': 13,
 'line_to': 22,
 'context': {'module': 'inverted_index',
  'file_path': 'lib/sparse/src/index/inverted_index/inverted_index_ram.rs',
  'file_name': 'inverted_index_ram.rs',
  'struct_name': None,
  'snippet': '/// Inverted flatten index from dimension id to posting list\n#[derive(Debug, Clone, PartialEq)]\npub struct InvertedIndexRam {\n    /// Posting lists for each dimension flattened (dimension id -> posting list)\n    /// Gaps are filled with empty posting lists\n    pub postings: Vec<PostingList>,\n    /// Number of unique indexed vectors\n    /// pre-computed on build and upsert to avoid having to traverse the posting lists.\n    pub vector_count: usize,\n}\n'}}

将代码转换为自然语言

编程语言语法并非自然语言,通用模型可能无法直接理解。可以去除代码特有形式,并加入模块、类、函数与文件名等上下文:

  1. 提取函数、方法或其他结构的签名。
  2. 把驼峰与下划线名称拆成单词。
  3. 取得文档字符串、注释和重要元数据。
  4. 按照预定义模板组装句子。
  5. 将特殊字符替换为空格。

下面用 inflection 定义 textify:

import inflection
import re

from typing import Dict, Any


def textify(chunk: Dict[str, Any]) -> str:

    # Get rid of all the camel case / snake case
    # - inflection.underscore changes the camel case to snake case
    # - inflection.humanize converts the snake case to human readable form
    name = inflection.humanize(inflection.underscore(chunk["name"]))
    signature = inflection.humanize(inflection.underscore(chunk["signature"]))

    # Check if docstring is provided
    docstring = ""
    if chunk["docstring"]:
        docstring = f"that does {chunk['docstring']} "

    # Extract the location of that snippet of code
    context = f"module {chunk['context']['module']} file {chunk['context']['file_name']}"
    if chunk["context"]["struct_name"]:
        struct_name = inflection.humanize(inflection.underscore(chunk["context"]["struct_name"]))
        context = f"defined in struct {struct_name} {context}"

    # Combine all the bits and pieces together
    text_representation = f"{chunk['code_type']} {name} {docstring}defined as {signature} {context}"

    # Remove any special characters and concatenate the tokens
    tokens = re.split(r"\W", text_representation)
    tokens = filter(lambda x: x, tokens)
    return " ".join(tokens)

转换全部代码块:

text_representations = list(map(textify, structures))

查看一个文本表示:

text_representations[1000]
'Function Hnsw discover precision that does Checks discovery search precision when using hnsw index this is different from the tests in defined as Fn hnsw discover precision module integration file hnsw_discover_test rs'

自然语言嵌入

from fastembed import TextEmbedding

batch_size = 5

nlp_model = TextEmbedding("sentence-transformers/all-MiniLM-L6-v2", threads=0)
nlp_embeddings = nlp_model.embed(text_representations, batch_size=batch_size)

代码嵌入

code_snippets = [structure["context"]["snippet"] for structure in structures]

code_model = TextEmbedding("jinaai/jina-embeddings-v2-base-code")

code_embeddings = code_model.embed(code_snippets, batch_size=batch_size)

建立 Qdrant 集合

Qdrant 支持内存、Docker、Qdrant Cloud 等多种部署方式,详情见安装说明。本文使用内存实例,它仅适合快速原型和测试,是服务器方法的 Python 实现。

创建保存向量的集合:

from qdrant_client import QdrantClient, models

COLLECTION_NAME = "qdrant-sources"

client = QdrantClient(":memory:")  # Use in-memory storage
# client = QdrantClient("http://locahost:6333")  # For Qdrant server

client.create_collection(
    COLLECTION_NAME,
    vectors_config={
        "text": models.VectorParams(
            size=384,
            distance=models.Distance.COSINE,
        ),
        "code": models.VectorParams(
            size=768,
            distance=models.Distance.COSINE,
        ),
    },
)

集合就绪后上传嵌入:

from tqdm import tqdm

points = []
total = len(structures)
print("Number of points to upload: ", total)

for id, (text_embedding, code_embedding, structure) in tqdm(
    enumerate(zip(nlp_embeddings, code_embeddings, structures)), total=total
):
    # FastEmbed returns generators. Embeddings are computed as consumed.
    points.append(
        models.PointStruct(
            id=id,
            vector={
                "text": text_embedding,
                "code": code_embedding,
            },
            payload=structure,
        )
    )

    # Upload points in batches
    if len(points) >= batch_size:
        client.upload_points(COLLECTION_NAME, points=points, wait=True)
        points = []

# Ensure any remaining points are uploaded
if points:
    client.upload_points(COLLECTION_NAME, points=points)

print(f"Total points in collection: {client.count(COLLECTION_NAME).count}")

上传的点立即可被搜索。接下来查询相关代码。

查询代码库

通过 Qdrant Query API,先使用文本嵌入搜索“How do I count points in a collection?”,即“如何统计集合中的点数?”:

query = "How do I count points in a collection?"

hits = client.query_points(
    COLLECTION_NAME,
    query=next(nlp_model.query_embed(query)).tolist(),
    using="text",
    limit=3,
).points

结果包含模块、文件名、得分,以及指向签名的链接:

模块文件名得分签名
operationstypes.rs0.5493385pub struct CountRequestInternal
map_indextypes.rs0.49973965fn get_points_with_value_count
map_indexmutable_map_index.rs0.49941066pub fn get_points_with_value_count

可以看到找到了相关代码结构。再使用代码嵌入:

hits = client.query_points(
    COLLECTION_NAME,
    query=next(code_model.query_embed(query)).tolist(),
    using="code",
    limit=3,
).points

输出:

模块文件名得分签名
field_indexgeo_index.rs0.7217579fn count_indexed_points
numeric_indexmod.rs0.7113214fn count_indexed_points
full_text_indextext_index.rs0.6993165fn count_indexed_points

不同模型的分数不能直接比较,但结果确实不同。文本与代码嵌入可以捕捉代码库的不同方面,因此可以同时查询,再合并结果,得到更相关的片段:

from qdrant_client import models

hits = client.query_points(
    collection_name=COLLECTION_NAME,
    prefetch=[
        models.Prefetch(
            query=next(nlp_model.query_embed(query)).tolist(),
            using="text",
            limit=5,
        ),
        models.Prefetch(
            query=next(code_model.query_embed(query)).tolist(),
            using="code",
            limit=5,
        ),
    ],
    query=models.FusionQuery(fusion=models.Fusion.RRF),
).points
>>> for hit in hits:
...     print(
...         "| ",
...         hit.payload["context"]["module"],
...         " | ",
...         hit.payload["context"]["file_path"],
...         " | ",
...         hit.score,
...         " | `",
...         hit.payload["signature"],
...         "` |",
...     )
|  operations  |  lib/collection/src/operations/types.rs  |  0.5  | ` # [doc = " Count Request"] # [doc = " Counts the number of points which satisfy the given filter."] # [doc = " If filter is not provided, the count of all points in the collection will be returned."] # [derive (Debug , Deserialize , Serialize , JsonSchema , Validate)] # [serde (rename_all = "snake_case")] pub struct CountRequestInternal &#123; # [doc = " Look only for points which satisfies this conditions"] # [validate] pub filter : Option < Filter > , # [doc = " If true, count exact number of points. If false, count approximate number of points faster."] # [doc = " Approximate count might be unreliable during the indexing process. Default: true"] # [serde (default = "default_exact_count")] pub exact : bool , } ` |
|  field_index  |  lib/segment/src/index/field_index/geo_index.rs  |  0.5  | ` fn count_indexed_points (& self) -> usize ` |
|  map_index  |  lib/segment/src/index/field_index/map_index/mod.rs  |  0.33333334  | ` fn get_points_with_value_count < Q > (& self , value : & Q) -> Option < usize > where Q : ? Sized , N : std :: borrow :: Borrow < Q > , Q : Hash + Eq , ` |
|  numeric_index  |  lib/segment/src/index/field_index/numeric_index/mod.rs  |  0.33333334  | ` fn count_indexed_points (& self) -> usize ` |
|  fixtures  |  lib/segment/src/fixtures/payload_context_fixture.rs  |  0.25  | ` fn total_point_count (& self) -> usize ` |
|  map_index  |  lib/segment/src/index/field_index/map_index/mutable_map_index.rs  |  0.25  | ` fn get_points_with_value_count < Q > (& self , value : & Q) -> Option < usize > where Q : ? Sized , N : std :: borrow :: Borrow < Q > , Q : Hash + Eq , ` |
|  id_tracker  |  lib/segment/src/id_tracker/simple_id_tracker.rs  |  0.2  | ` fn total_point_count (& self) -> usize ` |
|  map_index  |  lib/segment/src/index/field_index/map_index/mod.rs  |  0.2  | ` fn count_indexed_points (& self) -> usize ` |
|  map_index  |  lib/segment/src/index/field_index/map_index/mod.rs  |  0.16666667  | ` fn count_indexed_points (& self) -> usize ` |
|  field_index  |  lib/segment/src/index/field_index/stat_tools.rs  |  0.16666667  | ` fn number_of_selected_points (points : usize , values : usize) -> usize ` |

这只是不同模型结果融合的一种方式。实际应用还可能加入重排序、去重和其他后处理。

结果分组

按照 payload 属性分组,可以改进结果组织。本例可以按模块分组。代码嵌入可能给出来自同一模块的多个结果,因此下面限制每个模块只保留一个:

results = client.query_points_groups(
    COLLECTION_NAME,
    query=next(code_model.query_embed(query)).tolist(),
    using="code",
    group_by="context.module",
    limit=5,
    group_size=1,
)
>>> for group in results.groups:
...     for hit in group.hits:
...         print(
...             "| ",
...             hit.payload["context"]["module"],
...             " | ",
...             hit.payload["context"]["file_name"],
...             " | ",
...             hit.score,
...             " | `",
...             hit.payload["signature"],
...             "` |",
...         )
|  field_index  |  geo_index.rs  |  0.7217579  | ` fn count_indexed_points (& self) -> usize ` |
|  numeric_index  |  mod.rs  |  0.7113214  | ` fn count_indexed_points (& self) -> usize ` |
|  fixtures  |  payload_context_fixture.rs  |  0.6993165  | ` fn total_point_count (& self) -> usize ` |
|  map_index  |  mod.rs  |  0.68385994  | ` fn count_indexed_points (& self) -> usize ` |
|  full_text_index  |  text_index.rs  |  0.6660142  | ` fn count_indexed_points (& self) -> usize ` |

教程到此结束。向量嵌入的用途和改进空间仍然很多,欢迎继续实验,并与 Qdrant 团队分享成果。


原文:Code Search with Vector Embeddings and Qdrant。作者/维护者:Qdrant 团队。本文为原文的中文译文;代码保留原文内容。

© 版权声明
THE END
喜欢就支持一下吧
点赞0 分享
评论 抢沙发

请登录后发表评论

    暂无评论内容