Files
2026-09-19 12:54:45 +08:00

59 lines
1.9 KiB
Python

"""Output stability metric using Jaccard similarity."""
from typing import Dict, List, Set
from sqlalchemy.orm import Session
from app.analysis.parser import Finding, parse_output
from app.models import ExperimentRun
def _finding_key(finding: Finding) -> str:
parts = [finding.defect_type.lower()]
if finding.line_start is not None:
parts.append(str(finding.line_start))
if finding.line_end is not None:
parts.append(str(finding.line_end))
return "|".join(parts)
def compute_stability_score(db: Session, experiment_id: str, model_id: str, level: str) -> float:
"""Compute average pairwise Jaccard across three repeats for each sample."""
from uuid import UUID
runs = (
db.query(ExperimentRun)
.filter_by(experiment_id=UUID(experiment_id), model_id=model_id)
.all()
)
by_sample: Dict[str, List[Set[str]]] = {}
for run in runs:
if run.template_version.template.level != level:
continue
if not run.result or not run.result.raw_output:
continue
sample_id = str(run.sample_id)
findings = set(_finding_key(f) for f in parse_output(run.result.raw_output, level))
by_sample.setdefault(sample_id, []).append(findings)
scores = []
for sample_id, repeats in by_sample.items():
if len(repeats) < 2:
continue
# pairwise Jaccard for up to 3 repeats
pairs = [(0, 1), (0, 2), (1, 2)]
pair_scores = []
for i, j in pairs:
if i < len(repeats) and j < len(repeats):
a, b = repeats[i], repeats[j]
union = a | b
if not union:
pair_scores.append(1.0)
else:
pair_scores.append(len(a & b) / len(union))
if pair_scores:
scores.append(sum(pair_scores) / len(pair_scores))
return round(sum(scores) / len(scores), 4) if scores else 0.0