15 Risk-Aware Adaptive Agent¶
这一份 notebook 开始做真正的 病例级风险感知提问。
前面我们已经验证过两件事:
- 单纯在最后加安全门,会越来越保守
- 单纯把问题顺序固定成
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
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']
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}
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 都超不过,也没关系,这至少说明:
- 仅靠全局风险分数还不够
- 下一步必须引入更细的病例级类别信息,也就是更接近真正的信息增益策略