Files
John c1a6a105ef Add memory coach, transformer recall model, and training FAB.
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>
2026-06-04 18:11:49 -07:00

107 lines
2.9 KiB
Python

"""将 QuizRecord 序列编码为 Transformer 输入特征。"""
import math
from datetime import datetime, timezone
from typing import Optional
from models import QuizRecord, Word
FEATURE_DIM = 16
MAX_SEQ_LEN = 32
QUESTION_TYPES = ("en_to_zh", "zh_to_en", "spell", "memory_coach")
STATUS_ORDER = ("new", "learning", "mastered", "weak")
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 word_en(w: Word) -> str:
return w.target_text if w.source_lang == "zh" else w.source_text
def _log_hours(delta_hours: float) -> float:
return math.log1p(max(0.0, delta_hours)) / math.log1p(24 * 30)
def _one_hot(index: int, size: int) -> list[float]:
v = [0.0] * size
if 0 <= index < size:
v[index] = 1.0
return v
def _pad_features(feats: list[float]) -> list[float]:
row = feats[:FEATURE_DIM]
while len(row) < FEATURE_DIM:
row.append(0.0)
return row
def word_static_features(word: Word) -> list[float]:
en = word_en(word)
status_idx = STATUS_ORDER.index(word.status) if word.status in STATUS_ORDER else 1
feats = [
word.mastery_score / 100.0,
min(word.consecutive_correct_count, 10) / 10.0,
min(word.correct_count, 50) / 50.0,
min(word.wrong_count, 50) / 50.0,
min(len(en), 24) / 24.0,
]
feats.extend(_one_hot(status_idx, len(STATUS_ORDER)))
return _pad_features(feats)
def event_features(
record: QuizRecord,
prev_at: Optional[datetime],
at: datetime,
) -> list[float]:
q_idx = (
QUESTION_TYPES.index(record.question_type)
if record.question_type in QUESTION_TYPES
else 0
)
if prev_at is None:
delta_h = 0.0
else:
delta_h = max(0.0, (at - prev_at).total_seconds() / 3600.0)
dur = min(record.duration_seconds or 0, 600) / 600.0
feats = [
1.0 if record.is_correct else 0.0,
_log_hours(delta_h),
dur,
]
feats.extend(_one_hot(q_idx, len(QUESTION_TYPES)))
return _pad_features(feats)
def build_sequence_matrix(
word: Word,
records: list[QuizRecord],
now: Optional[datetime] = None,
) -> tuple[list[list[float]], int]:
"""
返回 (seq_features, valid_len)。
第 0 位为词项 CLS(静态),其后为按时间排序的练习事件(最多 MAX_SEQ_LEN-1)。
"""
now = now or datetime.now(timezone.utc)
ordered = sorted(records, key=lambda r: r.created_at)
seq: list[list[float]] = [word_static_features(word)]
prev_at: Optional[datetime] = parse_iso(word.created_at)
for r in ordered[-(MAX_SEQ_LEN - 1) :]:
at = parse_iso(r.created_at)
seq.append(event_features(r, prev_at, at))
prev_at = at
valid_len = len(seq)
while len(seq) < MAX_SEQ_LEN:
seq.append([0.0] * FEATURE_DIM)
return seq, valid_len