"""Analyzer for SFT (supervised fine-tuning) ID sets.""" import json from pathlib import Path from typing import Dict, Any, Set, List import logging from .base_analyzer import BaseAnalyzer logger = logging.getLogger(__name__) class SFTAnalyzer(BaseAnalyzer): """Analyze SFT ID sets and their overlap with training data.""" def __init__(self, output_dir): super().__init__("sft", output_dir) def load_sft_ids(self, sft_file_path: Path) -> Dict[str, Set[str]]: """Load SFT IDs from JSON file. Args: sft_file_path: Path to the SFT IDs JSON file Returns: Dictionary mapping subset names to ID sets """ with open(sft_file_path, "r") as f: sft_data = json.load(f) sft_sets = {} for key, id_list in sft_data.items(): if isinstance(id_list, list): sft_sets[key] = set(id_list) else: logger.warning(f"Unexpected type for SFT key {key}: {type(id_list)}") return sft_sets def analyze_sft_file(self, sft_file_path: Path, train_ids: Set[str]) -> Dict[str, Any]: """Analyze a single SFT file against training IDs. Args: sft_file_path: Path to SFT IDs file train_ids: Set of training IDs Returns: Analysis results """ sft_sets = self.load_sft_ids(sft_file_path) results = {"file": str(sft_file_path), "subsets": {}} total_sft_ids = set() for subset_name, subset_ids in sft_sets.items(): total_sft_ids.update(subset_ids) # Find missing IDs (in SFT but not in training) missing_ids = subset_ids - train_ids subset_info = { "total_ids": len(subset_ids), "found_in_train": len(subset_ids - missing_ids), "missing_from_train": len(missing_ids), "missing_ratio": len(missing_ids) / len(subset_ids) if subset_ids else 0, "missing_id_sample": sorted(missing_ids)[:20], } results["subsets"][subset_name] = subset_info # Overall statistics all_missing = total_sft_ids - train_ids results["summary"] = { "total_subsets": len(sft_sets), "total_unique_sft_ids": len(total_sft_ids), "total_found_in_train": len(total_sft_ids - all_missing), "total_missing_from_train": len(all_missing), "overall_missing_ratio": len(all_missing) / len(total_sft_ids) if total_sft_ids else 0, } return results def process_chunk(self, records: List[Dict[str, Any]]) -> Dict[str, Any]: """Not used for SFT analysis - we analyze the separate SFT files.""" return {} def aggregate(self, chunk_results: List[Dict[str, Any]]) -> Dict[str, Any]: """Not used for SFT analysis.""" return {}