D-SCANは、RAG(Retrieval-Augmented Generation)システムにおける文書ポイズニング攻撃を検出するための分析フレームワークです。生成中にLLMの内部状態(トークン確率と注意重み)を収集し、多次元特徴量を抽出し、クリーンな検索文書とポイズニングされた検索文書を区別する分類器を訓練します。
このプロジェクトはデフォルトで Llama-3.1-8B-Instruct を使用します。モデルをローカルパスにダウンロードし、collect_inner_state.py 内の MODEL_ID 変数を更新してください:
MODEL_ID = "/your/path/to/Llama-3.1-8B-Instruct"
ワークフロー全体は3つのステップで構成されています:内部状態の収集 → 特徴量の計算 → 分類器の訓練。 質問と関連する検索文書は https://huggingface.co/datasets/An998/D-SCAN で提供されています。
# Update configuration parameters in collect_inner_state.py, then run
python collect_inner_state.py
# Analyze both attack and clean data, compute all features, and save results
python compute_feature.py \
--attack_dir ./saved_reppl_weights_serial_query1_attack2_2wiki \
--clean_dir ./saved_reppl_weights_serial_query1_pure1_5_2wiki \
--output_dir ./analysis_results/2wiki_all \
--max_attack_samples 3000 \
--max_clean_samples 3000
fit_D-SCAN.ipynb を開き、セルを順番に実行して以下を行います:
collect_inner_state.py)各質問-文書入力に対して、LLMはマルチサンプル生成(デフォルト:10サンプル、temperature=1.0)を実行し、以下の内部状態を収集します:
| フィールド | 型 | 説明 |
|---|---|---|
outer_ppl_probs | List[Tensor] | 各サンプルのトークン生成確率 |
inner_ppl_matrix | List[List[Tensor]] | 各生成ステップにおける入力シーケンスに対する注意重み(層平均) |
doc_ranges | Dict[str, List[int]] | 入力シーケンス内の各文書のトークン位置範囲 |
generated_sequences | List[List[int]] | 各サンプルの生成トークンIDシーケンス |
MODEL_ID = "/path/to/model" # Model path
NUM_SAMPLES = 10 # Number of samples per question
MAX_NEW_TOKENS = 50 # Maximum generated tokens
TEMPERATURE = 1.0 # Sampling temperature
data_type = 'pure1_5' # Data type: 'pure1_5' (clean) or 'attack2' (attack)
dataset_name = 'hotpotqa' # Dataset: 'hotpotqa', '2wiki', 'musique'
出力は data_{batch_id}_reppl.pt ファイルと results_stats_*.jsonl 統計ファイルとして保存されます。
compute_feature.py)収集された内部状態から10カテゴリ、100次元以上の特徴量を抽出します。分類器はこれらの特徴量を用いて、特定のクエリに対して検索された文書の中にポイズニングされた文書が存在するかどうかを判定します。
| # | カテゴリ | クラス名 | 特徴量数 | 核心的な考え方 |
|---|---|---|---|---|
| 1 | 生成確率統計 | PerplexityMetrics | 8 | ポイズニングされた文書は生成中のモデルの不確実性を増加させ、確率分布の変化に反映される可能性がある |
| 2 | 注意エントロピー | AttentionEntropyMetrics | 4 | 高エントロピー=注意の分散=矛盾する情報の可能性;低エントロピー=注意の集中 |
| 3 | 注意集中度 | AttentionConcentrationMetrics | 8 | Top-K比率とジニ係数により、注意が少数のトークンに集中しているかを測定 |
| 4 | 文書注意密度 | DocumentAttentionDensityMetrics | 6 | 注意の合計を文書長で割ることで、注意配分における長さバイアスを排除 |
| 5 | マルチサンプル一貫性 | SampleConsistencyMetrics | 4 | ポイズニング下では、サンプル間の注意パターンが一貫しない可能性がある(コサイン類似度、JSダイバージェンス) |
| 6 | 注意ダイナミクス | AttentionDynamicsMetrics | 4 | 生成中の支配的文書の切り替え頻度とエントロピー変化量 |
| 7 | トークンレベル注意変動 | TokenLevelAttentionMetrics | 16 | 生成ステップ間における各入力トークン位置の注意の安定性(標準偏差、エントロピー) |
| 8 | 回答確率詳細統計 | AnswerProbabilityMetrics | 18 | 確率分位点、高/低確率トークン比率、対数確率、パープレキシティなど |
| 9 | 確率ダイナミクス | ProbabilityDynamicsMetrics | 14 | 生成シーケンスにおけるトレンド傾き、自己相関、ボラティリティ、スパイク比率 |
| 10 | サンプル間確率一貫性 | CrossSampleProbabilityConsistencyMetrics | 13 | サンプル間のコサイン類似度、ピアソン相関、MSE、ダイバージェンス指数 |
1. PerplexityMetrics(生成確率統計)
outer_ppl_probs(各生成トークンの確率)から計算:
ppl_mean_prob:全生成トークンの平均確率ppl_std_prob:確率の標準偏差ppl_min_prob:最小確率(極度の不確実性)ppl_low_prob_ratio:低確率トークン(<0.1)の比率ppl_cross_sample_var:サンプル間の平均確率の分散ppl_coef_variation:変動係数(CV = std/mean)ppl_skewness:確率分布の歪度ppl_kurtosis:確率分布の尖度2. AttentionEntropyMetrics(注意エントロピー)
inner_ppl_matrix(各ステップの注意重み)から計算:
attn_entropy_mean/std/max:注意分布エントロピーの平均、標準偏差、最大値attn_entropy_cv:注意エントロピーの変動係数3. AttentionConcentrationMetrics(注意集中度)
attn_top5/10/20_ratio_mean/std:Top-K%のトークンが占める注意の割合attn_gini_mean/std:注意分布のジニ係数(不平等度の指標)4. DocumentAttentionDensityMetrics(文書注意密度)
doc_attn_dens_std/range/max/min:文書間の注意密度の標準偏差、範囲、最大値、最小値doc_attn_dens_entropy:文書注意密度分布のエントロピーdoc_attn_dens_temporal_var_mean:文書注意密度の時間分散の平均5. SampleConsistencyMetrics(マルチサンプル一貫性)
sample_attn_consistency/std:トークンレベル注意のサンプル間コサイン類似度sample_doc_consistency:文書レベル注意のサンプル間コサイン類似度sample_doc_js_divergence:文書レベル注意のサンプル間JSダイバージェンス6. AttentionDynamicsMetrics(注意ダイナミクス)
attn_doc_switch_mean/max:支配的文書の切り替え回数attn_entropy_change_mean/std:ステップごとの注意エントロピー変化7. TokenLevelAttentionMetrics(トークンレベル注意変動)
tla_std_mean/std/max/median/p90/p99/high_ratio/cv:生成ステップ間におけるトークンごとの注意標準偏差の統計量tla_ent_mean/std/max/median/p90/p99/high_ratio/cv:生成ステップ間におけるトークンごとの注意エントロピーの統計量8. AnswerProbabilityMetrics(回答確率詳細統計)
aprob_p10/p25/p50/p75/p90/iqr:確率分位点aprob_high_ratio_05/08:高確率トークンの比率aprob_low_ratio_01/001:低確率トークンの比率aprob_log_mean/std/min:対数確率の統計量aprob_ppl_mean/std/max:シーケンスレベルのパープレキシティaprob_geometric_mean:確率の幾何平均aprob_distribution_entropy:確率ヒストグラムの情報エントロピー9. ProbabilityDynamicsMetrics(確率ダイナミクス)
pdyn_diff_mean/abs_diff_mean/abs_diff_std/abs_diff_max:確率差の統計量pdyn_max_drop/max_jump:単一ステップの最大下落/上昇pdyn_volatility_mean/std:ボラティリティpdyn_trend_slope_mean/std:線形トレンドの傾きpdyn_autocorr_mean/std:自己相関係数pdyn_spike_ratio_01/03:スパイク(変異点)比率10. CrossSampleProbabilityConsistencyMetrics(サンプル間確率一貫性)
cspc_mean_prob_std/cv/range:サンプル間の平均確率の一貫性cspc_ppl_std/cv/range:サンプル間のパープレキシティの一貫性cspc_min_prob_std/range:最小確率の一貫性cspc_seq_cosine_mean/std:確率シーケンスのサンプル間コサイン類似度cspc_seq_pearson_mean:確率シーケンスのサンプル間ピアソン相関cspc_seq_mse_mean:確率シーケンスのサンプル間MSEcspc_divergence_index:サンプル間ダイバージェンス指数compute_feature.py のコマンドライン引数python compute_feature.py \
--attack_dir <attack_data_directory> \
--clean_dir <clean_data_directory> \
--output_dir <output_directory> \
--max_attack_samples 3000 \
--max_clean_samples 3000 \
--min_correct_count 0 \
--min_attack_target_count 0 \
--num_use_samples 10 \
--model_path /path/to/model # Optional: load tokenizer for accuracy calculation