19 Conservative Safe Question Agent¶
这一份 notebook 把前面两条线合在一起:
- class-aware:不同类别组用不同提问强度
- safe gate:在最后决定是否给出诊断,还是拒答
这是一个更像“收尾版方法”的实验,因为它不再只管问什么,也开始管:
什么时候应该停下,并且不要给答案。
我们当前的三组类别定义仍然是:
potential_benefit:nv,melhigh_risk:bkl,df,akiecneutral:bcc,vasc
这一版的核心目标不是单纯追求最高 Macro-F1,而是看:
- 能不能在保持一定 coverage 的同时
- 降低危险高置信错误
- 让整体行为比前面的 class-aware 更像“临床上能接受的系统”
In [1]:
from pathlib import Path
from datetime import datetime
import json
import numpy as np
import pandas as pd
from sklearn.metrics import accuracy_score, balanced_accuracy_score, f1_score
PROJECT_ROOT = Path('/Users/applesues01/Documents/Medical_Agent')
DATA_DIR = PROJECT_ROOT / 'data' / 'HAM10000'
SPLIT_DIR = DATA_DIR / 'splits'
SUPPORT_DIR = PROJECT_ROOT / 'supports'
BASELINE_RESULTS_PATH = SUPPORT_DIR / 'baseline_results.csv'
MODEL_REGISTRY_PATH = SUPPORT_DIR / '2026-08-03_194704_all_saved_metadata_models.csv'
IMAGE_ONLY_CASE_PATH = SUPPORT_DIR / 'image_only_case_level_predictions.csv'
SAFE_TRADEOFF_PATH = SUPPORT_DIR / '2026-08-04_141020_safe_agent_tradeoff_summary.csv'
VALIDATED_AGENT_PATH = SUPPORT_DIR / '2026-08-03_185733_validated_agent_comparison.csv'
print(BASELINE_RESULTS_PATH)
print(MODEL_REGISTRY_PATH)
print(IMAGE_ONLY_CASE_PATH)
print(SAFE_TRADEOFF_PATH)
print(VALIDATED_AGENT_PATH)
/Users/applesues01/Documents/Medical_Agent/supports/baseline_results.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-03_194704_all_saved_metadata_models.csv /Users/applesues01/Documents/Medical_Agent/supports/image_only_case_level_predictions.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_141020_safe_agent_tradeoff_summary.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-03_185733_validated_agent_comparison.csv
1. 读取数据和结果表¶
In [2]:
test_df = pd.read_csv(SPLIT_DIR / 'test.csv')
baseline_df = pd.read_csv(BASELINE_RESULTS_PATH)
model_registry_df = pd.read_csv(MODEL_REGISTRY_PATH)
image_only_case_df = pd.read_csv(IMAGE_ONLY_CASE_PATH)
safe_tradeoff_df = pd.read_csv(SAFE_TRADEOFF_PATH)
validated_agent_df = pd.read_csv(VALIDATED_AGENT_PATH)
image_only_row = baseline_df[baseline_df['Method'] == 'Image Only'].iloc[0]
policy_model_scores = {
'image_only': {'macro_f1': float(image_only_row['Macro-F1'])},
}
report_model_scores = {
'image_only': {
'accuracy': float(image_only_row['Accuracy']),
'balanced_accuracy': float(image_only_row['Balanced Accuracy']),
'macro_f1': float(image_only_row['Macro-F1']),
}
}
for _, row in model_registry_df.iterrows():
key = str(row['Method'])
policy_model_scores[key] = {'macro_f1': float(row['Best Val Macro-F1'])}
report_model_scores[key] = {
'accuracy': float(row['Test Accuracy']),
'balanced_accuracy': float(row['Test Balanced Accuracy']),
'macro_f1': float(row['Test Macro-F1']),
}
case_df = test_df.merge(image_only_case_df, on=['image_id'], how='left')
len(case_df)
Out[2]:
1481
2. 提取病例级结构¶
In [3]:
prob_cols = [c for c in case_df.columns if c.startswith('prob_')]
class_names = [c.replace('prob_', '') for c in prob_cols]
def extract_case_structure(row):
probs = np.array([row[c] for c in prob_cols], dtype=float)
order = np.argsort(-probs)
top1_idx = int(order[0])
top2_idx = int(order[1])
top1_cls = class_names[top1_idx]
top2_cls = class_names[top2_idx]
top1_prob = float(probs[top1_idx])
top2_prob = float(probs[top2_idx])
margin = top1_prob - top2_prob
entropy = float(-(probs * np.log(probs + 1e-12)).sum())
return pd.Series({
'top1_cls': top1_cls,
'top2_cls': top2_cls,
'top1_prob': top1_prob,
'top2_prob': top2_prob,
'margin': margin,
'entropy': entropy,
})
case_df = pd.concat([case_df, case_df.apply(extract_case_structure, axis=1)], axis=1)
case_df[['image_id', 'dx', 'pred_label', 'top1_cls', 'top2_cls', 'margin', 'entropy']].head()
Out[3]:
| image_id | dx | pred_label | top1_cls | top2_cls | margin | entropy | |
|---|---|---|---|---|---|---|---|
| 0 | ISIC_0025837 | bkl | bkl | bkl | mel | 0.953966 | 0.150215 |
| 1 | ISIC_0025209 | bkl | bkl | bkl | akiec | 0.186130 | 1.487569 |
| 2 | ISIC_0029161 | bkl | bkl | bkl | nv | 0.665526 | 0.669199 |
| 3 | ISIC_0026273 | bkl | bkl | bkl | mel | 0.707970 | 0.723783 |
| 4 | ISIC_0025819 | bkl | bkl | bkl | nv | 0.969386 | 0.123900 |
3. 定义类别分组与 class-aware 提问顺序¶
In [4]:
CLASS_GROUPS = {
'potential_benefit': ['nv', 'mel'],
'high_risk': ['bkl', 'df', 'akiec'],
'neutral': ['bcc', 'vasc'],
}
def get_class_group(pred_cls):
for group_name, classes in CLASS_GROUPS.items():
if pred_cls in classes:
return group_name
return 'neutral'
# Conservative version: remove sex as an early question.
GROUP_POLICY = {
'potential_benefit': ['age', 'location'],
'high_risk': ['age'],
'neutral': ['age', 'location'],
}
case_df['predicted_group'] = case_df['top1_cls'].apply(get_class_group)
case_df[['image_id', 'top1_cls', 'predicted_group']].head()
Out[4]:
| image_id | top1_cls | predicted_group | |
|---|---|---|---|
| 0 | ISIC_0025837 | bkl | high_risk |
| 1 | ISIC_0025209 | bkl | high_risk |
| 2 | ISIC_0029161 | bkl | high_risk |
| 3 | ISIC_0026273 | bkl | high_risk |
| 4 | ISIC_0025819 | bkl | high_risk |
4. 定义 group-specific safe gate¶
这一版我们给不同组设置不同的诊断阈值:
potential_benefit: 稍微宽松一些neutral: 中等high_risk: 更严格
第一版先只用 image-only 的病例级结构来做 gate:
top1_probmargin
意思很直观:
- 高风险类需要更高置信度和更大 margin 才允许诊断
- 潜在受益类可以稍微放宽一点
In [5]:
GROUP_GATE = {
'potential_benefit': {'prob_thr': 0.85, 'margin_thr': 0.20},
'neutral': {'prob_thr': 0.90, 'margin_thr': 0.25},
'high_risk': {'prob_thr': 0.95, 'margin_thr': 0.30},
}
GROUP_GATE
Out[5]:
{'potential_benefit': {'prob_thr': 0.85, 'margin_thr': 0.2},
'neutral': {'prob_thr': 0.9, 'margin_thr': 0.25},
'high_risk': {'prob_thr': 0.95, 'margin_thr': 0.3}}
5. 状态函数¶
In [6]:
def build_initial_state(row):
return {
'image_id': row['image_id'],
'true_label': row['dx'],
'known_metadata': {},
'asked_questions': [],
'done': False,
'top1_cls': row['top1_cls'],
'top2_cls': row['top2_cls'],
'top1_prob': float(row['top1_prob']),
'top2_prob': float(row['top2_prob']),
'margin': float(row['margin']),
'entropy': float(row['entropy']),
'predicted_group': row['predicted_group'],
}
def ask_question(state, row, question):
new_state = dict(state)
new_state['known_metadata'] = dict(state['known_metadata'])
new_state['asked_questions'] = list(state['asked_questions'])
if question == 'age':
new_state['known_metadata']['age'] = row['age']
elif question == 'sex':
new_state['known_metadata']['sex'] = row['sex']
elif question == 'location':
new_state['known_metadata']['location'] = row['localization']
new_state['asked_questions'].append(question)
return new_state
def get_model_key_from_known_set(known_set):
if known_set == set():
return 'image_only'
if known_set == {'age'}:
return 'image_age'
if known_set == {'sex'}:
return 'image_sex'
if known_set == {'location'}:
return 'image_location'
if known_set == {'age', 'sex'}:
return 'image_age_sex'
if known_set == {'age', 'location'}:
return 'image_age_location'
if known_set == {'sex', 'location'}:
return 'image_sex_location'
if known_set == {'age', 'sex', 'location'}:
return 'image_all_metadata'
raise ValueError(f'Unknown known_set: {known_set}')
def get_model_key_from_state(state):
return get_model_key_from_known_set(set(state['known_metadata'].keys()))
def score_state(state):
model_key = get_model_key_from_state(state)
return policy_model_scores[model_key]['macro_f1']
6. 定义 class-aware safe gate policy¶
这一版 policy 的逻辑是:
- 如果当前病例已经达到所在组的安全条件,就直接诊断
- 否则根据所在组的提问顺序继续问
- 问到预算上限后,如果还不满足 gate,就拒答(abstain)
这次终于出现第三个动作:
askdiagnoseabstain
In [7]:
def class_aware_safe_gate_policy(state, max_questions=2):
group_name = state['predicted_group']
gate = GROUP_GATE[group_name]
if (state['top1_prob'] >= gate['prob_thr']) and (state['margin'] >= gate['margin_thr']):
return 'diagnose'
if len(state['asked_questions']) >= max_questions:
return 'abstain'
preferred_order = GROUP_POLICY[group_name]
for question in preferred_order:
if question not in state['asked_questions']:
return question
return 'abstain'
7. 先看单个病例轨迹¶
In [8]:
def run_class_aware_safe_episode(row, max_questions=2):
state = build_initial_state(row)
trajectory = []
while not state['done']:
action = class_aware_safe_gate_policy(state, max_questions=max_questions)
trajectory.append({
'step': len(trajectory),
'action': action,
'predicted_group': state['predicted_group'],
'model_key_before_action': get_model_key_from_state(state),
'score_before_action': score_state(state),
'known_metadata_before_action': dict(state['known_metadata']),
})
if action in ['diagnose', 'abstain']:
state['done'] = True
state['final_action'] = action
state['final_model_key'] = get_model_key_from_state(state)
state['final_score'] = score_state(state)
break
state = ask_question(state, row, action)
return trajectory, state
In [9]:
trajectory, final_state = run_class_aware_safe_episode(case_df.iloc[0], max_questions=2)
trajectory, final_state
Out[9]:
([{'step': 0,
'action': 'diagnose',
'predicted_group': 'high_risk',
'model_key_before_action': 'image_only',
'score_before_action': 0.6083234281767498,
'known_metadata_before_action': {}}],
{'image_id': 'ISIC_0025837',
'true_label': 'bkl',
'known_metadata': {},
'asked_questions': [],
'done': True,
'top1_cls': 'bkl',
'top2_cls': 'mel',
'top1_prob': 0.972960650920868,
'top2_prob': 0.0189945641905069,
'margin': 0.9539660867303611,
'entropy': 0.15021531890606907,
'predicted_group': 'high_risk',
'final_action': 'diagnose',
'final_model_key': 'image_only',
'final_score': 0.6083234281767498})
8. 全测试集评估¶
这里我们会同时输出:
- coverage
- abstain_rate
- diagnosed 子集上的 selective metrics
- 以及整体 expected metrics(只作为参考)
In [10]:
def evaluate_class_aware_safe_agent(max_questions=2):
records = []
for idx in range(len(case_df)):
row = case_df.iloc[idx]
trajectory, final_state = run_class_aware_safe_episode(row, max_questions=max_questions)
records.append({
'image_id': row['image_id'],
'true_label': row['dx'],
'predicted_group': row['predicted_group'],
'asked_questions': list(final_state['asked_questions']),
'num_questions': len(final_state['asked_questions']),
'final_action': final_state['final_action'],
'final_model_key': final_state['final_model_key'],
'final_policy_score': final_state['final_score'],
})
cases_df = pd.DataFrame(records)
diagnosed_df = cases_df[cases_df['final_action'] == 'diagnose'].copy()
abstained_df = cases_df[cases_df['final_action'] == 'abstain'].copy()
if len(diagnosed_df) > 0:
summary_df = diagnosed_df.groupby('final_model_key').size().reset_index(name='count')
summary_df['ratio_within_diagnosed'] = summary_df['count'] / len(diagnosed_df)
summary_df['test_accuracy'] = summary_df['final_model_key'].map(lambda k: report_model_scores[k]['accuracy'])
summary_df['test_balanced_accuracy'] = summary_df['final_model_key'].map(lambda k: report_model_scores[k]['balanced_accuracy'])
summary_df['test_macro_f1'] = summary_df['final_model_key'].map(lambda k: report_model_scores[k]['macro_f1'])
selective_accuracy = float((summary_df['ratio_within_diagnosed'] * summary_df['test_accuracy']).sum())
selective_bal_acc = float((summary_df['ratio_within_diagnosed'] * summary_df['test_balanced_accuracy']).sum())
selective_macro_f1 = float((summary_df['ratio_within_diagnosed'] * summary_df['test_macro_f1']).sum())
else:
summary_df = pd.DataFrame()
selective_accuracy = np.nan
selective_bal_acc = np.nan
selective_macro_f1 = np.nan
result = {
'agent_name': 'conservative_safe_question_agent',
'max_questions': max_questions,
'avg_questions': float(cases_df['num_questions'].mean()),
'coverage': float(len(diagnosed_df) / len(cases_df)),
'abstain_rate': float(len(abstained_df) / len(cases_df)),
'diagnosed_cases': int(len(diagnosed_df)),
'selective_accuracy': selective_accuracy,
'selective_balanced_accuracy': selective_bal_acc,
'selective_macro_f1': selective_macro_f1,
}
return cases_df, diagnosed_df, abstained_df, summary_df, result
In [11]:
safe_cases_df, safe_diagnosed_df, safe_abstained_df, safe_summary_df, safe_result = evaluate_class_aware_safe_agent(max_questions=2)
safe_result
Out[11]:
{'agent_name': 'conservative_safe_question_agent',
'max_questions': 2,
'avg_questions': 0.7434166103983795,
'coverage': 0.550303848750844,
'abstain_rate': 0.44969615124915596,
'diagnosed_cases': 815,
'selective_accuracy': 0.7771775827143822,
'selective_balanced_accuracy': 0.6372819453123182,
'selective_macro_f1': 0.6083234281767498}
9. 和前面的安全 trade-off 结果放到一起看¶
In [12]:
safe_tradeoff_df[['policy_label', 'coverage', 'selective_accuracy', 'selective_balanced_accuracy', 'selective_macro_f1']]
Out[12]:
| policy_label | coverage | selective_accuracy | selective_balanced_accuracy | selective_macro_f1 | |
|---|---|---|---|---|---|
| 0 | high_safety | 0.4112 | 0.9787 | 0.9175 | 0.9064 |
| 1 | balanced | 0.5165 | 0.9464 | 0.8591 | 0.8375 |
| 2 | high_coverage | 0.5881 | 0.9288 | 0.8274 | 0.7960 |
In [13]:
comparison_df = pd.concat([
safe_tradeoff_df[['policy_label', 'coverage', 'selective_accuracy', 'selective_balanced_accuracy', 'selective_macro_f1']].rename(columns={'policy_label': 'agent_name'}),
pd.DataFrame([safe_result])
], ignore_index=True)
comparison_df
Out[13]:
| agent_name | coverage | selective_accuracy | selective_balanced_accuracy | selective_macro_f1 | max_questions | avg_questions | abstain_rate | diagnosed_cases | |
|---|---|---|---|---|---|---|---|---|---|
| 0 | high_safety | 0.411200 | 0.978700 | 0.917500 | 0.906400 | NaN | NaN | NaN | NaN |
| 1 | balanced | 0.516500 | 0.946400 | 0.859100 | 0.837500 | NaN | NaN | NaN | NaN |
| 2 | high_coverage | 0.588100 | 0.928800 | 0.827400 | 0.796000 | NaN | NaN | NaN | NaN |
| 3 | conservative_safe_question_agent | 0.550304 | 0.777178 | 0.637282 | 0.608323 | 2.0 | 0.743417 | 0.449696 | 815.0 |
10. 保存结果¶
In [14]:
timestamp = datetime.now().strftime('%Y-%m-%d_%H%M%S')
cases_path = SUPPORT_DIR / f'{timestamp}_conservative_safe_question_cases.csv'
diagnosed_path = SUPPORT_DIR / f'{timestamp}_conservative_safe_question_diagnosed.csv'
abstained_path = SUPPORT_DIR / f'{timestamp}_conservative_safe_question_abstained.csv'
summary_path = SUPPORT_DIR / f'{timestamp}_conservative_safe_question_summary.csv'
result_path = SUPPORT_DIR / f'{timestamp}_conservative_safe_question_result.json'
comparison_path = SUPPORT_DIR / f'{timestamp}_conservative_safe_question_comparison.csv'
safe_cases_df.to_csv(cases_path, index=False)
safe_diagnosed_df.to_csv(diagnosed_path, index=False)
safe_abstained_df.to_csv(abstained_path, index=False)
safe_summary_df.to_csv(summary_path, index=False)
comparison_df.to_csv(comparison_path, index=False)
with open(result_path, 'w', encoding='utf-8') as f:
json.dump(safe_result, f, ensure_ascii=False, indent=2)
print(cases_path)
print(diagnosed_path)
print(abstained_path)
print(summary_path)
print(result_path)
print(comparison_path)
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204659_conservative_safe_question_cases.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204659_conservative_safe_question_diagnosed.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204659_conservative_safe_question_abstained.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204659_conservative_safe_question_summary.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204659_conservative_safe_question_result.json /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204659_conservative_safe_question_comparison.csv
11. 排查:为什么当前 selective 指标看起来退回了 image-only¶
当前 notebook 里的 evaluate_class_aware_safe_agent() 并没有对每个病例在提问后重新跑对应模型。
它做的是:
- 先根据策略决定病例最后落到哪个
final_model_key; - 再直接去拿这个模型在整个测试集上的总指标;
- 最后按病例占比做加权平均。
所以这一步得到的是模型级混合估计,不是病例级真实结果。 下面我们补一个真实逐例排查,直接看提问后的模型切换和预测变化。
In [15]:
# 修正 19 号 notebook 里的真实模型配置
combination_model_config = {
frozenset(): {
"name": "image_only",
"type": "image_only",
"checkpoint": str(PROJECT_ROOT / "checkpoints" / "resnet50_image_only_finetuned_best.pth"),
"metadata_features": [],
"metadata_embed_dim": None,
},
frozenset({"age"}): {
"name": "image_age",
"type": "fusion",
"checkpoint": str(PROJECT_ROOT / "checkpoints" / "image_age_best.pth"),
"metadata_features": ["age"],
"metadata_embed_dim": 256,
},
frozenset({"sex"}): {
"name": "image_sex",
"type": "fusion",
"checkpoint": str(PROJECT_ROOT / "checkpoints" / "image_sex_best.pth"),
"metadata_features": ["sex"],
"metadata_embed_dim": 32,
},
frozenset({"location"}): {
"name": "image_location",
"type": "fusion",
"checkpoint": str(PROJECT_ROOT / "checkpoints" / "image_location_best.pth"),
"metadata_features": ["location"],
"metadata_embed_dim": 64,
},
frozenset({"age", "sex"}): {
"name": "image_age_sex",
"type": "fusion",
"checkpoint": str(PROJECT_ROOT / "checkpoints" / "image_age_sex_best.pth"),
"metadata_features": ["age", "sex"],
"metadata_embed_dim": 32,
},
frozenset({"age", "location"}): {
"name": "image_age_location",
"type": "fusion",
"checkpoint": str(PROJECT_ROOT / "checkpoints" / "image_age_location_best.pth"),
"metadata_features": ["age", "location"],
"metadata_embed_dim": None,
},
frozenset({"sex", "location"}): {
"name": "image_sex_location",
"type": "fusion",
"checkpoint": str(PROJECT_ROOT / "checkpoints" / "image_sex_location_best.pth"),
"metadata_features": ["sex", "location"],
"metadata_embed_dim": 128,
},
frozenset({"age", "sex", "location"}): {
"name": "image_all_metadata",
"type": "fusion",
"checkpoint": str(PROJECT_ROOT / "checkpoints" / "image_all_metadata_best.pth"),
"metadata_features": ["age", "sex", "location"],
"metadata_embed_dim": 32,
},
}
loaded_models = {}
print("combination_model_config fixed")
combination_model_config fixed
In [16]:
import torch
import torch.nn as nn
from PIL import Image
from torchvision import transforms
from torchvision.models import resnet50
IMAGE_DIR1 = DATA_DIR / 'HAM10000_images_part_1'
IMAGE_DIR2 = DATA_DIR / 'HAM10000_images_part_2'
CLASS_NAMES = ['akiec', 'bcc', 'bkl', 'df', 'mel', 'nv', 'vasc']
def resolve_image_path(image_id):
filename = f'{image_id}.jpg'
path1 = IMAGE_DIR1 / filename
path2 = IMAGE_DIR2 / filename
return path1 if path1.exists() else path2
eval_transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
train_df_for_meta = pd.read_csv(SPLIT_DIR / 'train.csv')
train_age_mean = train_df_for_meta['age'].mean()
LOCATIONS = [
'scalp', 'ear', 'face', 'back', 'trunk', 'chest',
'upper extremity', 'abdomen', 'unknown', 'lower extremity',
'genital', 'neck', 'hand', 'foot', 'acral'
]
SEX_MAP = {
'male': [1.0, 0.0, 0.0],
'female': [0.0, 1.0, 0.0],
'unknown': [0.0, 0.0, 1.0],
}
def build_metadata_vector_from_row(meta_row, selected_features):
feats = []
if 'age' in selected_features:
age = meta_row['age']
if pd.isna(age):
age = train_age_mean
feats.append(float(age) / 100.0)
if 'sex' in selected_features:
sex_key = meta_row['sex'] if meta_row['sex'] in SEX_MAP else 'unknown'
feats.extend(SEX_MAP[sex_key])
if 'location' in selected_features:
loc_key = meta_row['localization'] if meta_row['localization'] in LOCATIONS else 'unknown'
loc_vector = [0.0] * len(LOCATIONS)
loc_vector[LOCATIONS.index(loc_key)] = 1.0
feats.extend(loc_vector)
return torch.tensor(feats, dtype=torch.float32).unsqueeze(0)
class ResNet50FeatureExtractor(nn.Module):
def __init__(self, backbone):
super().__init__()
self.features = nn.Sequential(*list(backbone.children())[:-1])
def forward(self, x):
x = self.features(x)
return torch.flatten(x, 1)
class RawFusionClassifier(nn.Module):
def __init__(self, metadata_input_dim, num_classes=7):
super().__init__()
self.classifier = nn.Sequential(
nn.Linear(2048 + metadata_input_dim, 512),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(512, 128),
nn.ReLU(),
nn.Linear(128, num_classes),
)
def forward(self, image_features, metadata):
fused = torch.cat([image_features, metadata], dim=1)
return self.classifier(fused)
class MetadataEncoderFusionClassifier(nn.Module):
def __init__(self, metadata_input_dim, metadata_embed_dim=64, num_classes=7):
super().__init__()
self.metadata_encoder = nn.Sequential(
nn.Linear(metadata_input_dim, metadata_embed_dim),
nn.BatchNorm1d(metadata_embed_dim),
nn.ReLU(),
nn.Dropout(0.2),
)
self.classifier = nn.Sequential(
nn.Linear(2048 + metadata_embed_dim, 512),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(512, 128),
nn.ReLU(),
nn.Linear(128, num_classes),
)
def forward(self, image_features, metadata):
metadata_features = self.metadata_encoder(metadata)
fused = torch.cat([image_features, metadata_features], dim=1)
return self.classifier(fused)
device = torch.device('mps' if torch.backends.mps.is_available() else 'cpu')
image_only_backbone = resnet50(weights=None)
image_only_backbone.fc = nn.Linear(image_only_backbone.fc.in_features, 7)
image_only_backbone.load_state_dict(torch.load(PROJECT_ROOT / 'checkpoints' / 'resnet50_image_only_finetuned_best.pth', map_location=device))
image_only_backbone = image_only_backbone.to(device)
image_only_backbone.eval()
feature_extractor = ResNet50FeatureExtractor(image_only_backbone).to(device)
feature_extractor.eval()
for p in feature_extractor.parameters():
p.requires_grad = False
combination_model_config = {frozenset(): {
'name': 'image_only',
'type': 'image_only',
'checkpoint': str(PROJECT_ROOT / 'checkpoints' / 'resnet50_image_only_finetuned_best.pth'),
'metadata_features': [],
'metadata_embed_dim': None,
}}
for _, row in model_registry_df.iterrows():
key = frozenset([x.strip() for x in str(row['Features']).split(',') if x.strip()])
combination_model_config[key] = {
'name': str(row['Method']),
'type': 'fusion',
'checkpoint': str(row['Checkpoint Path']),
'metadata_features': [x.strip() for x in str(row['Features']).split(',') if x.strip()],
# Do not trust the registry dim here; infer it from the checkpoint below.
'metadata_embed_dim': None,
}
loaded_models = {}
def infer_model_shape_from_checkpoint(checkpoint_path):
state_dict = torch.load(checkpoint_path, map_location='cpu')
if 'metadata_encoder.0.weight' in state_dict:
metadata_embed_dim = int(state_dict['metadata_encoder.0.weight'].shape[0])
return state_dict, metadata_embed_dim, 'encoded'
return state_dict, None, 'raw'
def load_combination_model(known_keys):
key = frozenset(known_keys)
config = combination_model_config[key]
if key in loaded_models:
return loaded_models[key], config
if config['type'] == 'image_only':
model = resnet50(weights=None)
model.fc = nn.Linear(model.fc.in_features, 7)
model.load_state_dict(torch.load(config['checkpoint'], map_location=device))
model = model.to(device)
model.eval()
loaded_models[key] = model
return model, config
state_dict, inferred_embed_dim, inferred_mode = infer_model_shape_from_checkpoint(config['checkpoint'])
dummy_row = {'age': train_age_mean, 'sex': 'unknown', 'localization': 'unknown'}
metadata_dim = len(build_metadata_vector_from_row(dummy_row, config['metadata_features']).squeeze(0))
if inferred_mode == 'raw':
model = RawFusionClassifier(metadata_input_dim=metadata_dim, num_classes=7)
config['metadata_embed_dim'] = None
else:
model = MetadataEncoderFusionClassifier(metadata_input_dim=metadata_dim, metadata_embed_dim=inferred_embed_dim, num_classes=7)
config['metadata_embed_dim'] = inferred_embed_dim
model.load_state_dict(state_dict)
model = model.to(device)
model.eval()
loaded_models[key] = model
return model, config
@torch.no_grad()
def predict_with_known_metadata(image_tensor, meta_row, known_keys):
model, config = load_combination_model(known_keys)
image_tensor = image_tensor.unsqueeze(0).to(device)
if config['type'] == 'image_only':
logits = model(image_tensor)
else:
image_features = feature_extractor(image_tensor)
metadata_tensor = build_metadata_vector_from_row(meta_row, config['metadata_features']).to(device)
logits = model(image_features, metadata_tensor)
probs = torch.softmax(logits, dim=1).squeeze(0).cpu().numpy()
pred_idx = int(np.argmax(probs))
return {
'model_name': config['name'],
'pred_label': CLASS_NAMES[pred_idx],
'pred_idx': pred_idx,
'max_prob': float(np.max(probs)),
'prob_vector': probs,
}
metadata_lookup = {row['image_id']: row for _, row in test_df.iterrows()}
print('debug predictor ready')
debug predictor ready
In [17]:
asked_case_ids = safe_cases_df[safe_cases_df['num_questions'] > 0]['image_id'].head(10).tolist()
debug_rows = []
for image_id in asked_case_ids:
row = case_df[case_df['image_id'] == image_id].iloc[0]
meta_row = metadata_lookup[image_id]
raw_image = Image.open(resolve_image_path(image_id)).convert('RGB')
image_tensor = eval_transform(raw_image)
trajectory, final_state = run_class_aware_safe_episode(row, max_questions=2)
pred0 = predict_with_known_metadata(image_tensor, meta_row, [])
pred_age = predict_with_known_metadata(image_tensor, meta_row, ['age'])
pred_sex = predict_with_known_metadata(image_tensor, meta_row, ['sex'])
pred_location = predict_with_known_metadata(image_tensor, meta_row, ['location'])
pred_age_location = predict_with_known_metadata(image_tensor, meta_row, ['age', 'location'])
pred_age_sex = predict_with_known_metadata(image_tensor, meta_row, ['age', 'sex'])
pred_sex_location = predict_with_known_metadata(image_tensor, meta_row, ['sex', 'location'])
pred_all = predict_with_known_metadata(image_tensor, meta_row, ['age', 'sex', 'location'])
debug_rows.append({
'image_id': image_id,
'true_label': row['dx'],
'predicted_group': row['predicted_group'],
'trajectory': trajectory,
'final_action': final_state['final_action'],
'final_model_key': final_state['final_model_key'],
'image_only_label': pred0['pred_label'],
'image_only_prob': pred0['max_prob'],
'age_label': pred_age['pred_label'],
'age_prob': pred_age['max_prob'],
'sex_label': pred_sex['pred_label'],
'sex_prob': pred_sex['max_prob'],
'location_label': pred_location['pred_label'],
'location_prob': pred_location['max_prob'],
'age_sex_label': pred_age_sex['pred_label'],
'age_sex_prob': pred_age_sex['max_prob'],
'age_location_label': pred_age_location['pred_label'],
'age_location_prob': pred_age_location['max_prob'],
'sex_location_label': pred_sex_location['pred_label'],
'sex_location_prob': pred_sex_location['max_prob'],
'all_label': pred_all['pred_label'],
'all_prob': pred_all['max_prob'],
})
debug_cases_df = pd.DataFrame(debug_rows)
debug_cases_df[[
'image_id', 'true_label', 'predicted_group', 'final_action', 'final_model_key',
'image_only_label', 'image_only_prob',
'age_label', 'age_prob',
'sex_label', 'sex_prob',
'location_label', 'location_prob',
'age_location_label', 'age_location_prob',
'sex_location_label', 'sex_location_prob',
'all_label', 'all_prob'
]]
Out[17]:
| image_id | true_label | predicted_group | final_action | final_model_key | image_only_label | image_only_prob | age_label | age_prob | sex_label | sex_prob | location_label | location_prob | age_location_label | age_location_prob | sex_location_label | sex_location_prob | all_label | all_prob | |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 0 | ISIC_0025209 | bkl | high_risk | abstain | image_age | bkl | 0.407122 | bkl | 0.788691 | mel | 0.490441 | bkl | 0.519953 | bkl | 0.818753 | bkl | 0.793395 | bkl | 0.752879 |
| 1 | ISIC_0029161 | bkl | high_risk | abstain | image_age | bkl | 0.793358 | bkl | 0.941818 | bkl | 0.987926 | nv | 0.528874 | bkl | 0.912514 | bkl | 0.907503 | bkl | 0.951908 |
| 2 | ISIC_0026273 | bkl | high_risk | abstain | image_age | bkl | 0.802558 | bkl | 0.963898 | bkl | 0.961713 | bkl | 0.470342 | bkl | 0.945809 | bkl | 0.957953 | bkl | 0.965302 |
| 3 | ISIC_0032013 | bkl | high_risk | abstain | image_age | bkl | 0.762678 | bkl | 0.988941 | bkl | 0.998160 | bkl | 0.965312 | bkl | 0.950229 | bkl | 0.935912 | bkl | 0.964793 |
| 4 | ISIC_0029289 | bkl | potential_benefit | abstain | image_age_location | nv | 0.464676 | nv | 0.879213 | nv | 0.983204 | nv | 0.936176 | nv | 0.704972 | nv | 0.917013 | nv | 0.959388 |
| 5 | ISIC_0029912 | bkl | high_risk | abstain | image_age | bkl | 0.617092 | bkl | 0.892666 | nv | 0.963171 | nv | 0.737985 | bkl | 0.587242 | bkl | 0.600730 | nv | 0.934774 |
| 6 | ISIC_0033539 | bkl | high_risk | abstain | image_age | bkl | 0.706276 | bkl | 0.928126 | bkl | 0.973478 | bkl | 0.838161 | bkl | 0.917992 | bkl | 0.944586 | bkl | 0.973286 |
| 7 | ISIC_0029022 | bkl | high_risk | abstain | image_age | bkl | 0.822935 | mel | 0.596700 | bkl | 0.549541 | mel | 0.670840 | bkl | 0.665434 | bkl | 0.541340 | bkl | 0.655241 |
| 8 | ISIC_0027957 | bkl | potential_benefit | abstain | image_age_location | nv | 0.798028 | nv | 0.973491 | nv | 0.998547 | nv | 0.952253 | nv | 0.900361 | nv | 0.942456 | nv | 0.999041 |
| 9 | ISIC_0031212 | bkl | potential_benefit | abstain | image_age_location | nv | 0.540349 | nv | 0.641004 | nv | 0.907318 | mel | 0.640403 | nv | 0.520501 | nv | 0.648289 | nv | 0.951641 |
因为它不是在大量“纠正错误”,而是在大量“发现问完也不可靠,于是拒答”。¶
12. 真正的新策略:Risk-Checked Safe Agent¶
这一版不再只是按固定顺序问,而是:
- 先对候选问题做“试探性预览”;
- 计算每个问题带来的收益和风险;
- 只有当最优问题的风险检查通过,才真的提问;否则直接拒答。
In [18]:
ALLOWED_QUESTIONS = {
'potential_benefit': ['age', 'location'],
'high_risk': ['age'],
'neutral': ['age', 'location'],
}
RISK_CHECK_CONFIG = {
'confidence_gain_weight': 1.0,
'margin_gain_weight': 0.5,
'wrong_label_penalty': 1.2,
'very_high_conf_penalty': 0.8,
'min_score_to_ask': 0.03,
'unsafe_prob_thr': 0.90,
}
def build_case_level_state(row):
image_id = row['image_id']
meta_row = metadata_lookup[image_id]
raw_image = Image.open(resolve_image_path(image_id)).convert('RGB')
image_tensor = eval_transform(raw_image)
initial_pred = predict_with_known_metadata(image_tensor, meta_row, [])
probs = initial_pred['prob_vector']
sorted_probs = np.sort(probs)[::-1]
initial_pred['margin'] = float(sorted_probs[0] - sorted_probs[1])
return {
'image_id': image_id,
'true_label': row['dx'],
'meta_row': meta_row,
'image_tensor': image_tensor,
'known_keys': [],
'asked_questions': [],
'current_pred': initial_pred,
'predicted_group': row['predicted_group'],
'done': False,
}
def preview_question_effect(state, question):
next_keys = list(state['known_keys']) + [question]
pred = predict_with_known_metadata(state['image_tensor'], state['meta_row'], next_keys)
probs = pred['prob_vector']
sorted_probs = np.sort(probs)[::-1]
pred['margin'] = float(sorted_probs[0] - sorted_probs[1])
current = state['current_pred']
score = (
RISK_CHECK_CONFIG['confidence_gain_weight'] * (pred['max_prob'] - current['max_prob'])
+ RISK_CHECK_CONFIG['margin_gain_weight'] * (pred['margin'] - current.get('margin', 0.0))
)
# Penalize if the question changes prediction to a different class at very high confidence.
if pred['pred_label'] != current['pred_label']:
score -= RISK_CHECK_CONFIG['wrong_label_penalty']
if pred['max_prob'] >= RISK_CHECK_CONFIG['unsafe_prob_thr'] and pred['pred_label'] != current['pred_label']:
score -= RISK_CHECK_CONFIG['very_high_conf_penalty']
return pred, float(score)
def choose_risk_checked_question(state):
group_name = state['predicted_group']
candidates = [q for q in ALLOWED_QUESTIONS[group_name] if q not in state['asked_questions']]
if not candidates:
return None, None, None
best_q, best_pred, best_score = None, None, None
for q in candidates:
pred, score = preview_question_effect(state, q)
if (best_score is None) or (score > best_score):
best_q, best_pred, best_score = q, pred, score
if best_score is None or best_score < RISK_CHECK_CONFIG['min_score_to_ask']:
return None, best_pred, best_score
return best_q, best_pred, best_score
def risk_checked_safe_policy(state, max_questions=2):
group_name = state['predicted_group']
gate = GROUP_GATE[group_name]
current = state['current_pred']
if (current['max_prob'] >= gate['prob_thr']) and (current.get('margin', 0.0) >= gate['margin_thr']):
return 'diagnose', None, None, None
if len(state['asked_questions']) >= max_questions:
return 'abstain', None, None, None
best_q, best_pred, best_score = choose_risk_checked_question(state)
if best_q is None:
return 'abstain', None, best_pred, best_score
return best_q, best_pred, best_score, group_name
In [19]:
def run_risk_checked_safe_episode(row, max_questions=2):
state = build_case_level_state(row)
trajectory = []
while not state['done']:
action, preview_pred, preview_score, group_name = risk_checked_safe_policy(state, max_questions=max_questions)
trajectory.append({
'step': len(trajectory),
'action': action,
'predicted_group': state['predicted_group'],
'known_keys_before_action': list(state['known_keys']),
'current_label': state['current_pred']['pred_label'],
'current_prob': state['current_pred']['max_prob'],
'preview_label': None if preview_pred is None else preview_pred['pred_label'],
'preview_prob': None if preview_pred is None else preview_pred['max_prob'],
'preview_score': preview_score,
})
if action == 'diagnose':
state['done'] = True
state['final_action'] = 'diagnose'
state['final_pred'] = state['current_pred']
break
if action == 'abstain':
state['done'] = True
state['final_action'] = 'abstain'
state['final_pred'] = state['current_pred']
break
state['asked_questions'].append(action)
state['known_keys'].append(action)
state['current_pred'] = preview_pred
return trajectory, state
def evaluate_risk_checked_safe_agent(max_questions=2):
records = []
for idx in range(len(case_df)):
row = case_df.iloc[idx]
trajectory, final_state = run_risk_checked_safe_episode(row, max_questions=max_questions)
final_pred = final_state['final_pred']
final_correct = (final_pred['pred_label'] == row['dx'])
records.append({
'image_id': row['image_id'],
'true_label': row['dx'],
'predicted_group': row['predicted_group'],
'asked_questions': list(final_state['asked_questions']),
'num_questions': len(final_state['asked_questions']),
'final_action': final_state['final_action'],
'final_pred_label': final_pred['pred_label'],
'final_max_prob': final_pred['max_prob'],
'final_correct': final_correct,
'trajectory': trajectory,
})
cases_df = pd.DataFrame(records)
diagnosed_df = cases_df[cases_df['final_action'] == 'diagnose'].copy()
abstained_df = cases_df[cases_df['final_action'] == 'abstain'].copy()
if len(diagnosed_df) > 0:
selective_accuracy = float((diagnosed_df['final_correct']).mean())
selective_bal_acc = float(balanced_accuracy_score(diagnosed_df['true_label'], diagnosed_df['final_pred_label']))
selective_macro_f1 = float(f1_score(diagnosed_df['true_label'], diagnosed_df['final_pred_label'], average='macro'))
else:
selective_accuracy = np.nan
selective_bal_acc = np.nan
selective_macro_f1 = np.nan
result = {
'agent_name': 'risk_checked_safe_agent',
'max_questions': max_questions,
'avg_questions': float(cases_df['num_questions'].mean()),
'coverage': float(len(diagnosed_df) / len(cases_df)),
'abstain_rate': float(len(abstained_df) / len(cases_df)),
'diagnosed_cases': int(len(diagnosed_df)),
'selective_accuracy': selective_accuracy,
'selective_balanced_accuracy': selective_bal_acc,
'selective_macro_f1': selective_macro_f1,
}
return cases_df, diagnosed_df, abstained_df, result
risk_cases_df, risk_diagnosed_df, risk_abstained_df, risk_result = evaluate_risk_checked_safe_agent(max_questions=2)
risk_result
Out[19]:
{'agent_name': 'risk_checked_safe_agent',
'max_questions': 2,
'avg_questions': 0.3153274814314652,
'coverage': 0.7663740715732613,
'abstain_rate': 0.23362592842673868,
'diagnosed_cases': 1135,
'selective_accuracy': 0.8863436123348017,
'selective_balanced_accuracy': 0.7245939548463518,
'selective_macro_f1': 0.733477344970264}
In [20]:
comparison_df = pd.concat([
safe_tradeoff_df[['policy_label', 'coverage', 'selective_accuracy', 'selective_balanced_accuracy', 'selective_macro_f1']].rename(columns={'policy_label': 'agent_name'}),
pd.DataFrame([safe_result]),
pd.DataFrame([risk_result])
], ignore_index=True)
comparison_df
Out[20]:
| agent_name | coverage | selective_accuracy | selective_balanced_accuracy | selective_macro_f1 | max_questions | avg_questions | abstain_rate | diagnosed_cases | |
|---|---|---|---|---|---|---|---|---|---|
| 0 | high_safety | 0.411200 | 0.978700 | 0.917500 | 0.906400 | NaN | NaN | NaN | NaN |
| 1 | balanced | 0.516500 | 0.946400 | 0.859100 | 0.837500 | NaN | NaN | NaN | NaN |
| 2 | high_coverage | 0.588100 | 0.928800 | 0.827400 | 0.796000 | NaN | NaN | NaN | NaN |
| 3 | conservative_safe_question_agent | 0.550304 | 0.777178 | 0.637282 | 0.608323 | 2.0 | 0.743417 | 0.449696 | 815.0 |
| 4 | risk_checked_safe_agent | 0.766374 | 0.886344 | 0.724594 | 0.733477 | 2.0 | 0.315327 | 0.233626 | 1135.0 |
13. 风险检查策略的使用分布与拒答分析¶
In [21]:
question_usage_df = (
risk_cases_df.explode('asked_questions')
.dropna(subset=['asked_questions'])
.groupby('asked_questions')
.size()
.reset_index(name='count')
.rename(columns={'asked_questions': 'question'})
)
if len(question_usage_df) > 0:
question_usage_df['ratio_among_all_cases'] = question_usage_df['count'] / len(risk_cases_df)
question_usage_df['ratio_among_asked_cases'] = question_usage_df['count'] / max((risk_cases_df['num_questions'] > 0).sum(), 1)
question_usage_df
abstain_by_class_df = (
risk_cases_df.groupby('true_label')
.agg(
total_cases=('image_id', 'count'),
abstained_cases=('final_action', lambda x: int((x == 'abstain').sum())),
diagnosed_cases=('final_action', lambda x: int((x == 'diagnose').sum())),
)
.reset_index()
)
abstain_by_class_df['abstain_rate'] = abstain_by_class_df['abstained_cases'] / abstain_by_class_df['total_cases']
abstain_by_class_df['diagnose_rate'] = abstain_by_class_df['diagnosed_cases'] / abstain_by_class_df['total_cases']
abstain_by_class_df = abstain_by_class_df.sort_values('abstain_rate', ascending=False)
abstain_by_class_df
Out[21]:
| true_label | total_cases | abstained_cases | diagnosed_cases | abstain_rate | diagnose_rate | |
|---|---|---|---|---|---|---|
| 3 | df | 20 | 12 | 8 | 0.600000 | 0.400000 |
| 2 | bkl | 168 | 85 | 83 | 0.505952 | 0.494048 |
| 0 | akiec | 46 | 21 | 25 | 0.456522 | 0.543478 |
| 4 | mel | 165 | 59 | 106 | 0.357576 | 0.642424 |
| 6 | vasc | 19 | 6 | 13 | 0.315789 | 0.684211 |
| 1 | bcc | 71 | 21 | 50 | 0.295775 | 0.704225 |
| 5 | nv | 992 | 142 | 850 | 0.143145 | 0.856855 |
In [22]:
asked_cases_df = risk_cases_df[risk_cases_df['num_questions'] > 0].copy()
if len(asked_cases_df) > 0:
asked_cases_df['first_question'] = asked_cases_df['asked_questions'].apply(lambda x: x[0] if len(x) > 0 else None)
first_question_perf_df = (
asked_cases_df.groupby('first_question')
.agg(
cases=('image_id', 'count'),
final_accuracy=('final_correct', 'mean'),
avg_questions=('num_questions', 'mean'),
)
.reset_index()
.sort_values('cases', ascending=False)
)
else:
first_question_perf_df = pd.DataFrame()
first_question_perf_df
Out[22]:
| first_question | cases | final_accuracy | avg_questions | |
|---|---|---|---|---|
| 0 | age | 278 | 0.661871 | 1.028777 |
| 1 | location | 176 | 0.579545 | 1.028409 |
In [23]:
def to_jsonable(value):
if isinstance(value, dict):
return {k: to_jsonable(v) for k, v in value.items()}
if isinstance(value, list):
return [to_jsonable(v) for v in value]
if isinstance(value, tuple):
return [to_jsonable(v) for v in value]
if isinstance(value, np.ndarray):
return value.tolist()
if isinstance(value, np.integer):
return int(value)
if isinstance(value, np.floating):
return float(value)
if isinstance(value, np.bool_):
return bool(value)
return value
14. 保存 risk-checked safe agent 结果¶
In [24]:
timestamp = datetime.now().strftime('%Y-%m-%d_%H%M%S')
risk_cases_path = SUPPORT_DIR / f'{timestamp}_risk_checked_safe_cases.csv'
risk_diagnosed_path = SUPPORT_DIR / f'{timestamp}_risk_checked_safe_diagnosed.csv'
risk_abstained_path = SUPPORT_DIR / f'{timestamp}_risk_checked_safe_abstained.csv'
risk_result_path = SUPPORT_DIR / f'{timestamp}_risk_checked_safe_result.json'
risk_comparison_path = SUPPORT_DIR / f'{timestamp}_risk_checked_safe_comparison.csv'
question_usage_path = SUPPORT_DIR / f'{timestamp}_risk_checked_safe_question_usage.csv'
abstain_by_class_path = SUPPORT_DIR / f'{timestamp}_risk_checked_safe_abstain_by_class.csv'
first_question_perf_path = SUPPORT_DIR / f'{timestamp}_risk_checked_safe_first_question_perf.csv'
risk_cases_to_save = risk_cases_df.copy()
risk_cases_to_save['asked_questions'] = risk_cases_to_save['asked_questions'].apply(
lambda x: json.dumps(to_jsonable(x), ensure_ascii=False)
)
risk_cases_to_save['trajectory'] = risk_cases_to_save['trajectory'].apply(
lambda x: json.dumps(to_jsonable(x), ensure_ascii=False)
)
risk_cases_to_save.to_csv(risk_cases_path, index=False)
risk_diagnosed_to_save = risk_diagnosed_df.copy()
risk_diagnosed_to_save['asked_questions'] = risk_diagnosed_to_save['asked_questions'].apply(
lambda x: json.dumps(to_jsonable(x), ensure_ascii=False)
)
risk_diagnosed_to_save['trajectory'] = risk_diagnosed_to_save['trajectory'].apply(
lambda x: json.dumps(to_jsonable(x), ensure_ascii=False)
)
risk_diagnosed_to_save.to_csv(risk_diagnosed_path, index=False)
risk_abstained_to_save = risk_abstained_df.copy()
risk_abstained_to_save['asked_questions'] = risk_abstained_to_save['asked_questions'].apply(
lambda x: json.dumps(to_jsonable(x), ensure_ascii=False)
)
risk_abstained_to_save['trajectory'] = risk_abstained_to_save['trajectory'].apply(
lambda x: json.dumps(to_jsonable(x), ensure_ascii=False)
)
risk_abstained_to_save.to_csv(risk_abstained_path, index=False)
comparison_df.to_csv(risk_comparison_path, index=False)
question_usage_df.to_csv(question_usage_path, index=False)
abstain_by_class_df.to_csv(abstain_by_class_path, index=False)
first_question_perf_df.to_csv(first_question_perf_path, index=False)
with open(risk_result_path, 'w', encoding='utf-8') as f:
json.dump(risk_result, f, ensure_ascii=False, indent=2)
print(risk_cases_path)
print(risk_diagnosed_path)
print(risk_abstained_path)
print(risk_result_path)
print(risk_comparison_path)
print(question_usage_path)
print(abstain_by_class_path)
print(first_question_perf_path)
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204948_risk_checked_safe_cases.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204948_risk_checked_safe_diagnosed.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204948_risk_checked_safe_abstained.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204948_risk_checked_safe_result.json /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204948_risk_checked_safe_comparison.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204948_risk_checked_safe_question_usage.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204948_risk_checked_safe_abstain_by_class.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204948_risk_checked_safe_first_question_perf.csv
15. Calibrated Risk-Checked Agent¶
In [26]:
from sklearn.metrics import log_loss
import torch.optim as optim
val_df_cal = pd.read_csv(SPLIT_DIR / 'val.csv')
@torch.no_grad()
def predict_logits_with_known_metadata(image_tensor, meta_row, known_keys):
model, config = load_combination_model(known_keys)
image_tensor = image_tensor.unsqueeze(0).to(device)
if config['type'] == 'image_only':
logits = model(image_tensor)
else:
image_features = feature_extractor(image_tensor)
metadata_tensor = build_metadata_vector_from_row(meta_row, config['metadata_features']).to(device)
logits = model(image_features, metadata_tensor)
return logits.squeeze(0).detach().cpu()
def build_image_only_prediction_df(dataframe):
rows = []
for _, row in dataframe.iterrows():
img = Image.open(resolve_image_path(row['image_id'])).convert('RGB')
image_tensor = eval_transform(img)
pred = predict_with_known_metadata(image_tensor, row, [])
rows.append({
'image_id': row['image_id'],
'top1_cls': pred['pred_label'],
'max_prob': pred['max_prob'],
})
return pd.DataFrame(rows)
val_case_df = val_df_cal.merge(build_image_only_prediction_df(val_df_cal), on='image_id', how='left')
val_case_df['predicted_group'] = val_case_df['top1_cls'].apply(get_class_group)
def fit_temperature_for_known_keys(dataframe, known_keys):
logits_list, labels = [], []
for _, row in dataframe.iterrows():
img = Image.open(resolve_image_path(row['image_id'])).convert('RGB')
image_tensor = eval_transform(img)
logits = predict_logits_with_known_metadata(image_tensor, row, known_keys)
logits_list.append(logits)
labels.append(CLASS_NAMES.index(row['dx']))
logits = torch.stack(logits_list)
labels = torch.tensor(labels, dtype=torch.long)
temperature = torch.nn.Parameter(torch.ones(1) * 1.0)
optimizer = optim.LBFGS([temperature], lr=0.05, max_iter=50)
criterion = torch.nn.CrossEntropyLoss()
def closure():
optimizer.zero_grad()
loss = criterion(logits / temperature.clamp(min=1e-3), labels)
loss.backward()
return loss
optimizer.step(closure)
return float(temperature.detach().clamp(min=1e-3).item())
temperature_map = {
frozenset(): fit_temperature_for_known_keys(val_df_cal, []),
frozenset({'age'}): fit_temperature_for_known_keys(val_df_cal, ['age']),
frozenset({'location'}): fit_temperature_for_known_keys(val_df_cal, ['location']),
frozenset({'age','location'}): fit_temperature_for_known_keys(val_df_cal, ['age','location']),
}
temperature_map
Out[26]:
{frozenset(): 1.3430252075195312,
frozenset({'age'}): 1.9865938425064087,
frozenset({'location'}): 2.0149312019348145,
frozenset({'age', 'location'}): 1.5314850807189941}
In [27]:
@torch.no_grad()
def calibrated_predict_with_known_metadata(image_tensor, meta_row, known_keys):
logits = predict_logits_with_known_metadata(image_tensor, meta_row, known_keys)
T = temperature_map[frozenset(known_keys)]
probs = torch.softmax(logits / T, dim=0).numpy()
pred_idx = int(np.argmax(probs))
sorted_probs = np.sort(probs)[::-1]
return {
'pred_label': CLASS_NAMES[pred_idx],
'pred_idx': pred_idx,
'max_prob': float(np.max(probs)),
'prob_vector': probs,
'margin': float(sorted_probs[0] - sorted_probs[1]),
}
def evaluate_calibrated_risk_checked_agent(max_questions=2):
records = []
for _, row in case_df.iterrows():
raw_image = Image.open(resolve_image_path(row['image_id'])).convert('RGB')
image_tensor = eval_transform(raw_image)
state = {
'image_id': row['image_id'], 'true_label': row['dx'], 'meta_row': metadata_lookup[row['image_id']],
'image_tensor': image_tensor, 'known_keys': [], 'asked_questions': [],
'current_pred': calibrated_predict_with_known_metadata(image_tensor, metadata_lookup[row['image_id']], []),
'predicted_group': row['predicted_group'], 'done': False,
}
while not state['done']:
group_name = state['predicted_group']
gate = GROUP_GATE[group_name]
current = state['current_pred']
if (current['max_prob'] >= gate['prob_thr']) and (current['margin'] >= gate['margin_thr']):
final_pred = current; final_action = 'diagnose'; break
if len(state['asked_questions']) >= max_questions:
final_pred = current; final_action = 'abstain'; break
candidates = [q for q in ALLOWED_QUESTIONS[group_name] if q not in state['asked_questions']]
best_q, best_pred, best_score = None, None, None
for q in candidates:
next_keys = list(state['known_keys']) + [q]
pred = calibrated_predict_with_known_metadata(image_tensor, metadata_lookup[row['image_id']], next_keys)
score = (pred['max_prob'] - current['max_prob']) + 0.5 * (pred['margin'] - current['margin'])
if pred['pred_label'] != current['pred_label']:
score -= 1.2
if best_score is None or score > best_score:
best_q, best_pred, best_score = q, pred, score
if best_q is None or best_score < 0.03:
final_pred = current; final_action = 'abstain'; break
state['asked_questions'].append(best_q)
state['known_keys'].append(best_q)
state['current_pred'] = best_pred
records.append({
'image_id': row['image_id'], 'true_label': row['dx'], 'final_action': final_action,
'final_pred_label': final_pred['pred_label'], 'final_correct': final_pred['pred_label'] == row['dx'],
'num_questions': len(state['asked_questions']), 'asked_questions': list(state['asked_questions'])
})
out = pd.DataFrame(records)
diag = out[out['final_action']=='diagnose'].copy()
result = {
'agent_name': 'calibrated_risk_checked_agent',
'coverage': float(len(diag)/len(out)),
'avg_questions': float(out['num_questions'].mean()),
'selective_accuracy': float(diag['final_correct'].mean()) if len(diag) else np.nan,
'selective_balanced_accuracy': float(balanced_accuracy_score(diag['true_label'], diag['final_pred_label'])) if len(diag) else np.nan,
'selective_macro_f1': float(f1_score(diag['true_label'], diag['final_pred_label'], average='macro')) if len(diag) else np.nan,
}
return out, diag, result
cal_cases_df, cal_diag_df, cal_result = evaluate_calibrated_risk_checked_agent()
cal_result
Out[27]:
{'agent_name': 'calibrated_risk_checked_agent',
'coverage': 0.5813639432815665,
'avg_questions': 0.36191762322754895,
'selective_accuracy': 0.9500580720092915,
'selective_balanced_accuracy': 0.8581758613408753,
'selective_macro_f1': 0.8764710391574163}
In [28]:
timestamp = datetime.now().strftime('%Y-%m-%d_%H%M%S')
pd.DataFrame([cal_result]).to_csv(SUPPORT_DIR / f'{timestamp}_calibrated_risk_checked_result.csv', index=False)
cal_cases_df.to_csv(SUPPORT_DIR / f'{timestamp}_calibrated_risk_checked_cases.csv', index=False)
print(timestamp)
2026-08-04_211137