Process a single training batch. Args: batch: List of data items with 'question' and 'standard_answer' update_knowledge_base: Whether to update the knowledge base subject_retrieval_top_k: Optional override for subject retrieval top_k (def
(
self,
batch: List[Dict[str, Any]],
update_knowledge_base: bool = True,
subject_retrieval_top_k: Optional[int] = None,
subject_retrieval_threshold: Optional[float] = None,
is_retrieval_subject: bool = True,
)
| 390 | return guidance.strip() |
| 391 | |
| 392 | def _process_batch( |
| 393 | self, |
| 394 | batch: List[Dict[str, Any]], |
| 395 | update_knowledge_base: bool = True, |
| 396 | subject_retrieval_top_k: Optional[int] = None, |
| 397 | subject_retrieval_threshold: Optional[float] = None, |
| 398 | is_retrieval_subject: bool = True, |
| 399 | ) -> Optional[Dict[str, Any]]: |
| 400 | """ |
| 401 | Process a single training batch. |
| 402 | |
| 403 | Args: |
| 404 | batch: List of data items with 'question' and 'standard_answer' |
| 405 | update_knowledge_base: Whether to update the knowledge base |
| 406 | subject_retrieval_top_k: Optional override for subject retrieval top_k (default: self.subject_retrieval_top_k) |
| 407 | subject_retrieval_threshold: Optional override for subject retrieval threshold (default: self.subject_retrieval_threshold) |
| 408 | is_retrieval_subject: Whether to retrieve guidance based on subjects or questions |
| 409 | Returns: |
| 410 | Dictionary with batch metrics, or None if no new guidance was generated |
| 411 | """ |
| 412 | questions = [item["question"] for item in batch] |
| 413 | standard_answers = [item["answer"] for item in batch] |
| 414 | |
| 415 | # Step 1: Classify subjects for each question |
| 416 | subjects = self.llm_client.classify_subjects( |
| 417 | questions=questions |
| 418 | ) |
| 419 | if subjects is None: |
| 420 | logger.warning("All cases in batch failed, skipping batch") |
| 421 | return None |
| 422 | |
| 423 | # Step 2: Retrieve relevant guidance for each question (RAG-based) |
| 424 | # Use subject for retrieval since we already have subjects from classification |
| 425 | baseline_system_prompts = self._retrieve_guidance_for_batch( |
| 426 | subjects=subjects if is_retrieval_subject else questions, |
| 427 | ) |
| 428 | |
| 429 | # Step 3: Evaluate with retrieved guidance (baseline) |
| 430 | baseline_responses = self.llm_client.batch_generate( |
| 431 | prompts=questions, |
| 432 | system_prompt=baseline_system_prompts, # Each question gets its own relevant guidance |
| 433 | ) |
| 434 | |
| 435 | # Filter out failed cases (marked as None) |
| 436 | valid_indices = [i for i, resp in enumerate(baseline_responses) if resp is not None] |
| 437 | if len(valid_indices) < len(questions): |
| 438 | failed_count = len(questions) - len(valid_indices) |
| 439 | logger.warning(f"Filtered out {failed_count} failed case(s) from batch") |
| 440 | |
| 441 | # If all cases failed, skip this batch |
| 442 | if not valid_indices: |
| 443 | logger.warning("All cases in batch failed, skipping batch") |
| 444 | return None |
| 445 | |
| 446 | # Filter data to only include valid cases |
| 447 | questions = [questions[i] for i in valid_indices] |
| 448 | standard_answers = [standard_answers[i] for i in valid_indices] |
| 449 | subjects = [subjects[i] for i in valid_indices] |
no test coverage detected