ahmedtaha100's picture
Upload aggregate.py with huggingface_hub
124c010 verified
Raw
History Blame Contribute Delete
18.2 kB
from __future__ import annotations
import argparse
import contextlib
import csv
import json
import logging
from collections import Counter, defaultdict
from pathlib import Path
import numpy as np
logger = logging.getLogger(__name__)
def write_validation_results_csv(
output_path: Path,
pairs: list[dict],
reviews_by_reviewer: dict[str, list[dict]],
exclusions: dict,
) -> None:
pair_strata = {idx: classify_transformation(pair) for idx, pair in enumerate(pairs)}
implausible_any = set(exclusions["implausible_any_rater"])
implausible_majority = set(exclusions["implausible_majority"])
disputed_quality = set(exclusions["disputed_quality"])
final_excluded = set(exclusions["all_excluded"])
fieldnames = [
"pair_index",
"reviewer_id",
"clinically_plausible",
"pathology_preserved",
"quality_score",
"comments",
"timestamp",
"transform_type",
"excluded_implausible_any",
"excluded_implausible_majority",
"excluded_disputed_quality",
"excluded_final",
]
output_path.parent.mkdir(parents=True, exist_ok=True)
with open(output_path, "w", newline="") as f:
writer = csv.DictWriter(f, fieldnames=fieldnames, extrasaction="ignore")
writer.writeheader()
for reviewer, rows in sorted(reviews_by_reviewer.items()):
for row in sorted(rows, key=lambda r: int(r["pair_index"])):
pair_index = int(row["pair_index"])
writer.writerow(
{
"pair_index": pair_index,
"reviewer_id": row.get("reviewer_id", reviewer),
"clinically_plausible": row.get("clinically_plausible", ""),
"pathology_preserved": row.get("pathology_preserved", ""),
"quality_score": row.get("quality_score", ""),
"comments": row.get("comments", ""),
"timestamp": row.get("timestamp", ""),
"transform_type": pair_strata.get(pair_index, "unknown"),
"excluded_implausible_any": int(pair_index in implausible_any),
"excluded_implausible_majority": int(pair_index in implausible_majority),
"excluded_disputed_quality": int(pair_index in disputed_quality),
"excluded_final": int(pair_index in final_excluded),
}
)
def load_all_reviews(results_dir: Path) -> dict[str, list[dict]]:
reviews_by_reviewer: dict[str, list[dict]] = {}
for path in sorted(results_dir.glob("reviews_*.csv")):
reviewer = path.stem.replace("reviews_", "")
rows: list[dict] = []
with open(path) as f:
for row in csv.DictReader(f):
rows.append(row)
if rows:
reviews_by_reviewer[reviewer] = rows
return reviews_by_reviewer
def classify_transformation(pair: dict) -> str:
age_changed = pair.get("source_age") != pair.get("target_age")
sex_changed = pair.get("source_sex") != pair.get("target_sex")
if pair.get("_transform_type"):
return str(pair["_transform_type"])
if age_changed and sex_changed:
return "intersectional"
if age_changed:
return "age_only"
if sex_changed:
return "sex_only"
return "none"
def build_rating_matrix(
reviews_by_reviewer: dict[str, list[dict]],
field: str,
num_pairs: int,
value_map: dict[str, float] | None = None,
) -> np.ndarray:
reviewers = sorted(reviews_by_reviewer.keys())
matrix = np.full((num_pairs, len(reviewers)), np.nan)
for j, reviewer in enumerate(reviewers):
for row in reviews_by_reviewer[reviewer]:
idx = int(row["pair_index"])
if idx < num_pairs and row.get(field):
val = row[field]
if value_map and val in value_map:
matrix[idx, j] = value_map[val]
else:
with contextlib.suppress(ValueError, TypeError):
matrix[idx, j] = float(val)
return matrix
def fleiss_kappa(matrix: np.ndarray, categories: list[int | float]) -> float:
valid_rows = ~np.any(np.isnan(matrix), axis=1)
data = matrix[valid_rows]
n_subjects = data.shape[0]
n_raters = data.shape[1]
if n_subjects == 0 or n_raters < 2:
return float("nan")
counts = np.zeros((n_subjects, len(categories)))
for k_idx, k in enumerate(categories):
counts[:, k_idx] = np.sum(data == k, axis=1)
p_j = np.sum(counts, axis=0) / (n_subjects * n_raters)
p_i = (np.sum(counts**2, axis=1) - n_raters) / (n_raters * (n_raters - 1))
p_bar = np.mean(p_i)
p_e = np.sum(p_j**2)
if abs(1.0 - p_e) < 1e-10:
return 1.0
return float((p_bar - p_e) / (1.0 - p_e))
def build_stratum_rating_matrix(
reviews_by_reviewer: dict[str, list[dict]],
field: str,
indices: list[int],
value_map: dict[str, float] | None = None,
) -> np.ndarray:
index_set = set(indices)
reviewers = sorted(reviews_by_reviewer.keys())
idx_to_row = {idx: row_num for row_num, idx in enumerate(sorted(indices))}
matrix = np.full((len(indices), len(reviewers)), np.nan)
for j, reviewer in enumerate(reviewers):
for row in reviews_by_reviewer[reviewer]:
pair_idx = int(row["pair_index"])
if pair_idx in index_set and row.get(field):
val = row[field]
matrix_row = idx_to_row[pair_idx]
if value_map and val in value_map:
matrix[matrix_row, j] = value_map[val]
else:
with contextlib.suppress(ValueError, TypeError):
matrix[matrix_row, j] = float(val)
return matrix
def compute_stratum_kappa(
pairs: list[dict],
reviews_by_reviewer: dict[str, list[dict]],
) -> dict[str, dict[str, float | None]]:
pair_strata: dict[int, str] = {}
for i, pair in enumerate(pairs):
pair_strata[i] = classify_transformation(pair)
strata = ["age_only", "sex_only", "intersectional"]
result: dict[str, dict[str, float | None]] = {}
for stratum in strata:
indices = [i for i, s in pair_strata.items() if s == stratum]
if not indices:
result[stratum] = {
"fleiss_kappa_quality": None,
"fleiss_kappa_plausibility": None,
"fleiss_kappa_pathology_preservation": None,
}
continue
q_matrix = build_stratum_rating_matrix(reviews_by_reviewer, "quality_score", indices)
p_matrix = build_stratum_rating_matrix(
reviews_by_reviewer, "clinically_plausible", indices,
value_map={"Yes": 1.0, "No": 0.0},
)
pp_matrix = build_stratum_rating_matrix(
reviews_by_reviewer, "pathology_preserved", indices,
value_map={"Yes": 1.0, "No": 0.0, "Uncertain": 0.5},
)
kq = fleiss_kappa(q_matrix, categories=[1, 2, 3, 4, 5])
kp = fleiss_kappa(p_matrix, categories=[0, 1])
kpp = fleiss_kappa(pp_matrix, categories=[0, 0.5, 1])
result[stratum] = {
"fleiss_kappa_quality": round(kq, 4) if not np.isnan(kq) else None,
"fleiss_kappa_plausibility": round(kp, 4) if not np.isnan(kp) else None,
"fleiss_kappa_pathology_preservation": round(kpp, 4) if not np.isnan(kpp) else None,
}
return result
def compute_exclusions(
reviews_by_reviewer: dict[str, list[dict]],
num_pairs: int,
) -> dict:
plausible_counts: dict[int, dict[str, int]] = defaultdict(lambda: {"Yes": 0, "No": 0})
quality_scores: dict[int, list[int]] = defaultdict(list)
reviewer_counts: dict[int, int] = defaultdict(int)
for reviews in reviews_by_reviewer.values():
for row in reviews:
idx = int(row["pair_index"])
reviewer_counts[idx] += 1
p = row.get("clinically_plausible", "Yes")
if p in ("Yes", "No"):
plausible_counts[idx][p] += 1
if row.get("quality_score"):
quality_scores[idx].append(int(row["quality_score"]))
excluded_implausible: list[int] = []
excluded_majority: list[int] = []
excluded_disputed: list[int] = []
for idx in range(num_pairs):
if reviewer_counts[idx] == 0:
continue
no_votes = plausible_counts[idx]["No"]
if no_votes > 0:
excluded_implausible.append(idx)
total = reviewer_counts[idx]
if no_votes > total / 2:
excluded_majority.append(idx)
scores = quality_scores.get(idx, [])
if len(scores) >= 3 and (max(scores) - min(scores)) >= 3 and idx not in excluded_implausible:
score_counts = Counter(scores)
has_majority = any(c > len(scores) / 2 for c in score_counts.values())
if not has_majority:
excluded_disputed.append(idx)
all_excluded = sorted(set(excluded_implausible + excluded_disputed))
return {
"implausible_any_rater": excluded_implausible,
"implausible_majority": excluded_majority,
"disputed_quality": excluded_disputed,
"all_excluded": all_excluded,
}
def compute_stratum_stats(
pairs: list[dict],
reviews_by_reviewer: dict[str, list[dict]],
exclusions: dict,
) -> dict:
pair_strata: dict[int, str] = {}
for i, pair in enumerate(pairs):
pair_strata[i] = classify_transformation(pair)
all_reviews_by_pair: dict[int, list[dict]] = defaultdict(list)
for reviews in reviews_by_reviewer.values():
for row in reviews:
idx = int(row["pair_index"])
all_reviews_by_pair[idx].append(row)
strata = ["age_only", "sex_only", "intersectional"]
stats: dict[str, dict] = {}
excluded_set = set(exclusions["all_excluded"])
for stratum in strata:
indices = [i for i, s in pair_strata.items() if s == stratum]
if not indices:
stats[stratum] = {"total": 0}
continue
n_excluded = sum(1 for i in indices if i in excluded_set)
n_passed = len(indices) - n_excluded
plausible_yes = 0
plausible_total = 0
preserved_yes = 0
preserved_total = 0
quality_vals: list[int] = []
for idx in indices:
for row in all_reviews_by_pair.get(idx, []):
if row.get("clinically_plausible") in ("Yes", "No"):
plausible_total += 1
if row["clinically_plausible"] == "Yes":
plausible_yes += 1
if row.get("pathology_preserved") in ("Yes", "No"):
preserved_total += 1
if row["pathology_preserved"] == "Yes":
preserved_yes += 1
if row.get("quality_score"):
quality_vals.append(int(row["quality_score"]))
stats[stratum] = {
"total": len(indices),
"passed": n_passed,
"excluded": n_excluded,
"pass_rate": round(n_passed / len(indices), 3) if indices else 0,
"plausibility_rate": round(plausible_yes / plausible_total, 3) if plausible_total else None,
"pathology_preservation_rate": round(preserved_yes / preserved_total, 3) if preserved_total else None,
"mean_quality": round(float(np.mean(quality_vals)), 2) if quality_vals else None,
"quality_distribution": {
str(k): int(np.sum(np.array(quality_vals) == k))
for k in range(1, 6)
} if quality_vals else {},
}
return stats
def main() -> None:
logging.basicConfig(level=logging.INFO)
parser = argparse.ArgumentParser(description="Aggregate radiologist reviews")
parser.add_argument("--pairs", type=str, required=True)
parser.add_argument("--results-dir", type=str, default="results")
parser.add_argument("--output", type=str, default="validation_report.json")
parser.add_argument("--csv-output", type=str, default=None)
args = parser.parse_args()
with open(args.pairs) as f:
pairs = json.load(f)
num_pairs = len(pairs)
reviews_by_reviewer = load_all_reviews(Path(args.results_dir))
reviewers = sorted(reviews_by_reviewer.keys())
if not reviewers:
logger.warning("No review files found.")
return
quality_matrix = build_rating_matrix(
reviews_by_reviewer, "quality_score", num_pairs,
)
plausible_matrix = build_rating_matrix(
reviews_by_reviewer, "clinically_plausible", num_pairs,
value_map={"Yes": 1.0, "No": 0.0},
)
preserved_matrix = build_rating_matrix(
reviews_by_reviewer, "pathology_preserved", num_pairs,
value_map={"Yes": 1.0, "No": 0.0, "Uncertain": 0.5},
)
kappa_quality = fleiss_kappa(quality_matrix, categories=[1, 2, 3, 4, 5])
kappa_plausible = fleiss_kappa(plausible_matrix, categories=[0, 1])
kappa_preserved = fleiss_kappa(preserved_matrix, categories=[0, 0.5, 1])
exclusions = compute_exclusions(reviews_by_reviewer, num_pairs)
stratum_stats = compute_stratum_stats(pairs, reviews_by_reviewer, exclusions)
stratum_kappa = compute_stratum_kappa(pairs, reviews_by_reviewer)
total_reviewed = sum(len(rows) for rows in reviews_by_reviewer.values())
per_reviewer: dict[str, dict] = {}
all_quality: list[int] = []
for reviewer, rows in reviews_by_reviewer.items():
scores = [int(r["quality_score"]) for r in rows if r.get("quality_score")]
n_implausible = sum(1 for r in rows if r.get("clinically_plausible") == "No")
n_not_preserved = sum(1 for r in rows if r.get("pathology_preserved") == "No")
per_reviewer[reviewer] = {
"reviewed": len(rows),
"mean_quality": round(float(np.mean(scores)), 2) if scores else None,
"flagged_implausible": n_implausible,
"flagged_not_preserved": n_not_preserved,
}
all_quality.extend(scores)
quality_dist = {str(k): int(np.sum(np.array(all_quality) == k)) for k in range(1, 6)} if all_quality else {}
valid_quality = quality_matrix[~np.all(np.isnan(quality_matrix), axis=1)]
mean_quality = float(np.nanmean(valid_quality)) if valid_quality.size > 0 else None
report = {
"num_pairs": num_pairs,
"num_reviewers": len(reviewers),
"reviewers": reviewers,
"total_reviews": total_reviewed,
"inter_rater_reliability": {
"overall": {
"fleiss_kappa_quality": round(kappa_quality, 4) if not np.isnan(kappa_quality) else None,
"fleiss_kappa_plausibility": round(kappa_plausible, 4) if not np.isnan(kappa_plausible) else None,
"fleiss_kappa_pathology_preservation": round(kappa_preserved, 4) if not np.isnan(kappa_preserved) else None,
},
"by_transformation_type": stratum_kappa,
},
"overall": {
"mean_quality_score": round(mean_quality, 2) if mean_quality else None,
"quality_distribution": quality_dist,
"pass_rate": round(
(num_pairs - len(exclusions["all_excluded"])) / num_pairs, 3
) if num_pairs > 0 else 0,
"plausibility_rate": round(
float(np.nanmean(plausible_matrix[~np.all(np.isnan(plausible_matrix), axis=1)])), 3
) if plausible_matrix.size > 0 else None,
"pathology_preservation_rate": round(
float(np.nanmean(preserved_matrix[~np.all(np.isnan(preserved_matrix), axis=1)])), 3
) if preserved_matrix.size > 0 else None,
},
"by_transformation_type": stratum_stats,
"exclusions": {
"implausible_any_rater": len(exclusions["implausible_any_rater"]),
"implausible_majority_vote": len(exclusions["implausible_majority"]),
"disputed_quality": len(exclusions["disputed_quality"]),
"total_excluded": len(exclusions["all_excluded"]),
"excluded_pair_indices": exclusions["all_excluded"],
},
"per_reviewer": per_reviewer,
}
output_path = Path(args.output)
output_path.parent.mkdir(parents=True, exist_ok=True)
with open(output_path, "w") as f:
json.dump(report, f, indent=2)
exclusion_path = output_path.parent / "exclusion_list.json"
with open(exclusion_path, "w") as f:
json.dump({"excluded_pair_indices": exclusions["all_excluded"]}, f, indent=2)
csv_output_path = (
Path(args.csv_output)
if args.csv_output
else output_path.parent / "validation_results.csv"
)
write_validation_results_csv(csv_output_path, pairs, reviews_by_reviewer, exclusions)
logger.info("Reviewers: %d", len(reviewers))
logger.info("Total reviews: %d", total_reviewed)
irr = report["inter_rater_reliability"]["overall"]
logger.info("Fleiss kappa (quality): %s", irr["fleiss_kappa_quality"])
logger.info("Fleiss kappa (plausibility): %s", irr["fleiss_kappa_plausibility"])
logger.info("Fleiss kappa (pathology): %s", irr["fleiss_kappa_pathology_preservation"])
for stratum, sk in stratum_kappa.items():
logger.info(
" %s kappa: quality=%s plausibility=%s pathology=%s",
stratum, sk["fleiss_kappa_quality"], sk["fleiss_kappa_plausibility"],
sk["fleiss_kappa_pathology_preservation"],
)
logger.info("Overall pass rate: %s", report["overall"]["pass_rate"])
logger.info("Total excluded: %d", report["exclusions"]["total_excluded"])
for stratum, stats in stratum_stats.items():
if stats.get("total", 0) > 0:
logger.info(" %s: %d/%d passed (%s)", stratum, stats["passed"], stats["total"], stats["pass_rate"])
logger.info("Report: %s", output_path)
logger.info("Exclusion list: %s", exclusion_path)
logger.info("Validation CSV: %s", csv_output_path)
if __name__ == "__main__":
main()