MCPcopy Create free account
hub / github.com/Bairong-Xdynamics/MistakeNotebookLearning / _process_batch

Method _process_batch

mnl/trainer.py:392–622  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

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]

Callers 1

trainMethod · 0.95

Calls 9

classify_subjectsMethod · 0.80
batch_generateMethod · 0.80
update_entryMethod · 0.80
evaluate_batchMethod · 0.80
_save_entriesMethod · 0.80
itemsMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected