CG-IR / src /tracing /bkt.py
Lifunn's picture
Upload 10 files
9342074 verified
Raw
History Blame Contribute Delete
2.99 kB
from dataclasses import dataclass, field
from typing import Dict, List, Optional
@dataclass
class BKTParams:
p_init: float = 0.10
p_transit: float = 0.30
p_slip: float = 0.10
p_guess: float = 0.20
@dataclass
class ObservationRecord:
concept: str
correct: bool
mastery_before: float
mastery_after: float
bloom_level: int = 1
question_id: str = ""
class BayesianKnowledgeTracing:
def __init__(self, default_params: Optional[BKTParams] = None):
self.default_params = default_params or BKTParams()
self.concept_params: Dict[str, BKTParams] = {}
self._mastery: Dict[str, float] = {}
self.history: List[ObservationRecord] = []
def get_mastery(self, concept: str) -> float:
return self._mastery.get(
concept,
self.concept_params.get(concept, self.default_params).p_init
)
def set_params(self, concept: str, params: BKTParams):
self.concept_params[concept] = params
def update(
self,
concept: str,
correct: bool,
bloom_level: int = 1,
question_id: str = "",
) -> float:
params = self.concept_params.get(concept, self.default_params)
p_l = self.get_mastery(concept)
if correct:
p_obs_l1 = 1.0 - params.p_slip
p_obs_l0 = params.p_guess
else:
p_obs_l1 = params.p_slip
p_obs_l0 = 1.0 - params.p_guess
p_obs = p_obs_l1 * p_l + p_obs_l0 * (1.0 - p_l)
if p_obs < 1e-12:
p_obs = 1e-12
p_l_given_obs = (p_obs_l1 * p_l) / p_obs
p_l_next = p_l_given_obs + (1.0 - p_l_given_obs) * params.p_transit
p_l_next = min(max(p_l_next, 0.0), 1.0)
self._mastery[concept] = p_l_next
self.history.append(ObservationRecord(
concept=concept,
correct=correct,
mastery_before=p_l,
mastery_after=p_l_next,
bloom_level=bloom_level,
question_id=question_id,
))
return p_l_next
def get_topic_mastery(self, concepts: List[str]) -> float:
if not concepts:
return 0.0
return sum(self.get_mastery(c) for c in concepts) / len(concepts)
def all_masteries(self) -> Dict[str, float]:
return dict(self._mastery)
def concept_history(self, concept: str) -> List[ObservationRecord]:
return [r for r in self.history if r.concept == concept]
def summary(self) -> str:
lines = ["BKT Mastery Summary:"]
for concept, mastery in sorted(self._mastery.items(), key=lambda x: -x[1]):
hist = self.concept_history(concept)
attempts = len(hist)
correct = sum(1 for r in hist if r.correct)
lines.append(
f" {concept:<30} {mastery:.3f} ({correct}/{attempts} correct)"
)
return "\n".join(lines)