# -*- coding: utf-8 -*-
"""步骤一：专利摘要 -> BERT 最后四层 CLS 平均特征 -> MySQL 入库
依赖：pip install torch transformers sqlalchemy pymysql
模型：先用魔搭下载到本地再改 MODEL_DIR，或直接填本地目录
"""
import random, numpy as np, torch
from transformers import BertTokenizer, BertModel
from sqlalchemy import create_engine, text

# 改成你自己的本地模型目录（魔搭 snapshot_download 后得到的快照路径）
MODEL_DIR = "bert-base-chinese"
SEED = 42
DB_URL = "mysql+pymysql://root@localhost:3306/bert_demo?charset=utf8mb4"  # 按需改

# ── 1. 固定所有随机源 ────────────────────────────────────────
random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
print(f"[1/5] 随机种子已固定: seed={SEED} (random / numpy / torch)")

# ── 2. 从本地目录加载模型（离线，权重冻结） ──────────────────
tok = BertTokenizer.from_pretrained(MODEL_DIR)
model = BertModel.from_pretrained(MODEL_DIR)
model.eval()
print(f"[2/5] 本地模型已加载: {model.config.num_hidden_layers} 层, hidden={model.config.hidden_size}")

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

# ── 4. 特征提取：CLS 的最后四层平均 ─────────────────────────
def extract(texts, batch_size=3):
    feats = []
    for i in range(0, len(texts), batch_size):
        enc = tok(texts[i:i+batch_size], return_tensors="pt",
                  padding="max_length", truncation=True, max_length=128)
        with torch.no_grad():
            out = model(**enc, output_hidden_states=True)
        cls_states = torch.stack(out.hidden_states[-4:], dim=1)   # [B, 4, L, H]
        cls_vec = cls_states[:, :, 0, :]                           # 各层 CLS
        feats.append(cls_vec.mean(dim=1))                          # 四层平均
    return torch.cat(feats).numpy().astype(np.float32)

X = extract([p[2] for p in patents])
print(f"[4/5] 特征提取完成: 形状 {X.shape}, dtype {X.dtype}")
print(f"      P1 前 5 维: {np.round(X[0][:5], 4).tolist()}")

# ── 5. 存入 MySQL（LargeBinary 存字节流） ────────────────────
engine = create_engine(DB_URL)
with engine.begin() as conn:
    conn.execute(text("DROP TABLE IF EXISTS patent_features"))
    conn.execute(text("""CREATE TABLE patent_features (
        pid VARCHAR(8) PRIMARY KEY, field VARCHAR(16),
        dim SMALLINT NOT NULL, feat LONGBLOB NOT NULL, model_tag VARCHAR(64) NOT NULL
    ) CHARACTER SET utf8mb4"""))
    for (pid, field, _), v in zip(patents, X):
        conn.execute(text("INSERT INTO patent_features VALUES (:p,:f,:d,:b,:m)"),
                     {"p": pid, "f": field, "d": int(v.shape[0]),
                      "b": v.tobytes(), "m": "bert-base-chinese@local"})
print(f"[5/5] MySQL 写入完成: 表 patent_features, {len(patents)} 行, 特征共 {X.nbytes:,} 字节")
print("OK")
