c1a6a105ef
Introduce Q/K/V memory dialogue with coach APIs, a lightweight NumPy transformer for per-word forgetting prediction, and a floating training menu linking daily quiz, spell, and coach flows. Co-authored-by: Cursor <cursoragent@cursor.com>
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()
|