You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 

42 lines
6.2 KiB

"""Pilot experiment readiness diagnostics and transparent artifact export."""
from __future__ import annotations
import csv,json
from dataclasses import asdict
from datetime import datetime
from pathlib import Path
from .config_loader import PilotDiagnosticsSettings
from .improvement_recommender import ImprovementRecommender
from .models import ConditionBalanceDiagnostic,PilotDiagnosticMetric,PilotDiagnosticReport
from .session_comparator import SessionComparator
def _b(v)->bool:return str(v).lower() in {"true","1","yes"}
class PilotDiagnostics:
def __init__(self,config:PilotDiagnosticsSettings)->None:self.config=config
def _rows(self,root:Path|None)->list[dict[str,str]]:
rows=[]
if root and root.exists():
for p in root.rglob("*person_analysis_table.csv"):
with p.open(encoding="utf-8-sig",newline="") as f:rows.extend(csv.DictReader(f))
return rows
def run(self,analysis_dir:Path,logs_dir:Path|None=None)->PilotDiagnosticReport:
rows=self._rows(logs_dir);total=len(rows);pr=[r for r in rows if r.get("prompt_condition")=="prompt"];cr=[r for r in rows if r.get("prompt_condition")=="control"];pv=sum(_b(r.get("valid_for_voice_analysis")) for r in pr);cv=sum(_b(r.get("valid_for_voice_analysis")) for r in cr);valid=pv+cv
def rate(items,key,predicate=lambda r:True):return sum(predicate(r) for r in items)/len(items) if items else None
pvr=[r for r in pr if _b(r.get("valid_for_voice_analysis"))];cvr=[r for r in cr if _b(r.get("valid_for_voice_analysis"))];p_resp=rate(pvr,"",lambda r:_b(r.get("response_detected")));c_resp=rate(cvr,"",lambda r:_b(r.get("response_detected")));ne=rate(rows,"",lambda r:r.get("turn_level")=="not_evaluable");np=rate(rows,"",lambda r:not _b(r.get("pose_estimated_ever")));vr=valid/total if total else None;ratio=max(pv,cv)/min(pv,cv) if min(pv,cv)>0 else None;balanced=ratio is not None and ratio<=self.config.max_condition_valid_count_ratio;balance=ConditionBalanceDiagnostic(len(pr),len(cr),pv,cv,p_resp,c_resp,abs(pv-cv),ratio,balanced,"条件人数は許容範囲" if balanced else "条件間の有効人数差が大きい","不足条件を同じ環境で追加収集する")
camera_missing=any(not r.get("camera_position_note","").strip() for r in rows);metrics=[self._metric("valid_voice_analysis_rate",vr,self.config.min_valid_voice_analysis_rate,vr is not None and vr>=self.config.min_valid_voice_analysis_rate,"有効音声解析率","除外理由を確認する"),self._metric("not_evaluable_rate",ne,self.config.max_not_evaluable_rate,ne is not None and ne<=self.config.max_not_evaluable_rate,"評価不能率","カメラ位置と顔処理を確認する"),self._metric("no_pose_estimated_rate",np,self.config.max_no_pose_estimated_rate,np is not None and np<=self.config.max_no_pose_estimated_rate,"Pose未推定率","顔が映る位置へ調整する")]
sessions=SessionComparator().summarize(rows);recs=ImprovementRecommender().recommend(vr,ne,np,balance,pv,cv,camera_missing)
if not rows or (self.config.ready_requires_prompt_and_control and (not pr or not cr)) or min(pv,cv)<self.config.min_valid_records_per_condition:status="insufficient_data"
elif (vr is not None and vr<self.config.min_valid_voice_analysis_rate) or (ne is not None and ne>self.config.max_not_evaluable_rate) or (np is not None and np>self.config.max_no_pose_estimated_rate) or camera_missing:status="needs_major_adjustment"
elif min(pv,cv)<self.config.target_valid_records_per_condition or not balanced:status="needs_minor_adjustment"
else:status="ready"
return PilotDiagnosticReport(f"diagnostic_{datetime.now():%Y%m%d_%H%M%S}",datetime.now().astimezone().isoformat(timespec="seconds"),str(analysis_dir),str(logs_dir) if logs_dir else None,status,status=="ready",total,pv,cv,p_resp,c_resp,vr,ne,np,balance,sessions,metrics,recs,"診断は本実験移行の目安であり、研究妥当性や因果関係を保証しない")
def _metric(self,name,value,threshold,passed,message,action):return PilotDiagnosticMetric(name,value,threshold,passed,"info" if passed else "warning",message,action)
def export(self,report:PilotDiagnosticReport,output:Path)->list[Path]:
output.mkdir(parents=True,exist_ok=True);paths=[]
paths.append(self._csv(output/"pilot_diagnostic_metrics.csv",[asdict(x) for x in report.metrics],["metric_name","value","threshold","passed","severity","message","suggested_action"]))
recs=[]
for x in report.recommendations:d=asdict(x);d["related_metrics_json"]=json.dumps(d.pop("related_metrics"),ensure_ascii=False);recs.append(d)
paths.append(self._csv(output/"improvement_recommendations.csv",recs,["recommendation_id","category","priority","title","reason","suggested_action","related_metrics_json","expected_effect"]))
paths.append(self._csv(output/"session_quality_summary.csv",[asdict(x) for x in report.session_summaries],list(asdict(report.session_summaries[0]).keys()) if report.session_summaries else ["session_id","condition","total_tracks","valid_voice_analysis_tracks","valid_voice_analysis_rate","response_detected_count","response_rate","not_evaluable_count","not_evaluable_rate","no_pose_estimated_count","no_pose_estimated_rate","mean_reaction_time_sec","median_reaction_time_sec","camera_position_note","quality_level","main_issue","suggested_action"]))
md=output/"pilot_diagnostic_report.md";lines=["# パイロット実験診断",f"- overall_status: {report.overall_status}",f"- ready_for_main_experiment: {report.ready_for_main_experiment}",f"- Prompt有効人数: {report.prompt_valid_records}",f"- Control有効人数: {report.control_valid_records}",f"- valid_voice_analysis_rate: {report.valid_voice_analysis_rate}",f"- not_evaluable_rate: {report.not_evaluable_rate}",f"- no_pose_estimated_rate: {report.no_pose_estimated_rate}","","## 改善提案"]+[f"- [{r.priority}] {r.title}: {r.suggested_action}" for r in report.recommendations]+["","## 注意",report.notes];md.write_text("\n".join(lines)+"\n",encoding="utf-8");paths.insert(0,md);return paths
def _csv(self,path,rows,fields):
with path.open("w",encoding="utf-8-sig",newline="") as f:w=csv.DictWriter(f,fieldnames=fields);w.writeheader();w.writerows(rows)
return path