PerturbReason/PerturbReason_dataset_code
012
1"""Create centralized mini-0426 subsets for multiple model output folders.2 3This script reuses the mini_0426 benchmark's 1-based sample_idx values to4subset model outputs produced on the full 0426 benchmark. It supports multiple5source trees with different filename prefixes and context layouts, and writes6all extracted subsets into a single LLM_api output root.7"""8 9from __future__ import annotations10 11import argparse12from pathlib import Path13 14from extract_mini_outputs_0426 import read_jsonl, subset_records, write_jsonl15 16 17MINI_BASE_DEFAULT = Path(18 "/scratch/group/PROJECT_NAME/OMICS_DATA/baseline_collections/LLM_api/mini_0426"19)20OUTPUT_ROOT_DEFAULT = Path(21 "/scratch/group/PROJECT_NAME/OMICS_DATA/baseline_collections/LLM_api/mini_subset_outputs_0502"22)23 24JOBS = {25 "SFT": {26 "source_root": Path(27 "/scratch/group/PROJECT_NAME/OMICS_DATA/baseline_collections/SFT_result/our_model"28 ),29 "output_name": "SFT",30 "contexts": {31 "hidden_context": {32 "source_dir": "hidden",33 "mini_variant": "hidden_context",34 "prefix": "",35 },36 "noisy_context": {37 "source_dir": "noisy",38 "mini_variant": "noisy_context",39 "prefix": "",40 },41 },42 },43}44 45 46def extract_context(47 mini_base: Path,48 source_root: Path,49 source_dir: str,50 mini_variant: str,51 prefix: str,52 out_dir: Path,53) -> int:54 mini_variant_dir = mini_base / mini_variant55 if not mini_variant_dir.exists():56 raise FileNotFoundError(f"Missing mini variant directory: {mini_variant_dir}")57 58 full_variant_dir = source_root / source_dir59 if not full_variant_dir.exists():60 raise FileNotFoundError(f"Missing full output directory: {full_variant_dir}")61 62 written = 063 for split_dir in sorted(path for path in mini_variant_dir.iterdir() if path.is_dir()):64 split_written = 065 for mini_file in sorted(split_dir.glob("*.jsonl")):66 full_file = full_variant_dir / split_dir.name / f"{prefix}{mini_file.name}"67 if not full_file.exists():68 raise FileNotFoundError(f"Missing full output file: {full_file}")69 70 mini_records = read_jsonl(mini_file)71 sample_indices = get_sample_indices_from_records(mini_records, mini_file)72 full_records = read_jsonl(full_file)73 subset = subset_records_with_fallback(74 full_records=full_records,75 mini_records=mini_records,76 sample_indices=sample_indices,77 full_path=full_file,78 )79 80 out_file = out_dir / split_dir.name / full_file.name81 write_jsonl(subset, out_file)82 83 split_written += len(subset)84 written += len(subset)85 print(86 f"[{out_dir.parent.name}/{out_dir.name}/{split_dir.name}] {mini_file.name}: "87 f"{len(full_records)} -> {len(subset)} rows"88 )89 90 print(91 f"[{out_dir.parent.name}/{out_dir.name}/{split_dir.name}] wrote {split_written} rows"92 )93 94 return written95 96 97def get_sample_indices_from_records(mini_records: list[dict], mini_path: Path) -> list[int]:98 indices: list[int] = []99 for record in mini_records:100 sample_idx = record.get("sample_idx")101 if not isinstance(sample_idx, int) or sample_idx <= 0:102 raise ValueError(f"Invalid sample_idx in {mini_path}: {sample_idx!r}")103 indices.append(sample_idx)104 return indices105 106 107def subset_records_with_fallback(108 full_records: list[dict],109 mini_records: list[dict],110 sample_indices: list[int],111 full_path: Path,112) -> list[dict]:113 try:114 return subset_records(full_records, sample_indices, full_path)115 except IndexError:116 return subset_records_by_id(full_records, mini_records, full_path)117 118 119def subset_records_by_id(120 full_records: list[dict], mini_records: list[dict], full_path: Path121) -> list[dict]:122 full_records_by_id: dict[str, dict] = {}123 for record in full_records:124 record_id = record.get("id")125 if not isinstance(record_id, str) or not record_id:126 raise KeyError(f"Cannot fall back to id matching for {full_path}: missing id")127 if record_id in full_records_by_id:128 raise KeyError(f"Duplicate id {record_id!r} in {full_path}")129 full_records_by_id[record_id] = record130 131 subset: list[dict] = []132 missing_ids: list[str] = []133 for mini_record in mini_records:134 record_id = mini_record.get("id")135 if not isinstance(record_id, str) or not record_id:136 raise KeyError(f"Mini record missing id while matching {full_path}")137 if record_id not in full_records_by_id:138 missing_ids.append(record_id)139 continue140 141 record = dict(full_records_by_id[record_id])142 record["sample_idx"] = mini_record["sample_idx"]143 subset.append(record)144 145 if missing_ids:146 print(147 f" [warn] {full_path.name}: matched {len(subset)}/{len(mini_records)} rows by id; "148 f"{len(missing_ids)} ids missing"149 )150 151 return subset152 153 154def parse_args() -> argparse.Namespace:155 parser = argparse.ArgumentParser()156 parser.add_argument("--mini_base", type=Path, default=MINI_BASE_DEFAULT)157 parser.add_argument("--output_root", type=Path, default=OUTPUT_ROOT_DEFAULT)158 parser.add_argument(159 "--job",160 action="append",161 choices=sorted(JOBS),162 help="Run only the named job. Can be supplied multiple times.",163 )164 return parser.parse_args()165 166 167def main() -> None:168 args = parse_args()169 selected_jobs = args.job or list(JOBS)170 171 grand_total = 0172 for job_name in selected_jobs:173 job = JOBS[job_name]174 job_total = 0175 for context_name, context in job["contexts"].items():176 out_dir = args.output_root / job["output_name"] / context_name177 context_total = extract_context(178 mini_base=args.mini_base,179 source_root=job["source_root"],180 source_dir=context["source_dir"],181 mini_variant=context["mini_variant"],182 prefix=context["prefix"],183 out_dir=out_dir,184 )185 job_total += context_total186 print(f"[{job_name}/{context_name}] total rows written: {context_total}")187 188 grand_total += job_total189 print(f"[{job_name}] total rows written: {job_total}")190 191 print(f"Done. Total mini rows written: {grand_total}")192 193 194if __name__ == "__main__":195 main()