CoolFace
Datasetpublic

PerturbReason/PerturbReason_dataset_code

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes12downloads
extract_mini_output_0502.py195 linesDownload Raw Back to Baselines
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()