257 lines
8.7 KiB
Python
257 lines
8.7 KiB
Python
|
|
"""
|
|||
|
|
训练轻量 Transformer 记忆模型,导出 weights.json(NumPy 推理)。
|
|||
|
|
|
|||
|
|
用法:
|
|||
|
|
cd backend && source venv/bin/activate
|
|||
|
|
pip install torch numpy # 训练仅需本机安装
|
|||
|
|
python -m scripts.train_memory_transformer
|
|||
|
|
python -m scripts.train_memory_transformer --epochs 30 --from-db
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
from __future__ import annotations
|
|||
|
|
|
|||
|
|
import argparse
|
|||
|
|
import math
|
|||
|
|
import sys
|
|||
|
|
from datetime import datetime, timezone
|
|||
|
|
from pathlib import Path
|
|||
|
|
|
|||
|
|
import numpy as np
|
|||
|
|
|
|||
|
|
BACKEND = Path(__file__).resolve().parents[1]
|
|||
|
|
sys.path.insert(0, str(BACKEND))
|
|||
|
|
|
|||
|
|
from database import SessionLocal # noqa: E402
|
|||
|
|
from models import QuizRecord, Word # noqa: E402
|
|||
|
|
from services.memory_transformer.encoding import ( # noqa: E402
|
|||
|
|
FEATURE_DIM,
|
|||
|
|
MAX_SEQ_LEN,
|
|||
|
|
build_sequence_matrix,
|
|||
|
|
)
|
|||
|
|
from services.memory_transformer.network import MiniTransformer, init_random_weights # noqa: E402
|
|||
|
|
from services.memory_transformer.service import ( # noqa: E402
|
|||
|
|
DEFAULT_WEIGHTS_PATH,
|
|||
|
|
WEIGHTS_PATH,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
try:
|
|||
|
|
import torch
|
|||
|
|
import torch.nn as nn
|
|||
|
|
except ImportError:
|
|||
|
|
torch = None # type: ignore
|
|||
|
|
nn = None # type: ignore
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _build_torch_model():
|
|||
|
|
class TorchMiniTransformer(nn.Module):
|
|||
|
|
def __init__(self, d_model=48, n_heads=2, d_ff=96, n_layers=2):
|
|||
|
|
super().__init__()
|
|||
|
|
self.d_model = d_model
|
|||
|
|
self.n_heads = n_heads
|
|||
|
|
self.n_layers = n_layers
|
|||
|
|
self.in_proj = nn.Linear(FEATURE_DIM, d_model)
|
|||
|
|
self.ln_in = nn.LayerNorm(d_model)
|
|||
|
|
self.layers = nn.ModuleList(
|
|||
|
|
[
|
|||
|
|
nn.TransformerEncoderLayer(
|
|||
|
|
d_model=d_model,
|
|||
|
|
nhead=n_heads,
|
|||
|
|
dim_feedforward=d_ff,
|
|||
|
|
batch_first=True,
|
|||
|
|
dropout=0.1,
|
|||
|
|
)
|
|||
|
|
for _ in range(n_layers)
|
|||
|
|
]
|
|||
|
|
)
|
|||
|
|
self.head = nn.Linear(d_model, 1)
|
|||
|
|
|
|||
|
|
def forward(self, x, valid_lens):
|
|||
|
|
h = self.ln_in(self.in_proj(x))
|
|||
|
|
key_padding = torch.zeros(
|
|||
|
|
x.size(0), MAX_SEQ_LEN, dtype=torch.bool, device=x.device
|
|||
|
|
)
|
|||
|
|
for i, vl in enumerate(valid_lens):
|
|||
|
|
if vl < MAX_SEQ_LEN:
|
|||
|
|
key_padding[i, vl:] = True
|
|||
|
|
for layer in self.layers:
|
|||
|
|
h = layer(h, src_key_padding_mask=key_padding)
|
|||
|
|
cls = h[:, 0, :]
|
|||
|
|
return self.head(cls).squeeze(-1)
|
|||
|
|
|
|||
|
|
return TorchMiniTransformer()
|
|||
|
|
|
|||
|
|
|
|||
|
|
def export_torch_to_numpy(torch_model) -> MiniTransformer:
|
|||
|
|
"""将 PyTorch 权重映射到 NumPy MiniTransformer 命名。"""
|
|||
|
|
m = MiniTransformer(
|
|||
|
|
d_model=torch_model.d_model,
|
|||
|
|
n_heads=torch_model.n_heads,
|
|||
|
|
d_ff=96,
|
|||
|
|
n_layers=torch_model.n_layers,
|
|||
|
|
)
|
|||
|
|
w = m.weights
|
|||
|
|
w["in_proj"] = torch_model.in_proj.weight.detach().cpu().numpy().T
|
|||
|
|
w["in_bias"] = torch_model.in_proj.bias.detach().cpu().numpy()
|
|||
|
|
w["ln_in_g"] = torch_model.ln_in.weight.detach().cpu().numpy()
|
|||
|
|
w["ln_in_b"] = torch_model.ln_in.bias.detach().cpu().numpy()
|
|||
|
|
|
|||
|
|
for li, layer in enumerate(torch_model.layers):
|
|||
|
|
attn = layer.self_attn
|
|||
|
|
d = m.d_model
|
|||
|
|
# PyTorch MHA: in_proj_weight stacks Q,K,V
|
|||
|
|
in_w = attn.in_proj_weight.detach().cpu().numpy()
|
|||
|
|
Wq, Wk, Wv = in_w[:d], in_w[d : 2 * d], in_w[2 * d :]
|
|||
|
|
w[f"L{li}.Wq"] = Wq.T
|
|||
|
|
w[f"L{li}.Wk"] = Wk.T
|
|||
|
|
w[f"L{li}.Wv"] = Wv.T
|
|||
|
|
in_b = attn.in_proj_bias.detach().cpu().numpy()
|
|||
|
|
w[f"L{li}.Bq"] = in_b[:d]
|
|||
|
|
w[f"L{li}.Bk"] = in_b[d : 2 * d]
|
|||
|
|
w[f"L{li}.Bv"] = in_b[2 * d :]
|
|||
|
|
w[f"L{li}.Wo"] = attn.out_proj.weight.detach().cpu().numpy().T
|
|||
|
|
w[f"L{li}.Bo"] = attn.out_proj.bias.detach().cpu().numpy()
|
|||
|
|
w[f"L{li}.ln1_g"] = layer.norm1.weight.detach().cpu().numpy()
|
|||
|
|
w[f"L{li}.ln1_b"] = layer.norm1.bias.detach().cpu().numpy()
|
|||
|
|
w[f"L{li}.W1"] = layer.linear1.weight.detach().cpu().numpy().T
|
|||
|
|
w[f"L{li}.b1"] = layer.linear1.bias.detach().cpu().numpy()
|
|||
|
|
w[f"L{li}.W2"] = layer.linear2.weight.detach().cpu().numpy().T
|
|||
|
|
w[f"L{li}.b2"] = layer.linear2.bias.detach().cpu().numpy()
|
|||
|
|
w[f"L{li}.ln2_g"] = layer.norm2.weight.detach().cpu().numpy()
|
|||
|
|
w[f"L{li}.ln2_b"] = layer.norm2.bias.detach().cpu().numpy()
|
|||
|
|
|
|||
|
|
w["head_w"] = torch_model.head.weight.detach().cpu().numpy()[0]
|
|||
|
|
w["head_b"] = torch_model.head.bias.detach().cpu().numpy()[0]
|
|||
|
|
return m
|
|||
|
|
|
|||
|
|
|
|||
|
|
def load_samples_from_db(limit: int = 50000) -> tuple[list, list]:
|
|||
|
|
db = SessionLocal()
|
|||
|
|
try:
|
|||
|
|
records = (
|
|||
|
|
db.query(QuizRecord)
|
|||
|
|
.order_by(QuizRecord.created_at.asc())
|
|||
|
|
.limit(limit)
|
|||
|
|
.all()
|
|||
|
|
)
|
|||
|
|
word_cache: dict[int, Word] = {}
|
|||
|
|
xs, ys = [], []
|
|||
|
|
for r in records:
|
|||
|
|
if r.word_id not in word_cache:
|
|||
|
|
word_cache[r.word_id] = db.query(Word).filter(Word.id == r.word_id).first()
|
|||
|
|
word = word_cache.get(r.word_id)
|
|||
|
|
if not word:
|
|||
|
|
continue
|
|||
|
|
hist = (
|
|||
|
|
db.query(QuizRecord)
|
|||
|
|
.filter(QuizRecord.word_id == r.word_id, QuizRecord.created_at < r.created_at)
|
|||
|
|
.order_by(QuizRecord.created_at.asc())
|
|||
|
|
.all()
|
|||
|
|
)
|
|||
|
|
seq, vl = build_sequence_matrix(word, hist, parse_iso(r.created_at))
|
|||
|
|
xs.append(seq)
|
|||
|
|
ys.append(1.0 if r.is_correct else 0.0)
|
|||
|
|
return xs, ys
|
|||
|
|
finally:
|
|||
|
|
db.close()
|
|||
|
|
|
|||
|
|
|
|||
|
|
def parse_iso(s: str) -> datetime:
|
|||
|
|
s = s.replace("Z", "+00:00")
|
|||
|
|
try:
|
|||
|
|
return datetime.fromisoformat(s)
|
|||
|
|
except ValueError:
|
|||
|
|
return datetime.strptime(s[:19], "%Y-%m-%dT%H:%M:%S").replace(tzinfo=timezone.utc)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def synthetic_samples(n: int = 4000) -> tuple[list, list]:
|
|||
|
|
rng = np.random.default_rng(0)
|
|||
|
|
xs, ys = [], []
|
|||
|
|
for _ in range(n):
|
|||
|
|
vl = int(rng.integers(2, 12))
|
|||
|
|
seq = rng.normal(0, 0.5, (MAX_SEQ_LEN, FEATURE_DIM)).tolist()
|
|||
|
|
seq[0][0] = rng.uniform(0, 1)
|
|||
|
|
seq[0][1] = rng.uniform(0, 1)
|
|||
|
|
last_ok = seq[-1][0] if vl > 1 else 0.5
|
|||
|
|
y = 1.0 if rng.random() < 0.4 + 0.4 * last_ok + 0.1 * seq[0][0] else 0.0
|
|||
|
|
xs.append(seq)
|
|||
|
|
ys.append(y)
|
|||
|
|
return xs, ys
|
|||
|
|
|
|||
|
|
|
|||
|
|
def train_and_export(
|
|||
|
|
epochs: int = 20,
|
|||
|
|
batch_size: int = 64,
|
|||
|
|
from_db: bool = False,
|
|||
|
|
out_path: Path = WEIGHTS_PATH,
|
|||
|
|
) -> None:
|
|||
|
|
if torch is None:
|
|||
|
|
print("请先安装 PyTorch: pip install torch")
|
|||
|
|
print("将写入随机初始化 default_weights 供推理回退…")
|
|||
|
|
m = MiniTransformer()
|
|||
|
|
m.weights = init_random_weights()
|
|||
|
|
DEFAULT_WEIGHTS_PATH.parent.mkdir(parents=True, exist_ok=True)
|
|||
|
|
m.save_json(DEFAULT_WEIGHTS_PATH)
|
|||
|
|
print(f"已保存 {DEFAULT_WEIGHTS_PATH}")
|
|||
|
|
return
|
|||
|
|
|
|||
|
|
xs, ys = load_samples_from_db() if from_db else synthetic_samples()
|
|||
|
|
if len(xs) < 32:
|
|||
|
|
print("样本不足,使用合成数据")
|
|||
|
|
xs, ys = synthetic_samples()
|
|||
|
|
|
|||
|
|
X = torch.tensor(xs, dtype=torch.float32)
|
|||
|
|
y = torch.tensor(ys, dtype=torch.float32)
|
|||
|
|
valid_lens = [min(MAX_SEQ_LEN, sum(1 for row in s if any(v != 0 for v in row))) for s in xs]
|
|||
|
|
for i, s in enumerate(xs):
|
|||
|
|
vl = 1
|
|||
|
|
for j in range(1, MAX_SEQ_LEN):
|
|||
|
|
if any(v != 0 for v in s[j]):
|
|||
|
|
vl = j + 1
|
|||
|
|
valid_lens[i] = max(1, vl)
|
|||
|
|
|
|||
|
|
model = _build_torch_model()
|
|||
|
|
opt = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4)
|
|||
|
|
loss_fn = nn.BCEWithLogitsLoss()
|
|||
|
|
|
|||
|
|
n = len(xs)
|
|||
|
|
for ep in range(epochs):
|
|||
|
|
perm = torch.randperm(n)
|
|||
|
|
total_loss = 0.0
|
|||
|
|
steps = 0
|
|||
|
|
for start in range(0, n, batch_size):
|
|||
|
|
idx = perm[start : start + batch_size]
|
|||
|
|
bx = X[idx]
|
|||
|
|
by = y[idx]
|
|||
|
|
vl_batch = [valid_lens[i] for i in idx.tolist()]
|
|||
|
|
opt.zero_grad()
|
|||
|
|
logits = model(bx, vl_batch)
|
|||
|
|
loss = loss_fn(logits, by)
|
|||
|
|
loss.backward()
|
|||
|
|
opt.step()
|
|||
|
|
total_loss += loss.item()
|
|||
|
|
steps += 1
|
|||
|
|
acc = 0.0
|
|||
|
|
with torch.no_grad():
|
|||
|
|
pred = (torch.sigmoid(model(X, valid_lens)) > 0.5).float()
|
|||
|
|
acc = (pred == y).float().mean().item()
|
|||
|
|
print(f"epoch {ep + 1}/{epochs} loss={total_loss / max(steps, 1):.4f} acc={acc:.3f}")
|
|||
|
|
|
|||
|
|
numpy_model = export_torch_to_numpy(model)
|
|||
|
|
out_path.parent.mkdir(parents=True, exist_ok=True)
|
|||
|
|
numpy_model.save_json(out_path)
|
|||
|
|
print(f"已导出 {out_path}")
|
|||
|
|
|
|||
|
|
|
|||
|
|
def main():
|
|||
|
|
parser = argparse.ArgumentParser()
|
|||
|
|
parser.add_argument("--epochs", type=int, default=20)
|
|||
|
|
parser.add_argument("--from-db", action="store_true")
|
|||
|
|
parser.add_argument("--out", type=str, default="")
|
|||
|
|
args = parser.parse_args()
|
|||
|
|
out = Path(args.out) if args.out else WEIGHTS_PATH
|
|||
|
|
train_and_export(epochs=args.epochs, from_db=args.from_db, out_path=out)
|
|||
|
|
|
|||
|
|
|
|||
|
|
if __name__ == "__main__":
|
|||
|
|
main()
|