19 Conservative Safe Question Agent¶

这一份 notebook 把前面两条线合在一起:

  • class-aware:不同类别组用不同提问强度
  • safe gate:在最后决定是否给出诊断,还是拒答

这是一个更像“收尾版方法”的实验,因为它不再只管问什么,也开始管:

什么时候应该停下,并且不要给答案。

我们当前的三组类别定义仍然是:

  • potential_benefit: nv, mel
  • high_risk: bkl, df, akiec
  • neutral: 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_prob
  • margin

意思很直观:

  • 高风险类需要更高置信度和更大 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 的逻辑是:

  1. 如果当前病例已经达到所在组的安全条件,就直接诊断
  2. 否则根据所在组的提问顺序继续问
  3. 问到预算上限后,如果还不满足 gate,就拒答(abstain)

这次终于出现第三个动作:

  • ask
  • diagnose
  • abstain
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_204712_conservative_safe_question_cases.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204712_conservative_safe_question_diagnosed.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204712_conservative_safe_question_abstained.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204712_conservative_safe_question_summary.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204712_conservative_safe_question_result.json
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_204712_conservative_safe_question_comparison.csv

11. 排查:为什么当前 selective 指标看起来退回了 image-only¶

当前 notebook 里的 evaluate_class_aware_safe_agent() 并没有对每个病例在提问后重新跑对应模型。 它做的是:

  1. 先根据策略决定病例最后落到哪个 final_model_key;
  2. 再直接去拿这个模型在整个测试集上的总指标;
  3. 最后按病例占比做加权平均。

所以这一步得到的是模型级混合估计,不是病例级真实结果。 下面我们补一个真实逐例排查,直接看提问后的模型切换和预测变化。

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¶

这一版不再只是按固定顺序问,而是:

  1. 先对候选问题做“试探性预览”;
  2. 计算每个问题带来的收益和风险;
  3. 只有当最优问题的风险检查通过,才真的提问;否则直接拒答。
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_205018_risk_checked_safe_cases.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_205018_risk_checked_safe_diagnosed.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_205018_risk_checked_safe_abstained.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_205018_risk_checked_safe_result.json
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_205018_risk_checked_safe_comparison.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_205018_risk_checked_safe_question_usage.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_205018_risk_checked_safe_abstain_by_class.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_205018_risk_checked_safe_first_question_perf.csv

15. Class-Conditional Gate Agent¶

In [26]:
val_df_gate = pd.read_csv(SPLIT_DIR / 'val.csv')

def build_image_only_prediction_df(dataframe):
    rows = []
    for _, row in dataframe.iterrows():
        raw_image = Image.open(resolve_image_path(row['image_id'])).convert('RGB')
        image_tensor = eval_transform(raw_image)
        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_gate.merge(build_image_only_prediction_df(val_df_gate), on='image_id', how='left')

class_thresholds = {}
for cls in sorted(val_case_df['top1_cls'].dropna().unique()):
    sub = val_case_df[val_case_df['top1_cls'] == cls].copy()
    best_thr, best_util = 0.5, -1
    for thr in np.arange(0.4, 0.96, 0.05):
        mask = sub['max_prob'] >= thr
        cov = mask.mean() if len(sub) else 0
        if mask.sum() == 0:
            continue
        acc = (sub.loc[mask, 'top1_cls'] == sub.loc[mask, 'dx']).mean()
        util = cov * acc
        if util > best_util:
            best_thr, best_util = thr, util
    class_thresholds[cls] = float(best_thr)
class_thresholds
Out[26]:
{'akiec': 0.4,
 'bcc': 0.4,
 'bkl': 0.4,
 'df': 0.4,
 'mel': 0.4,
 'nv': 0.4,
 'vasc': 0.4}
In [27]:
def evaluate_class_conditional_gate_agent():
    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)
        meta_row = metadata_lookup[row['image_id']]
        p0 = predict_with_known_metadata(image_tensor, meta_row, [])
        cls_thr = class_thresholds.get(p0['pred_label'], 0.75)
        if p0['max_prob'] >= cls_thr:
            final_action='diagnose'; final_pred=p0; asked=[]
        else:
            p_age = predict_with_known_metadata(image_tensor, meta_row, ['age'])
            p_loc = predict_with_known_metadata(image_tensor, meta_row, ['location'])
            final_pred = p_age if p_age['max_prob'] >= p_loc['max_prob'] else p_loc
            final_action = 'diagnose' if final_pred['max_prob'] >= cls_thr else 'abstain'
            asked = ['age'] if final_pred is p_age else ['location']
        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(asked), 'asked_questions': asked
        })
    out = pd.DataFrame(records)
    diag = out[out['final_action']=='diagnose'].copy()
    result = {
        'agent_name': 'class_conditional_gate_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

cc_cases_df, cc_diag_df, cc_result = evaluate_class_conditional_gate_agent()
cc_result
Out[27]:
{'agent_name': 'class_conditional_gate_agent',
 'coverage': 1.0,
 'avg_questions': 0.022957461174881837,
 'selective_accuracy': 0.7819041188386225,
 'selective_balanced_accuracy': 0.6397335931268693,
 'selective_macro_f1': 0.6126982557752435}
In [28]:
timestamp = datetime.now().strftime('%Y-%m-%d_%H%M%S')
pd.DataFrame([cc_result]).to_csv(SUPPORT_DIR / f'{timestamp}_class_conditional_gate_result.csv', index=False)
cc_cases_df.to_csv(SUPPORT_DIR / f'{timestamp}_class_conditional_gate_cases.csv', index=False)
pd.DataFrame([class_thresholds]).to_csv(SUPPORT_DIR / f'{timestamp}_class_conditional_gate_thresholds.csv', index=False)
print(timestamp)
2026-08-04_210605