19 Class-Aware Safe Gate Agent¶
这一份 notebook 把前面两条线合在一起:
- class-aware:不同类别组用不同提问强度
- safe gate:在最后决定是否给出诊断,还是拒答
这是一个更像“收尾版方法”的实验,因为它不再只管问什么,也开始管:
什么时候应该停下,并且不要给答案。
我们当前的三组类别定义仍然是:
potential_benefit:nv,melhigh_risk:bkl,df,akiecneutral:bcc,vasc
这一版的核心目标不是单纯追求最高 Macro-F1,而是看:
- 能不能在保持一定 coverage 的同时
- 降低危险高置信错误
- 让整体行为比前面的 class-aware 更像“临床上能接受的系统”
In [21]:
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 [22]:
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[22]:
1481
2. 提取病例级结构¶
In [23]:
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[23]:
| 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 [24]:
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'
GROUP_POLICY = {
'potential_benefit': ['age', 'location', 'sex'],
'high_risk': ['age'],
'neutral': ['age', 'sex', 'location'],
}
case_df['predicted_group'] = case_df['top1_cls'].apply(get_class_group)
case_df[['image_id', 'top1_cls', 'predicted_group']].head()
Out[24]:
| 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 [25]:
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[25]:
{'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 [26]:
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 [27]:
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 [28]:
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 [29]:
trajectory, final_state = run_class_aware_safe_episode(case_df.iloc[0], max_questions=2)
trajectory, final_state
Out[29]:
([{'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 [30]:
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': 'class_aware_safe_gate_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 [31]:
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[31]:
{'agent_name': 'class_aware_safe_gate_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 [32]:
safe_tradeoff_df[['policy_label', 'coverage', 'selective_accuracy', 'selective_balanced_accuracy', 'selective_macro_f1']]
Out[32]:
| 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 [33]:
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[33]:
| 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 | class_aware_safe_gate_agent | 0.550304 | 0.777178 | 0.637282 | 0.608323 | 2.0 | 0.743417 | 0.449696 | 815.0 |
10. 保存结果¶
In [34]:
timestamp = datetime.now().strftime('%Y-%m-%d_%H%M%S')
cases_path = SUPPORT_DIR / f'{timestamp}_class_aware_safe_gate_cases.csv'
diagnosed_path = SUPPORT_DIR / f'{timestamp}_class_aware_safe_gate_diagnosed.csv'
abstained_path = SUPPORT_DIR / f'{timestamp}_class_aware_safe_gate_abstained.csv'
summary_path = SUPPORT_DIR / f'{timestamp}_class_aware_safe_gate_summary.csv'
result_path = SUPPORT_DIR / f'{timestamp}_class_aware_safe_gate_result.json'
comparison_path = SUPPORT_DIR / f'{timestamp}_class_aware_safe_gate_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_201139_class_aware_safe_gate_cases.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_201139_class_aware_safe_gate_diagnosed.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_201139_class_aware_safe_gate_abstained.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_201139_class_aware_safe_gate_summary.csv /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_201139_class_aware_safe_gate_result.json /Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_201139_class_aware_safe_gate_comparison.csv
11. 排查:为什么当前 selective 指标看起来退回了 image-only¶
当前 notebook 里的 evaluate_class_aware_safe_agent() 并没有对每个病例在提问后重新跑对应模型。
它做的是:
- 先根据策略决定病例最后落到哪个
final_model_key; - 再直接去拿这个模型在整个测试集上的总指标;
- 最后按病例占比做加权平均。
所以这一步得到的是模型级混合估计,不是病例级真实结果。 下面我们补一个真实逐例排查,直接看提问后的模型切换和预测变化。
In [35]:
# 修正 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 [36]:
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 [37]:
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[37]:
| 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 |