"""将 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