| 416 | |
| 417 | @dataclass |
| 418 | class UserDiscourseState: |
| 419 | profile_text: str |
| 420 | profile_vector: Dict[str, float] = field(default_factory=dict) |
| 421 | selected_sums: Dict[str, Dict[str, float]] = field(default_factory=lambda: defaultdict(dict)) |
| 422 | selected_facets: List[Dict[str, str]] = field(default_factory=list) |
| 423 | selected_count: int = 0 |
| 424 | |
| 425 | def feedback_texts(self) -> List[str]: |
| 426 | texts: List[str] = [] |
| 427 | for facets in self.selected_facets: |
| 428 | texts.extend(facets.values()) |
| 429 | return texts |
| 430 | |
| 431 | def prepare_vectors(self, idf: Dict[str, float]) -> None: |
| 432 | self.profile_vector = vectorize_text(self.profile_text, idf) |
| 433 | self.selected_sums = defaultdict(dict) |
| 434 | for facets in self.selected_facets: |
| 435 | for facet, text in facets.items(): |
| 436 | weight = FACET_WEIGHTS.get(facet, 0.5) |
| 437 | add_weighted_vector(self.selected_sums[facet], vectorize_text(text, idf), weight) |
| 438 | |
| 439 | def selected_centroid(self, facet: str) -> Dict[str, float]: |
| 440 | return normalize_vector(self.selected_sums.get(facet, {})) |
| 441 | |
| 442 | def update_selected(self, row: Dict[str, Any]) -> None: |
| 443 | self.selected_facets.append(extract_discourse_facets(row)) |
| 444 | self.selected_count += 1 |
| 445 | |
| 446 | |
| 447 | def facet_similarity_score(candidate_facets: Dict[str, str], idf: Dict[str, float], state: UserDiscourseState) -> float: |
no outgoing calls
no test coverage detected