# -*- coding: utf-8 -*-
"""步骤三：下游应用——从 MySQL 读回特征，计算文本间余弦相似度"""
import numpy as np
from sqlalchemy import create_engine, text

DB_URL = "mysql+pymysql://root@localhost:3306/bert_demo?charset=utf8mb4"

engine = create_engine(DB_URL)
with engine.connect() as conn:
    rows = conn.execute(text("SELECT pid, field, feat FROM patent_features ORDER BY pid")).fetchall()

pids  = [r[0] for r in rows]
names = [r[1] for r in rows]
X = np.stack([np.frombuffer(r[2], dtype=np.float32) for r in rows])
print(f"从 MySQL 读回特征: {X.shape}  (来源表 patent_features)")

Xn = X / np.linalg.norm(X, axis=1, keepdims=True)   # L2 归一化后内积 = 余弦相似度
S = Xn @ Xn.T

print("      " + "".join(f"{p:>7}" for p in pids))
for i, pid in enumerate(pids):
    print(f"{pid:>4}  " + "".join(f"{S[i,j]:7.3f}" for j in range(len(pids))))

iu = np.triu_indices(len(pids), k=1)
order = np.argsort(S[iu])[::-1]
print("\n最相似的三对:")
for k in order[:3]:
    i, j = iu[0][k], iu[1][k]
    print(f"  {pids[i]}({names[i]}) - {pids[j]}({names[j]}): {S[i,j]:.3f}")
i, j = iu[0][order[-1]], iu[1][order[-1]]
print(f"最不相似的一对: {pids[i]}({names[i]}) - {pids[j]}({names[j]}): {S[i,j]:.3f}")
print("OK")
