Skip to content

Repository files navigation

When Context Bites: Detecting RAG Poisoning via Document-Level Attention Collapse

D-SCAN is an analysis framework for detecting document poisoning attacks in RAG (Retrieval-Augmented Generation) systems. It collects LLM internal states during generation (token probabilities and attention weights), extracts multi-dimensional features, and trains classifiers to distinguish between clean and poisoned retrieved documents.

1. Installation & Requirements

Prerequisites

  • Python >= 3.9
  • CUDA-compatible GPU (>= 24GB VRAM recommended for 8B-parameter model inference)

Model Preparation

This project uses Llama-3.1-8B-Instruct by default. Download the model to a local path and update the MODEL_ID variable in collect_inner_state.py:

MODEL_ID = "/your/path/to/Llama-3.1-8B-Instruct"

2. Quick Start

The full workflow consists of three steps: Collect Internal States → Compute Features → Train Classifier. The question and related retrieved documents are provided in https://huggingface.co/datasets/An998/D-SCAN.

Step 1: Collect Model Internal States

# Update configuration parameters in collect_inner_state.py, then run
python collect_inner_state.py

Step 2: Compute Feature Metrics

# 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

Step 3: Train Classifier and Evaluate

Open fit_D-SCAN.ipynb and execute cells sequentially to:

  1. Load the feature data generated in Step 2
  2. Train Logistic Regression / Random Forest classifiers (5-Fold CV)
  3. Evaluate transfer performance on the test set

3. Pipeline Details

3.1 Data Collection (collect_inner_state.py)

For each question-document input, the LLM performs multi-sample generation (default: 10 samples, temperature=1.0) and collects the following internal states:

Field Type Description
outer_ppl_probs List[Tensor] Token generation probability for each sample
inner_ppl_matrix List[List[Tensor]] Attention weights over input sequence at each generation step (layer-averaged)
doc_ranges Dict[str, List[int]] Token position range for each document in the input sequence
generated_sequences List[List[int]] Generated token ID sequences for each sample

Key Configuration Parameters

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'

Outputs are saved as data_{batch_id}_reppl.pt files and results_stats_*.jsonl statistics files.


3.2 Feature Computation (compute_feature.py)

Extracts 10 categories with 100+ dimensional features from the collected internal states. The classifier uses these features to determine whether poisoned documents exist among the retrieved documents for a given query.

Metric Overview

# Category Class Name # Features Core Idea
1 Generation Probability Stats PerplexityMetrics 8 Poisoned docs may increase model uncertainty during generation, reflected in probability distribution changes
2 Attention Entropy AttentionEntropyMetrics 4 High entropy = dispersed attention = possible conflicting information; Low entropy = focused attention
3 Attention Concentration AttentionConcentrationMetrics 8 Measures whether attention is concentrated on a few tokens via Top-K ratio and Gini coefficient
4 Document Attention Density DocumentAttentionDensityMetrics 6 Attention sum divided by document length, eliminating length bias in attention allocation
5 Multi-Sample Consistency SampleConsistencyMetrics 4 Under poisoning, attention patterns across samples may be inconsistent (cosine similarity, JS divergence)
6 Attention Dynamics AttentionDynamicsMetrics 4 Frequency of dominant-document switches and entropy change magnitude during generation
7 Token-Level Attention Fluctuation TokenLevelAttentionMetrics 16 Stability of attention at each input token position across generation steps (std, entropy)
8 Answer Probability Deep Stats AnswerProbabilityMetrics 18 Probability quantiles, high/low probability token ratios, log probability, perplexity, etc.
9 Probability Dynamics ProbabilityDynamicsMetrics 14 Trend slope, autocorrelation, volatility, spike ratio in the generation sequence
10 Cross-Sample Probability Consistency CrossSampleProbabilityConsistencyMetrics 13 Cosine similarity, Pearson correlation, MSE, divergence index across samples

Detailed Metric Descriptions

1. PerplexityMetrics (Generation Probability Statistics)

Computed from outer_ppl_probs (probability of each generated token):

  • ppl_mean_prob: Mean probability across all generated tokens
  • ppl_std_prob: Probability standard deviation
  • ppl_min_prob: Minimum probability (extreme uncertainty)
  • ppl_low_prob_ratio: Ratio of low-probability tokens (<0.1)
  • ppl_cross_sample_var: Variance of mean probabilities across samples
  • ppl_coef_variation: Coefficient of variation (CV = std/mean)
  • ppl_skewness: Probability distribution skewness
  • ppl_kurtosis: Probability distribution kurtosis

2. AttentionEntropyMetrics (Attention Entropy)

Computed from inner_ppl_matrix (attention weights at each step):

  • attn_entropy_mean/std/max: Mean, standard deviation, and maximum of attention distribution entropy
  • attn_entropy_cv: Coefficient of variation of attention entropy

3. AttentionConcentrationMetrics (Attention Concentration)

  • attn_top5/10/20_ratio_mean/std: Attention share captured by the Top-K% tokens
  • attn_gini_mean/std: Gini coefficient of the attention distribution (inequality measure)

4. DocumentAttentionDensityMetrics (Document Attention Density)

  • doc_attn_dens_std/range/max/min: Std, range, max, and min of attention density across documents
  • doc_attn_dens_entropy: Entropy of document attention density distribution
  • doc_attn_dens_temporal_var_mean: Mean temporal variance of document attention density

5. SampleConsistencyMetrics (Multi-Sample Consistency)

  • sample_attn_consistency/std: Cross-sample cosine similarity of token-level attention
  • sample_doc_consistency: Cross-sample cosine similarity of document-level attention
  • sample_doc_js_divergence: Cross-sample JS divergence of document-level attention

6. AttentionDynamicsMetrics (Attention Dynamics)

  • attn_doc_switch_mean/max: Dominant document switch count
  • attn_entropy_change_mean/std: Step-wise attention entropy change

7. TokenLevelAttentionMetrics (Token-Level Attention Fluctuation)

  • tla_std_mean/std/max/median/p90/p99/high_ratio/cv: Statistics of per-token attention std across generation steps
  • tla_ent_mean/std/max/median/p90/p99/high_ratio/cv: Statistics of per-token attention entropy across generation steps

8. AnswerProbabilityMetrics (Answer Probability Deep Statistics)

  • aprob_p10/p25/p50/p75/p90/iqr: Probability quantiles
  • aprob_high_ratio_05/08: Ratio of high-probability tokens
  • aprob_low_ratio_01/001: Ratio of low-probability tokens
  • aprob_log_mean/std/min: Log-probability statistics
  • aprob_ppl_mean/std/max: Sequence-level perplexity
  • aprob_geometric_mean: Geometric mean of probabilities
  • aprob_distribution_entropy: Information entropy of the probability histogram

9. ProbabilityDynamicsMetrics (Probability Dynamics)

  • pdyn_diff_mean/abs_diff_mean/abs_diff_std/abs_diff_max: Probability difference statistics
  • pdyn_max_drop/max_jump: Maximum single-step drop/jump
  • pdyn_volatility_mean/std: Volatility
  • pdyn_trend_slope_mean/std: Linear trend slope
  • pdyn_autocorr_mean/std: Autocorrelation coefficient
  • pdyn_spike_ratio_01/03: Spike (mutation point) ratio

10. CrossSampleProbabilityConsistencyMetrics (Cross-Sample Probability Consistency)

  • cspc_mean_prob_std/cv/range: Consistency of mean probabilities across samples
  • cspc_ppl_std/cv/range: Consistency of perplexity across samples
  • cspc_min_prob_std/range: Consistency of minimum probabilities
  • cspc_seq_cosine_mean/std: Cross-sample cosine similarity of probability sequences
  • cspc_seq_pearson_mean: Cross-sample Pearson correlation of probability sequences
  • cspc_seq_mse_mean: Cross-sample MSE of probability sequences
  • cspc_divergence_index: Inter-sample divergence index

compute_feature.py Command-Line Arguments

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
Argument Default Description
--attack_dir - Attack data directory (containing data_*_reppl.pt files)
--clean_dir - Clean data directory
--output_dir - Output directory
--max_attack_samples 3000 Maximum number of attack samples
--max_clean_samples 3000 Maximum number of clean samples
--min_correct_count 0 Minimum correct answer count in clean data (for filtering)
--min_attack_target_count 0 Minimum target answer hit count in attack data
--num_use_samples None Number of samples per question for metric computation (default: all)
--model_path None Model path (for loading tokenizer to decode generated sequences)

Output Files

Filename Description
single_metric_analysis.json AUC, p-value, Cohen's d, etc. for each metric
full_analysis_results.json Complete feature matrix + labels
document_detailed_metrics.json Per-document detailed metrics (JSON summary)
document_detailed_metrics_full.pkl Full document-level metrics (incl. step-level attention)

3.3 Classifier Training (fit_D-SCAN.ipynb)

Notebook workflow:

  1. Load Data: Read full_analysis_results.json output from compute_feature.py
  2. Feature Selection: Use all features or filter by category (e.g., attention-only features)
  3. Train Classifiers:
    • Logistic Regression (L2 regularization): Linear classifier, suitable for fewer features
    • Random Forest (100 trees, max_depth=5): Non-linear classifier, captures feature interactions
  4. 5-Fold Cross-Validation: Reports AUC, Accuracy, Precision, Recall, F1
  5. Feature Importance Analysis: Outputs Top-10 important features
  6. Test Set Evaluation: Uses training scaler and classifier to evaluate transfer performance on an independent test set
  7. Visualization: ROC curves, feature importance bar charts, Train vs Test comparison

Feature Selection for Classifiers

The use_features variable in the notebook provides flexible control over the feature subset used by classifiers:

# Use all features
use_features = [f for f in feature_names if f not in exclude_cols]

# Use only attention-related features
use_features = [f for f in use_features if 'attn' in f]

# Combine by metric category (example)
use_features = [f for f in feature_names if f.startswith(('ppl_', 'doc_', 'sample_'))]

Per-category classifier performance is also printed during compute_feature.py execution:

Feature Group Prefix / Keyword
perplexity ppl_*
attention_entropy *entropy* (excl. doc and aprob)
attention_concentration *top*, *gini*
document_attention doc_*
sample_consistency *sample*, *consistency*
attention_dynamics *switch*, *change*
token_level_attention tla_*
answer_probability aprob_*
probability_dynamics pdyn_*
cross_sample_prob_consistency cspc_*

About

No description, website, or topics provided.

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages