토큰·멀티벡터 임베딩

토큰·멀티벡터 임베딩 (Pooling: Token Embed)

토큰 단위 임베딩과 멀티벡터 검색 임베딩을 vLLM으로 생성하는 예제입니다. Jina v4·v3, ColBERT 같은 모델을 활용해 문장·토큰 벡터를 얻고, 오프라인·온라인 RAG 파이프라인에 연결합니다.

출처: 문서

본문

llm.embed()pooling_method를 지정해 토큰 단위 임베딩을 얻거나, 멀티벡터 검색 예제로 문서·쿼리를 여러 벡터로 표현합니다. 온라인 예제는 OpenAI 호환 /v1/embeddings를 사용합니다.

jina_embeddings_v4_offline.py

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import torch

from vllm import LLM
from vllm.config import PoolerConfig
from vllm.inputs import TextPrompt
from vllm.multimodal.utils import fetch_image


def main():
    # Initialize model
    model = LLM(
        model="jinaai/jina-embeddings-v4-vllm-text-matching",
        pooler_config=PoolerConfig(task="token_embed"),
        runner="pooling",
        max_model_len=1024,
        gpu_memory_utilization=0.8,
    )

    # Create text prompts
    text1 = "Ein wunderschöner Sonnenuntergang am Strand"
    text1_prompt = TextPrompt(prompt=f"Query: {text1}")

    text2 = "浜辺に沈む美しい夕日"
    text2_prompt = TextPrompt(prompt=f"Query: {text2}")

    # Create image prompt
    image = fetch_image(
        "https://vllm-public-assets.s3.us-west-2.amazonaws.com/multimodal_asset/eskimo.jpg"  # noqa: E501
    )
    image_prompt = TextPrompt(
        prompt="<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>Describe the image.<|im_end|>\n",  # noqa: E501
        multi_modal_data={"image": image},
    )

    # Encode all prompts
    prompts = [text1_prompt, text2_prompt, image_prompt]
    outputs = model.encode(prompts, pooling_task="token_embed")

    def get_embeddings(outputs):
        VISION_START_TOKEN_ID, VISION_END_TOKEN_ID = 151652, 151653

        embeddings = []
        for output in outputs:
            if VISION_START_TOKEN_ID in output.prompt_token_ids:
                # Gather only vision tokens
                img_start_pos = torch.where(
                    torch.tensor(output.prompt_token_ids) == VISION_START_TOKEN_ID
                )[0][0]
                img_end_pos = torch.where(
                    torch.tensor(output.prompt_token_ids) == VISION_END_TOKEN_ID
                )[0][0]
                embeddings_tensor = output.outputs.data.detach().clone()[
                    img_start_pos : img_end_pos + 1
                ]
            else:
                # Use all tokens for text-only prompts
                embeddings_tensor = output.outputs.data.detach().clone()

            # Pool and normalize embeddings
            pooled_output = (
                embeddings_tensor.sum(dim=0, dtype=torch.float32)
                / embeddings_tensor.shape[0]
            )
            embeddings.append(torch.nn.functional.normalize(pooled_output, dim=-1))
        return embeddings

    embeddings = get_embeddings(outputs)

    for embedding in embeddings:
        print(embedding.shape)


if __name__ == "__main__":
    main()

multi_vector_retrieval_offline.py

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

from argparse import Namespace

from vllm import LLM, EngineArgs
from vllm.config import PoolerConfig
from vllm.utils.argparse_utils import FlexibleArgumentParser


def parse_args():
    parser = FlexibleArgumentParser()
    parser = EngineArgs.add_cli_args(parser)
    # Set example specific arguments
    parser.set_defaults(
        model="BAAI/bge-m3",
        pooler_config=PoolerConfig(task="token_embed"),
        runner="pooling",
        enforce_eager=True,
    )
    return parser.parse_args()


def main(args: Namespace):
    # Sample prompts.
    prompts = [
        "Hello, my name is",
        "The president of the United States is",
        "The capital of France is",
        "The future of AI is",
    ]

    # Create an LLM.
    # You should pass runner="pooling" for embedding models
    llm = LLM(**vars(args))

    # Generate embedding for each token. The output is a list of PoolingRequestOutput.
    outputs = llm.encode(prompts, pooling_task="token_embed")

    # Print the outputs.
    print("\nGenerated Outputs:\n" + "-" * 60)
    for prompt, output in zip(prompts, outputs):
        multi_vector = output.outputs.data
        print(multi_vector.shape)

    query = "What is the capital of France?"
    documents = [
        "The capital of Brazil is Brasilia.",
        "The capital of France is Paris.",
    ]
    # Generate scores.
    outputs = llm.score(query, documents)
    # Print the outputs.
    print("\nGenerated Outputs:\n" + "-" * 60)
    for document, output in zip(documents, outputs):
        score = output.outputs.score
        print(f"Pair: {[query, document]!r} \nScore: {score}")
        print("-" * 60)


if __name__ == "__main__":
    args = parse_args()
    main(args)

multi_vector_retrieval_online.py

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

"""Example online usage of Pooling API for multi vector retrieval.

Run `vllm serve <model> --runner pooling`
to start up the server in vLLM. e.g.

vllm serve BAAI/bge-m3 --pooler-config.task token_embed
"""

import argparse
import pprint

import requests
import torch


def post_http_request(prompt: dict, api_url: str) -> requests.Response:
    headers = {"User-Agent": "Test Client"}
    response = requests.post(api_url, headers=headers, json=prompt)
    return response


def parse_args():
    parser = argparse.ArgumentParser()
    parser.add_argument("--host", type=str, default="localhost")
    parser.add_argument("--port", type=int, default=8000)
    parser.add_argument("--model", type=str, default="BAAI/bge-m3")

    return parser.parse_args()


def main(args):
    pooling_url = f"http://{args.host}:{args.port}/pooling"
    score_url = f"http://{args.host}:{args.port}/score"
    model_name = args.model

    prompts = [
        "Hello, my name is",
        "The president of the United States is",
        "The capital of France is",
        "The future of AI is",
    ]
    prompt = {"model": model_name, "input": prompts}

    pooling_response = post_http_request(prompt=prompt, api_url=pooling_url)
    for output in pooling_response.json()["data"]:
        multi_vector = torch.tensor(output["data"])
        print(multi_vector.shape)

    queries = "What is the capital of France?"
    documents = [
        "The capital of Brazil is Brasilia.",
        "The capital of France is Paris.",
    ]
    prompt = {"model": model_name, "queries": queries, "documents": documents}
    score_response = post_http_request(prompt=prompt, api_url=score_url)
    print("\nPrompt when queries is string and documents is a list:")
    pprint.pprint(prompt)
    print("\nScore Response:")
    pprint.pprint(score_response.json())


if __name__ == "__main__":
    args = parse_args()
    main(args)

jina_reranker_v3_offline.py

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# ruff: noqa: E501

import torch.nn.functional as F

from vllm import LLM

query = "What are the health benefits of green tea?"
documents = [
    "Green tea contains antioxidants called catechins that may help reduce inflammation and protect cells from damage.",
    "El precio del café ha aumentado un 20% este año debido a problemas en la cadena de suministro.",
    "Studies show that drinking green tea regularly can improve brain function and boost metabolism.",
    "Basketball is one of the most popular sports in the United States.",
    "绿茶富含儿茶素等抗氧化剂,可以降低心脏病风险,还有助于控制体重。",
    "Le thé vert est riche en antioxydants et peut améliorer la fonction cérébrale.",
]


def main():
    # Initialize model
    llm = LLM(
        model="jinaai/jina-reranker-v3",
        runner="pooling",
    )

    # Generate scores.
    outputs = llm.score(query, documents)

    # Print the outputs.
    print("\nGenerated Outputs:\n" + "-" * 60)
    for document, output in zip(documents, outputs):
        score = output.outputs.score
        print(f"Pair: {[query, document]!r} \nScore: {score}")
        print("-" * 60)

    # Generate embeddings.
    # The JinaForRanking model concatenates docs first, then query.
    # Let's stay consistent with this novel design.
    outputs = llm.encode(documents + [query], pooling_task="token_embed")
    embeds = outputs[0].outputs.data.float()

    doc_embeds = embeds[:-1]
    query_embeds = embeds[-1]

    scores = F.cosine_similarity(query_embeds, doc_embeds)

    # Print the outputs.
    print("\nGenerated Outputs:\n" + "-" * 60)
    for document, score in zip(documents, scores):
        print(f"Pair: {[query, document]!r} \nScore: {score}")
        print("-" * 60)


if __name__ == "__main__":
    main()

jina_reranker_v3_online.py

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# ruff: noqa: E501

"""Example online usage of the Jina Reranker v3 score and rerank APIs with a task
instruction.

Run `vllm serve jinaai/jina-reranker-v3 --runner pooling` to start up the
server in vLLM.
"""

import argparse
import json

import requests


def post_http_request(prompt: dict, api_url: str) -> requests.Response:
    headers = {"User-Agent": "Test Client"}
    response = requests.post(api_url, headers=headers, json=prompt)
    return response


def print_response(name: str, prompt: dict, response: requests.Response) -> None:
    print(f"\n{name} request:")
    print(json.dumps(prompt, indent=2))
    print(f"\n{name} response:")
    print(json.dumps(response.json(), indent=2))


def parse_args():
    parser = argparse.ArgumentParser()
    parser.add_argument("--host", type=str, default="localhost")
    parser.add_argument("--port", type=int, default=8000)
    parser.add_argument("--model", type=str, default="jinaai/jina-reranker-v3")
    return parser.parse_args()


def main(args):
    score_url = f"http://{args.host}:{args.port}/score"
    rerank_url = f"http://{args.host}:{args.port}/rerank"
    model_name = args.model

    query = "Which passage is about sports?"
    documents = [
        "Basketball is played by two teams on a court.",
        "Green tea contains antioxidants and may support metabolism.",
    ]
    instruction = "Rank passages about sports higher than passages about nutrition."

    score_prompt = {
        "model": model_name,
        "queries": query,
        "documents": documents,
        "instruction": instruction,
    }
    score_response = post_http_request(prompt=score_prompt, api_url=score_url)
    print_response("Score", score_prompt, score_response)

    rerank_prompt = {
        "model": model_name,
        "query": query,
        "documents": documents,
        "instruction": instruction,
    }
    rerank_response = post_http_request(prompt=rerank_prompt, api_url=rerank_url)
    print_response("Rerank", rerank_prompt, rerank_response)


if __name__ == "__main__":
    args = parse_args()
    main(args)

더 알아보기 (Learn more)