15 Risk-Aware Adaptive Agent¶

这一份 notebook 开始做真正的 病例级风险感知提问。

前面我们已经验证过两件事:

  1. 单纯在最后加安全门,会越来越保守
  2. 单纯把问题顺序固定成 age -> ...,效果也有限

所以这一版的核心变化是:

不同病例,根据当前状态,动态决定下一步先问 age / sex / location 中的哪一个。

这一版先做 第一代可解释版本:

  • 看当前已经知道哪些信息
  • 看还剩哪些问题没问
  • 用一个简单的风险感知打分,给每个候选问题打分
  • 选分数最高的那个问题

这一版不是最终最强方法,但它会把我们从 fixed-order 正式推进到 adaptive questioning。

In [3]:
from pathlib import Path
from datetime import datetime
import ast
import json

import numpy as np
import pandas as pd

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'
QUESTION_SCORE_PATH = SUPPORT_DIR / '2026-08-04_163052_risk_aware_score_table.csv'

print(BASELINE_RESULTS_PATH)
print(MODEL_REGISTRY_PATH)
print(QUESTION_SCORE_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/2026-08-04_163052_risk_aware_score_table.csv

1. 读取数据与模型分数¶

这里继续沿用我们之前 validated agent 的思路:

  • agent 做决策时,用 validation 分数
  • 最终汇报时,用 test 分数
In [4]:
train_df = pd.read_csv(SPLIT_DIR / 'train.csv')
val_df = pd.read_csv(SPLIT_DIR / 'val.csv')
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)
question_score_df = pd.read_csv(QUESTION_SCORE_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']),
    }

len(train_df), len(val_df), len(test_df)
Out[4]:
(7002, 1532, 1481)

2. 读取我们已经统计好的问题风险分数¶

这一步就把前面 13 号 notebook 跑出来的全局风险收益统计接进来。

In [5]:
question_score_df
Out[5]:
question cases improved_rate worsened_rate unsafe_rate final_correct_rate avg_confidence_gain risk_aware_score
0 age 122 0.188525 0.081967 0.278689 0.639344 0.288338 -0.450820
1 sex 344 0.119186 0.127907 0.334302 0.537791 0.324628 -0.677326
2 location 81 0.123457 0.123457 0.407407 0.469136 0.265222 -0.814815
In [6]:
global_question_score = {
    row['question']: float(row['risk_aware_score'])
    for _, row in question_score_df.iterrows()
}

global_question_score
Out[6]:
{'age': -0.4508196721311475,
 'sex': -0.6773255813953488,
 'location': -0.8148148148148148}

3. 状态与工具函数¶

这一部分继续沿用前面的 agent 状态定义。

In [7]:
def build_initial_state(row):
    return {
        'image_id': row['image_id'],
        'true_label': row['dx'],
        'known_metadata': {},
        'asked_questions': [],
        'done': False,
    }


def ask_question(state, row, question):
    assert question in ['age', 'sex', 'location']
    new_state = {
        'image_id': state['image_id'],
        'true_label': state['true_label'],
        'known_metadata': dict(state['known_metadata']),
        'asked_questions': list(state['asked_questions']),
        'done': state['done'],
    }

    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']

4. 第一代 risk-aware adaptive policy¶

这一版 policy 很简单,但已经是“病例级 adaptive”了:

基础原则¶

  • 还没问过的问题,才有资格继续选
  • 问题分数 = 当前状态下的潜在收益 + 全局风险感知先验

当前状态下的潜在收益¶

如果问某个问题后,对应模型的 validation Macro-F1 变高,说明它对当前状态是有帮助的。

全局风险感知先验¶

如果某个问题在全局统计上更危险,就对它做惩罚。

In [8]:
def risk_aware_question_score(state, question, alpha=1.0):
    current_known = set(state['known_metadata'].keys())
    if question in current_known:
        return -1e9

    current_key = get_model_key_from_known_set(current_known)
    future_key = get_model_key_from_known_set(current_known | {question})

    local_gain = policy_model_scores[future_key]['macro_f1'] - policy_model_scores[current_key]['macro_f1']
    risk_prior = global_question_score[question]

    return local_gain + alpha * risk_prior


def risk_aware_adaptive_policy(state, max_questions=2, alpha=1.0):
    if len(state['asked_questions']) >= max_questions:
        return 'diagnose'

    candidates = [q for q in ['age', 'sex', 'location'] if q not in state['asked_questions']]
    if not candidates:
        return 'diagnose'

    best_question = None
    best_score = -1e18

    for q in candidates:
        s = risk_aware_question_score(state, q, alpha=alpha)
        if s > best_score:
            best_score = s
            best_question = q

    if best_question is None:
        return 'diagnose'
    return best_question

5. 先看单个病例轨迹¶

先不要急着全量跑,我们先看看 agent 在一个病例上会怎么决策。

In [9]:
def run_risk_aware_episode(row, max_questions=2, alpha=1.0):
    state = build_initial_state(row)
    trajectory = []

    while not state['done']:
        action = risk_aware_adaptive_policy(state, max_questions=max_questions, alpha=alpha)

        trajectory.append({
            'step': len(trajectory),
            'action': action,
            '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 == 'diagnose':
            state['done'] = True
            state['final_action'] = 'diagnose'
            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 [10]:
sample_row = test_df.iloc[0]
trajectory, final_state = run_risk_aware_episode(sample_row, max_questions=2, alpha=1.0)

trajectory, final_state
Out[10]:
([{'step': 0,
   'action': 'age',
   'model_key_before_action': 'image_only',
   'score_before_action': 0.6083234281767498,
   'known_metadata_before_action': {}},
  {'step': 1,
   'action': 'sex',
   'model_key_before_action': 'image_age',
   'score_before_action': 0.5926280556025649,
   'known_metadata_before_action': {'age': np.float64(70.0)}},
  {'step': 2,
   'action': 'diagnose',
   'model_key_before_action': 'image_age_sex',
   'score_before_action': 0.5964635518225668,
   'known_metadata_before_action': {'age': np.float64(70.0),
    'sex': 'female'}}],
 {'image_id': 'ISIC_0025837',
  'true_label': 'bkl',
  'known_metadata': {'age': np.float64(70.0), 'sex': 'female'},
  'asked_questions': ['age', 'sex'],
  'done': True,
  'final_action': 'diagnose',
  'final_model_key': 'image_age_sex',
  'final_score': 0.5964635518225668})

6. 全测试集跑第一版 risk-aware adaptive agent¶

In [11]:
def evaluate_risk_aware_agent(alpha=1.0, max_questions=2):
    records = []
    for idx in range(len(test_df)):
        row = test_df.iloc[idx]
        trajectory, final_state = run_risk_aware_episode(
            row,
            max_questions=max_questions,
            alpha=alpha,
        )
        records.append({
            'image_id': row['image_id'],
            'true_label': row['dx'],
            'asked_questions': list(final_state['asked_questions']),
            'num_questions': len(final_state['asked_questions']),
            'final_model_key': final_state['final_model_key'],
            'final_policy_score': final_state['final_score'],
            'known_metadata': dict(final_state['known_metadata']),
        })

    cases_df = pd.DataFrame(records)
    summary_df = cases_df.groupby('final_model_key').size().reset_index(name='count')
    summary_df['ratio'] = summary_df['count'] / len(cases_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'])

    result = {
        'agent_name': 'risk_aware_adaptive',
        'alpha': alpha,
        'max_questions': max_questions,
        'avg_questions': float(cases_df['num_questions'].mean()),
        'expected_accuracy': float((summary_df['ratio'] * summary_df['test_accuracy']).sum()),
        'expected_balanced_accuracy': float((summary_df['ratio'] * summary_df['test_balanced_accuracy']).sum()),
        'expected_macro_f1': float((summary_df['ratio'] * summary_df['test_macro_f1']).sum()),
    }
    return cases_df, summary_df, result
In [12]:
risk_cases_df, risk_summary_df, risk_result = evaluate_risk_aware_agent(alpha=1.0, max_questions=2)
risk_result
Out[12]:
{'agent_name': 'risk_aware_adaptive',
 'alpha': 1.0,
 'max_questions': 2,
 'avg_questions': 2.0,
 'expected_accuracy': 0.7886563133018231,
 'expected_balanced_accuracy': 0.5856809903171467,
 'expected_macro_f1': 0.5837050128198616}

7. 和之前 baseline agent 对比¶

这里先和我们已经有的 3 个 baseline 比:

  • fixed_order
  • uncertainty
  • lookahead
In [13]:
validated_agent_df = pd.read_csv(SUPPORT_DIR / '2026-08-03_185733_validated_agent_comparison.csv')
comparison_df = pd.concat([
    validated_agent_df,
    pd.DataFrame([risk_result])
], ignore_index=True)
comparison_df
Out[13]:
agent_name max_questions avg_questions expected_accuracy expected_balanced_accuracy expected_macro_f1 threshold alpha
0 fixed_order 2 2.0 0.793383 0.582281 0.581666 NaN NaN
1 uncertainty 2 2.0 0.793383 0.582281 0.581666 0.605 NaN
2 lookahead 2 2.0 0.792708 0.598458 0.596497 NaN NaN
3 risk_aware_adaptive 2 2.0 0.788656 0.585681 0.583705 NaN 1.0

8. 保存结果¶

In [14]:
timestamp = datetime.now().strftime('%Y-%m-%d_%H%M%S')

cases_path = SUPPORT_DIR / f'{timestamp}_risk_aware_adaptive_cases.csv'
summary_path = SUPPORT_DIR / f'{timestamp}_risk_aware_adaptive_summary.csv'
result_path = SUPPORT_DIR / f'{timestamp}_risk_aware_adaptive_result.json'
comparison_path = SUPPORT_DIR / f'{timestamp}_risk_aware_adaptive_comparison.csv'

risk_cases_df.to_csv(cases_path, index=False)
risk_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(risk_result, f, ensure_ascii=False, indent=2)

print(cases_path)
print(summary_path)
print(result_path)
print(comparison_path)
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_191116_risk_aware_adaptive_cases.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_191116_risk_aware_adaptive_summary.csv
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_191116_risk_aware_adaptive_result.json
/Users/applesues01/Documents/Medical_Agent/supports/2026-08-04_191116_risk_aware_adaptive_comparison.csv

9. 这一版怎么看¶

如果这一版能超过 fixed_order,说明:

风险感知 + 病例级动态提问 是有希望的。

如果连 fixed_order 都超不过,也没关系,这至少说明:

  • 仅靠全局风险分数还不够
  • 下一步必须引入更细的病例级类别信息,也就是更接近真正的信息增益策略

个人感觉跑了一版没啥用的代码¶