Files

107 lines
2.9 KiB
Python
Raw Permalink Normal View History

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