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>
107 lines
2.9 KiB
Python
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
|