Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
104 changes: 66 additions & 38 deletions scratch/verify_real_data.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
import os
import sys
import time
import base64
Expand All @@ -11,8 +10,7 @@
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))

import torch
from app.models import get_model, VisualizedBGEEmbeddingModel
from app.config import EMBEDDING_MODELS, RERANK_MODELS
from app.models import get_model


def cosine_similarity(v1: list[float], v2: list[float]) -> float:
Expand All @@ -24,12 +22,16 @@ def cosine_similarity(v1: list[float], v2: list[float]) -> float:
return dot / (norm1 * norm2)


def create_sample_image(text: str = "システム構成図", color: tuple = (70, 130, 180)) -> Image.Image:
def create_sample_image(
text: str = "システム構成図", color: tuple = (70, 130, 180)
) -> Image.Image:
"""実データのテスト用画像を動的に生成"""
img = Image.new("RGB", (256, 256), color=(240, 244, 248))
draw = ImageDraw.Draw(img)
draw.rectangle([20, 20, 236, 100], fill=color, outline=(30, 60, 90), width=2)
draw.rectangle([40, 140, 216, 220], fill=(220, 230, 242), outline=(30, 60, 90), width=2)
draw.rectangle(
[40, 140, 216, 220], fill=(220, 230, 242), outline=(30, 60, 90), width=2
)
draw.line([(128, 100), (128, 140)], fill=(30, 60, 90), width=3)
return img

Expand All @@ -38,7 +40,7 @@ def run_text_embedding_verification():
print("\n========================================================")
print("1. テキスト埋め込み実データ検証 (RURI, BGE-M3)")
print("========================================================")

test_queries = [
"検索クエリ: 日本の首都はどこですか?",
"検索ドキュメント: 日本の首都は東京都であり、政治・経済・文化の中心地です。",
Expand All @@ -55,26 +57,32 @@ def run_text_embedding_verification():
t0 = time.perf_counter()
embeddings = model.encode(test_queries, normalize_embeddings=True)
infer_time = time.perf_counter() - t0

dim = len(embeddings[0])
sim_relevant = cosine_similarity(embeddings[0].tolist(), embeddings[1].tolist())
sim_irrelevant = cosine_similarity(embeddings[0].tolist(), embeddings[2].tolist())

sim_irrelevant = cosine_similarity(
embeddings[0].tolist(), embeddings[2].tolist()
)

print(f" ✓ 埋め込み次元数: {dim}")
print(f" ✓ 推論時間 (3文): {infer_time*1000:.1f}ms")
print(f" ✓ 推論時間 (3文): {infer_time * 1000:.1f}ms")
print(f" ✓ 関連ドキュメントとの類似度: {sim_relevant:.4f}")
print(f" ✓ 無関係ドキュメントとの類似度: {sim_irrelevant:.4f}")

assert dim > 0, "次元数が不正です"
assert sim_relevant > sim_irrelevant, f"関連文書の類似度({sim_relevant})が無関係文書({sim_irrelevant})を下回っています"
print(f" 🎯 判定: 合格 (関連度判定正常: {sim_relevant:.4f} > {sim_irrelevant:.4f})")
assert sim_relevant > sim_irrelevant, (
f"関連文書の類似度({sim_relevant})が無関係文書({sim_irrelevant})を下回っています"
)
print(
f" 🎯 判定: 合格 (関連度判定正常: {sim_relevant:.4f} > {sim_irrelevant:.4f})"
)


def run_multimodal_verification():
print("\n========================================================")
print("2. マルチモーダル実データ検証 (bge-visualized-m3)")
print("========================================================")

t0 = time.perf_counter()
model = get_model("bge-visualized-m3", device="cpu")
print(f" ✓ ロード完了 ({time.perf_counter() - t0:.2f}秒)")
Expand All @@ -97,30 +105,38 @@ def run_multimodal_verification():
t0 = time.perf_counter()
embeddings = model.encode_multimodal(items)
infer_time = time.perf_counter() - t0

dim = len(embeddings[0])
print(f" ✓ マルチモーダル埋め込み次元数: {dim}")
print(f" ✓ 推論時間 ({len(items)}アイテム): {infer_time*1000:.1f}ms")
print(f" ✓ 推論時間 ({len(items)}アイテム): {infer_time * 1000:.1f}ms")

# 類似度評価
# システム構成図(画像+テキスト) と クラウドインフラ構成図(テキスト)
sim_diagram = cosine_similarity(embeddings[0], embeddings[2])
# システム構成図(画像+テキスト) と 青空と緑の草原(テキスト)
sim_mismatch = cosine_similarity(embeddings[0], embeddings[3])

print(f" ✓ アーキテクチャ図(画像+文) vs クラウドインフラ(文) 類似度: {sim_diagram:.4f}")
print(
f" ✓ アーキテクチャ図(画像+文) vs クラウドインフラ(文) 類似度: {sim_diagram:.4f}"
)
print(f" ✓ アーキテクチャ図(画像+文) vs 草原風景(文) 類似度: {sim_mismatch:.4f}")

assert dim == 1024, f"bge-visualized-m3 の次元数は 1024 である必要があります (実際: {dim})"
assert sim_diagram > sim_mismatch, f"画像-テキスト間のセマンティック類似度が期待を満たしていません ({sim_diagram} vs {sim_mismatch})"
print(f" 🎯 判定: 合格 (マルチモーダル類似度正常: {sim_diagram:.4f} > {sim_mismatch:.4f})")

assert dim == 1024, (
f"bge-visualized-m3 の次元数は 1024 である必要があります (実際: {dim})"
)
assert sim_diagram > sim_mismatch, (
f"画像-テキスト間のセマンティック類似度が期待を満たしていません ({sim_diagram} vs {sim_mismatch})"
)
print(
f" 🎯 判定: 合格 (マルチモーダル類似度正常: {sim_diagram:.4f} > {sim_mismatch:.4f})"
)


def run_reranker_verification():
print("\n========================================================")
print("3. リランカー実データ検証 (ruri-v3-reranker-310m)")
print("========================================================")

t0 = time.perf_counter()
model = get_model("cl-nagoya/ruri-v3-reranker-310m", device="cpu")
print(f" ✓ ロード完了 ({time.perf_counter() - t0:.2f}秒)")
Expand All @@ -131,26 +147,30 @@ def run_reranker_verification():
"日本の温泉地ランキングでは、草津温泉や別府温泉、有馬温泉などが上位に選ばれています。",
"ニューラルネットワークの汎化性能向上のため、学習データのバリデーション分割やクロスバリデーションが推奨されます。",
]

pairs = [[query, p] for p in passages]
t0 = time.perf_counter()
scores = model.predict(pairs)
infer_time = time.perf_counter() - t0
print(f" ✓ 推論時間 ({len(pairs)}ペア): {infer_time*1000:.1f}ms")

print(f" ✓ 推論時間 ({len(pairs)}ペア): {infer_time * 1000:.1f}ms")
for i, (p, score) in enumerate(zip(passages, scores)):
print(f" [{i+1}] スコア: {score:+.4f} | 内容: {p[:35]}...")
print(f" [{i + 1}] スコア: {score:+.4f} | 内容: {p[:35]}...")

assert scores[0] > scores[1], "過学習対策ドキュメントのスコアが温泉ドキュメントを下回っています"
assert scores[2] > scores[1], "汎化性能ドキュメントのスコアが温泉ドキュメントを下回っています"
assert scores[0] > scores[1], (
"過学習対策ドキュメントのスコアが温泉ドキュメントを下回っています"
)
assert scores[2] > scores[1], (
"汎化性能ドキュメントのスコアが温泉ドキュメントを下回っています"
)
print(" 🎯 判定: 合格 (リランキング順位スコア正常)")


def run_device_switching_verification():
print("\n========================================================")
print("4. デバイス切り替え・フォールバック検証 (CPU / CUDA)")
print("========================================================")

cuda_available = torch.cuda.is_available()
print(f" 現在のCUDA利用可能性: {cuda_available}")

Expand Down Expand Up @@ -200,7 +220,9 @@ def run_fastapi_endpoints_real_verification():
data = res.json()
assert len(data["data"]) == 2
assert len(data["data"][0]["embedding"]) > 0
print(f" ✓ ステータス 200, 次元数: {len(data['data'][0]['embedding'])}, Usage: {data['usage']}")
print(
f" ✓ ステータス 200, 次元数: {len(data['data'][0]['embedding'])}, Usage: {data['usage']}"
)

# 2. /v1/embeddings (マルチモーダル Base64)
print("\n [Endpoint 2] POST /v1/embeddings (マルチモーダル Base64画像)")
Expand All @@ -223,7 +245,9 @@ def run_fastapi_endpoints_real_verification():
data = res.json()
assert len(data["data"]) == 1
assert len(data["data"][0]["embedding"]) == 1024
print(f" ✓ ステータス 200, マルチモーダル埋め込み次元数: {len(data['data'][0]['embedding'])}")
print(
f" ✓ ステータス 200, マルチモーダル埋め込み次元数: {len(data['data'][0]['embedding'])}"
)

# 3. /v1/rerank (リランキング)
print("\n [Endpoint 3] POST /v1/rerank (テキストリランキング)")
Expand All @@ -247,22 +271,26 @@ def run_fastapi_endpoints_real_verification():
assert len(results) == 2
print(f" ✓ ステータス 200, Top-{len(results)} 返却:")
for r in results:
print(f" - Doc {r['document']}: score={r['score']:+.4f} | text={r.get('text', '')[:35]}...")
print(
f" - Doc {r['document']}: score={r['score']:+.4f} | text={r.get('text', '')[:35]}..."
)
assert results[0]["document"] in (0, 2)
print(" 🎯 判定: 合格 (全APIエンドポイント実データ推論正常)")


if __name__ == "__main__":
t_start = time.perf_counter()
print("🚀 実データ・デバイス動作検証テストを開始します")

run_text_embedding_verification()
run_multimodal_verification()
run_reranker_verification()
run_device_switching_verification()
run_fastapi_endpoints_real_verification()

total_sec = time.perf_counter() - t_start
print(f"\n========================================================")
print(f"🎉 全ての実データ・デバイス検証テストに合格しました! (総所要時間: {total_sec:.2f}秒)")
print(f"========================================================")
print("\n========================================================")
print(
f"🎉 全ての実データ・デバイス検証テストに合格しました! (総所要時間: {total_sec:.2f}秒)"
)
print("========================================================")
Loading
Loading