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>
This commit is contained in:
@@ -0,0 +1,106 @@
|
||||
"""将 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
|
||||
Reference in New Issue
Block a user