# -*- coding: utf-8 -*-
"""步骤二：可复现性验证——两次独立提取逐位比对 + 与 MySQL 读回比对"""
import gc, numpy as np, torch
from transformers import BertTokenizer, BertModel
from sqlalchemy import create_engine, text

MODEL_DIR = "bert-base-chinese"
DB_URL = "mysql+pymysql://root@localhost:3306/bert_demo?charset=utf8mb4"

# 与 extract_store.py 保持相同的示例文本
ABSTRACTS = [
    "一种锂离子电池正极材料及其制备方法，通过镍钴锰三元前驱体与锂源混合烧结，掺杂铝元素改善晶体结构稳定性，提升循环寿命与热稳定性，适用于动力电池领域。",
    "一种城市污水深度处理装置，包括厌氧池、缺氧池与膜生物反应器串联结构，通过间歇曝气控制溶解氧浓度，实现脱氮除磷与污泥减量，出水达到一级A排放标准。",
    "基于深度学习的工业缺陷检测方法，采集产品表面图像后输入卷积神经网络，经数据增强与迁移学习训练的分类模型输出缺陷类别与位置，检测准确率显著提升。",
    "一种可降解食品包装膜及其制备工艺，以聚乳酸与纤维素纳米晶为基体，添加植物源增塑剂共混吹塑成型，薄膜兼具透气调节功能与力学强度，废弃后可完全降解。",
    "一种光伏逆变器最大功率点跟踪控制方法，采用改进扰动观察法结合自适应步长，在局部遮阴条件下通过全局扫描定位最大功率点，提高发电效率并抑制功率振荡。",
    "一种输电线路无人机自主巡检系统，搭载双光相机与激光雷达，通过航线规划算法与图像拼接识别绝缘子缺陷，巡检数据经边缘计算单元压缩后回传监控平台。",
]

def run_once():
    torch.manual_seed(42); np.random.seed(42)
    tok = BertTokenizer.from_pretrained(MODEL_DIR)
    model = BertModel.from_pretrained(MODEL_DIR); model.eval()
    feats = []
    for i in range(0, len(ABSTRACTS), 3):
        enc = tok(ABSTRACTS[i:i+3], return_tensors="pt",
                  padding="max_length", truncation=True, max_length=128)
        with torch.no_grad():
            out = model(**enc, output_hidden_states=True)
        cls = torch.stack(out.hidden_states[-4:], dim=1)[:, :, 0, :]
        feats.append(cls.mean(dim=1).numpy().astype(np.float32))
    del model; gc.collect()
    return np.concatenate(feats)

X1 = run_once(); print(f"第 1 次运行: 提取完成 {X1.shape}")
X2 = run_once(); print(f"第 2 次运行: 重新加载模型后提取完成 {X2.shape}")
diff = float(np.abs(X1 - X2).max())
print(f"两次独立运行最大绝对差 = {diff:.10f}  {'<-- 完全一致' if diff == 0 else '<-- 存在差异!'}")

engine = create_engine(DB_URL)
with engine.connect() as conn:
    rows = conn.execute(text("SELECT pid, feat, model_tag FROM patent_features ORDER BY pid")).fetchall()
same = 0
for (pid, blob, tag), v in zip(rows, X1):
    same += int(np.array_equal(np.frombuffer(blob, dtype=np.float32), v))
print(f"MySQL 读回比对: {same}/{len(rows)} 篇字节级一致 (model_tag={rows[0][2]})")
print("OK" if diff == 0 and same == len(rows) else "FAIL")
