1
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

【無料】ローカルLLMでRAGのRerankingをしてみる

1
Posted at

【無料】ローカルLLMでRAGのRerankingをしてみる

動作環境・使用するツールや言語

  • Windows 11 Home 25H2
  • RAM 16.0GB
  • NVIDIA GeForce RTX 3050 Ti Laptop GPU (4 GB)
  • Oracle 26ai free
  • Python

はじめに

AIの学習の手始めとして、無料かつローカルで気軽に試せる環境で動作確認していきます。
初回は青空文庫の小説を要約しようとしてハルシネーションを起こしているところまで確認し、
前回は精度を向上させるためにRAGの精度の評価の仕組みを実装しました。
前回の流れは以下の通りです。

1.テキストを読み込んでチャンクに分割し、Oracleにベクトルを保存

2.チャンクごとに評価用データ(質問と回答)をLLMで作成

3.作成した質問で本文のチャンクをベクトル検索してコサイン類似度が近いものを10行取得、作成元のチャンクが取得できていればOK(Retrieval評価)

4.作成した質問で本文のチャンクをベクトル検索してコサイン類似度が近いものを3行取得、これを根拠資料としてLLMに同じ質問をして回答を生成する

5.4で生成した回答と2で生成した回答(正解)を比較する(Generation評価)
 ・BERTScore
 ・Embedding Similarity
 ・LLM Judge

3と4でベクトル検索していますが、ベクトル検索は意味が近い文字列を取得しているだけです。
これだけでは質問の回答になる文章を取得できるか、というと、ある程度は当たりを引くこともありますが高精度とは言えません。
質問文と似たような意味のある単語が多く出現している文章に回答の根拠がある可能性が高い、というのは我々人間の国語の読解の試験のアプローチでもよくあることです。
国語の試験ではまずざっくり関係がありそうな文章の候補を探して、その中から本当に質問の回答になりそうな文章を吟味します。
この吟味にあたるものがRerankingです。
Rerankingは時間がかかる処理なので、まずベクトル検索で候補を絞ってからその中でRerankingするのが一般的です。単にソートするだけでもいいし、順位の低いものを捨てて絞ってもいいです。
これまでのベクトル検索はBi-Encoderと呼ばれており、質問と本文のチャンクを別々にエンコードしてから類似度を評価しています。
それに対して、RerankingではCross-Encoderと呼ばれるものを使用しており、質問と本文のチャンクを組み合わせてエンコードし、関連性を評価します。組み合わせなのでチャンクが多いと計算量が爆発するため、絞る必要があるわけです。

準備

短編だとチャンク数が少なすぎるので今回は中編として芥川龍之介の「地獄変」としました。

  • テキストのチャンク化
  • 評価データ生成
    までは前回と同様なので割愛します。

Reranking前にSQLでRetrieval評価すると以下のようになりました。
前回と比較して下がっていますが、単純にチャンク数が多くなるとヒット率も下がると思われます。

"HIT1","HIT5","HIT10","HIT20"
"27.3%","59.5%","72.6%","80.9%"

Reranking

情報を格納するテーブルを再作成しておきます。
前回と比較するとBi_EncoderとCross_Encoderのスコアを分けただけです。

DROP TABLE RAG_EVALUATION_CANDIDATE;
CREATE TABLE RAG_EVALUATION_CANDIDATE( 
    EVALUATION_ID NUMBER,
    CHUNK_ID NUMBER,
    CONTENT CLOB,
    BI_SCORE NUMBER,
    CROSS_SCORE NUMBER,
    RERANK NUMBER    
)
;

質問文一つに対して、質問と本文のチャンクをペアにしてエンコードします。専用のEncoderが必要です。
結果は関連性を示す数値として出力されるので、これを降順でソートすればRerankになります。

import oracledb
from sentence_transformers import SentenceTransformer
from sentence_transformers import CrossEncoder
import json
if __name__ == "__main__":
    # Oracle接続
    conn = oracledb.connect(
        user="rag",
        password="rag",
        dsn="127.0.0.1:1521/FREEPDB1"
    )
    # Embeddingモデル_Bi-Encoder
    bi_model = SentenceTransformer('intfloat/multilingual-e5-base')
    # Embeddingモデル_Cross_Encoder
    cross_model = CrossEncoder('hotchpotch/japanese-reranker-cross-encoder-xsmall-v1')
    select_cursor = conn.cursor()
    select_cursor2 = conn.cursor()
    insert_cursor = conn.cursor()

    try:
        rows = select_cursor.execute(
        """
        SELECT
            evaluation_id,
            question,
            source_chunk_id        
        FROM rag_evaluation
        """
        )
        # 質問1つごとに処理
        for evaluation_id, question, source_chunk_id in rows:
            # 質問をembeddingしてリストにする
            query_embedding = bi_model.encode(
                f"query: {question}"
            ).tolist()
            # ベクトル検索で質問文に近いcontentを検索する
            select_sql2 = """
            SELECT
                id,
                DBMS_LOB.SUBSTR(content,4000,1),
                VECTOR_DISTANCE(
                    embedding,
                    TO_VECTOR(:vec),
                    COSINE
                ) score
            FROM text_documents
            ORDER BY score
            FETCH FIRST 20 ROWS ONLY
            """
            # json文字列にして実行
            params = [
                json.dumps(query_embedding)
            ]
            select_cursor2.execute(
                select_sql2,
                params
            )

            # 質問1つに対してtop20をReranking
            pairs = []
            chunks = []
            for chunk_id, content, bi_score in select_cursor2:
                # 質問と本文のチャンクのペアを作成する                
                pairs.append((question,content))
                # チャンクの情報
                chunks.append((chunk_id,content,bi_score))
            # ペアの関連度を数値で出力する
            scores = cross_model.predict(pairs)
            # 関連度で降順にソートし、チャンクの情報も追記しておく
            # [(chunk7,0.98),
            # (chunk2,0.96),
            # (chunk15,0.93),
            # ...]
            # のような形式になる
            ranked = sorted(zip(chunks, scores), key=lambda x: x[1], reverse=True)
            # 降順にソートしたのでReRankできたのでrank1から順番にinsertしていく
            for rank,((chunk_id, content, bi_score),cross_score) in enumerate(ranked,1):
                if chunk_id == source_chunk_id:
                    print(
                        evaluation_id,
                        source_chunk_id,
                        rank
                    )
                insert_cursor.execute(
                    """
                    INSERT INTO rag_evaluation_candidate
                    (
                        evaluation_id,
                        chunk_id,
                        content,
                        bi_score,
                        cross_score,
                        rerank
                    )
                    VALUES
                    (
                        :1,
                        :2,
                        :3,
                        :4,
                        :5,
                        :6
                    )
                    """,
                    [
                        evaluation_id,
                        chunk_id,
                        content,
                        float(bi_score),
                        float(cross_score),
                        rank
                    ]
                )
        conn.commit() 
    except Exception as e:
        print(f"エラー: {e}")
    finally:
        select_cursor.close()
        select_cursor2.close()
        insert_cursor.close()
        conn.close()

Retrieval評価

では再度SQLでRetrieval評価してみましょう。

SELECT
    TRUNC(HIT1/TOTAL*100,1) || '%' AS HIT1,
    TRUNC(HIT5/TOTAL*100,1) || '%' AS HIT5,
    TRUNC(HIT10/TOTAL*100,1) || '%' AS HIT10,
    TRUNC(HIT20/TOTAL*100,1) || '%' AS HIT20    
FROM (
SELECT
    COUNT(*) total,
    SUM(
        CASE
            WHEN c.rerank <= 1 THEN 1
            ELSE 0
        END
    ) hit1,
    SUM(
        CASE
            WHEN c.rerank <= 5 THEN 1
            ELSE 0
        END
    ) hit5,
    SUM(
        CASE
            WHEN c.rerank <= 10 THEN 1
            ELSE 0
        END
    ) hit10,
    SUM(
        CASE
            WHEN c.rerank <= 20 THEN 1
            ELSE 0
        END
    ) hit20    
FROM rag_evaluation e
LEFT JOIN rag_evaluation_candidate c
ON e.evaluation_id=c.evaluation_id
AND e.source_chunk_id=c.chunk_id
);

出力は下記の通りでした。

"HIT1","HIT5","HIT10","HIT20"
"57.1%","76.1%","77.3%","80.9%"

Reranking前が

"HIT1","HIT5","HIT10","HIT20"
"27.3%","59.5%","72.6%","80.9%"

だったので、HIT1,HIT5,HIT10の精度が向上していることがわかります。

Generation評価

では次にGeneration評価をしてみましょう。
Reranking前は以下の通りです。

BERT F1 AVG: 0.63920975
EMB AVG: 0.8459409
                 bert_f1  emb_similarity    scores
bert_f1         1.000000        0.759018 -0.132907
emb_similarity  0.759018        1.000000 -0.104807
scores         -0.132907       -0.104807  1.000000
      bert_f1  scores
0    0.715700       3
1    0.546182       5
2    0.648605       2
3    0.661391       1
4    0.646020       1
..        ...     ...
163  0.543624       2
164  0.660829       1
165  0.642491       3
166  0.702824       1
167  0.611802       5

Rerankしたものをベースにして回答を生成させます。

from ollama import chat
import oracledb

if __name__ == "__main__":
    # Oracle接続
    conn = oracledb.connect(
        user="rag",
        password="rag",
        dsn="127.0.0.1:1521/FREEPDB1"
    )
    # 質問の取得
    select_cursor = conn.cursor()
    update_cursor = conn.cursor()
    select_sql = """
    SELECT
        evaluation_id,
        question        
    FROM rag_evaluation
    """
    rows = select_cursor.execute(select_sql)
    try:
        chunks = []
        i = 0
        # 質問1つごとに処理
        for evaluation_id, question in rows:
            select_cursor2 = conn.cursor()
            select_sql2 = """
            SELECT
                chunk_id,
                content
            FROM rag_evaluation_candidate
            WHERE evaluation_id = :1
            ORDER BY rerank 
            FETCH FIRST 3 ROWS ONLY
            """
            select_cursor2.execute(select_sql2,[evaluation_id])
            res = select_cursor2.fetchall()
            # select結果をchunksに格納する
            for r in res:
                chunks.append(r[1].read())

            prompt = f"""
            以下の文書を根拠に質問に回答してください。

            文書1:
            {chunks[0]}
            文書2:
            {chunks[1]}
            文書3:
            {chunks[2]}

            質問:
            {question}

            回答:
            """

            response = chat(
                model="gemma3:4b",
                messages=[
                    {
                        "role": "user",
                        "content": prompt
                    }
                ]
            )
            generated_answer = response["message"]["content"]
            print("====================")
            print(evaluation_id)
            print(question)
            print(generated_answer)
            print("====================")
            update_cursor.execute(
            """
            UPDATE RAG_EVALUATION
            SET GENERATED_ANSWER = :1
            WHERE EVALUATION_ID = :2
            """,
            [
                generated_answer,
                evaluation_id
            ]
            )
            i = i + 1
            if i % 10 == 0:
                conn.commit() 
        conn.commit()
    except Exception as e:
        print(f"エラー: {e}")
    finally:
        select_cursor.close()
        select_cursor2.close()
        update_cursor.close()
        conn.close()

スコアの取得は前回と同様です。

BERT F1 AVG: 0.6406883
EMB AVG: 0.8460993
                 bert_f1  emb_similarity    scores
bert_f1         1.000000        0.802711 -0.067856
emb_similarity  0.802711        1.000000  0.051284
scores         -0.067856        0.051284  1.000000
      bert_f1  scores
0    0.705810       4
1    0.541913       5
2    0.666397       3
3    0.659048       1
4    0.674465       1
..        ...     ...
163  0.585741       1
164  0.691808       2
165  0.661878       2
166  0.721374       2
167  0.625732       5

結果はほとんど変わりません。
Retrieval評価に比べてGeneration評価があまり改善しないのはよくあることのようです。
例えば
1位:A
2位:B
3位:C

1位:C
2位:A
3位:B
にRerankingしたとしても、LLMに渡す根拠資料3つは変わっていないので回答はほぼ変わりません。

1
0
0

Register as a new user and use Qiita more conveniently

  1. You get articles that match your needs
  2. You can efficiently read back useful information
  3. You can use dark theme
What you can do with signing up
1
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?