Skip to content

Latest commit

Β 

History

3 Commits

Folders and files

NameName
Last commit message
Last commit date
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 

Repository files navigation

InfoSteer: Steering Information Utility in Key-Value Memory for Language Model Post-Training

License: MIT Python 3.8+ PyTorch

[NeurIPS 2025] Official implementation of "Steering Information Utility in Key-Value Memory for Language Model Post-Training"

πŸ” Overview

InfoSteer introduces novel techniques for steering information utility in transformer-based language models during post-training. The framework focuses on manipulating key-value memory mechanisms through:

  1. Entropy Regularization: Controlling activation entropy in MLP layers to optimize information flow
  2. Weight Projection Analysis: Clustering and analyzing down-projection weights to understand model behavior
  3. Activation Resuscitation: Reviving dormant neurons by modifying low-activation values

πŸš€ Key Features

  • Multiple Training Modes: Standard, entropy-regularized, and resuscitation-based training
  • Comprehensive Evaluation: Support for mathematical reasoning (GSM8K, AddSub, MAWPS, etc.) and other NLP tasks
  • Weight Analysis Tools: Clustering and visualization of model weights
  • Scalable Training: DeepSpeed integration for large model training
  • Flexible Architecture: Support for various transformer models (Qwen, Llama, etc.)

πŸ“‹ Requirements

  • Python 3.8+
  • PyTorch 2.0+
  • Transformers 4.30+
  • DeepSpeed (for large models)
  • vLLM (for efficient inference)
  • scikit-learn
  • matplotlib
  • wandb (for experiment tracking)

πŸ› οΈ Installation

  1. Clone the repository:
git clone https://github.com/your-username/InfoSteer.git
cd InfoSteer
  1. Install dependencies:
pip install torch transformers datasets deepspeed vllm scikit-learn matplotlib wandb tqdm numpy
  1. Set up Weights & Biases (optional but recommended):
wandb login

πŸ“– Usage

Training Models

1. Standard Training

python train.py \
    --model_name "Qwen/Qwen2.5-0.5B-Instruct" \
    --task "gsm8k" \
    --trainer_type "standard" \
    --batch_size 4 \
    --learning_rate 5e-5 \
    --output_dir "./models/standard"

2. Entropy Regularization Training

python train.py \
    --model_name "Qwen/Qwen2.5-0.5B-Instruct" \
    --task "gsm8k" \
    --trainer_type "entropy" \
    --entropy_weight 0.01 \
    --layer_group "all" \
    --batch_size 4 \
    --learning_rate 5e-5 \
    --output_dir "./models/entropy"

Layer Group Options:

  • early: Apply regularization to early layers only
  • intermediate: Apply to middle layers
  • final: Apply to final layers
  • all: Apply to all MLP layers
  • contrastive: Maximize final layer entropy while minimizing early layer entropy

3. Resuscitation Training

python train.py \
    --model_name "Qwen/Qwen2.5-0.5B-Instruct" \
    --task "gsm8k" \
    --trainer_type "resuscitation" \
    --resuscitation_percentile 1.0 \
    --batch_size 4 \
    --learning_rate 5e-5 \
    --output_dir "./models/resuscitation"

4. Large Model Training with DeepSpeed

python train.py \
    --model_name "Qwen/Qwen2.5-7B-Instruct" \
    --task "gsm8k" \
    --trainer_type "entropy" \
    --use_deepspeed \
    --gradient_checkpointing \
    --batch_size 2 \
    --output_dir "./models/large"

Evaluation

Evaluate trained models on mathematical reasoning tasks:

python eval.py \
    --model_path "./models/entropy/Qwen2.5-0.5B-Instruct_gsm8k_entropy0.01_alllayers" \
    --dataset "gsm8k" \
    --batch_size 8 \
    --output_file "./results/evaluation_results.json"

The evaluation script automatically tests on multiple math datasets: GSM8K, AddSub, MAWPS, MultiArith, SingleEq, and SVAMP.

Weight Analysis

1. Extract Model Weights

cd weight_projection
python fetch_weights.py

This will extract down-projection weights from models in the predefined pool and save them locally.

2. Cluster and Analyze Weights

python clustering.py

This performs k-means clustering on the extracted weights and generates:

  • Cluster labels for each layer
  • Visualization plots showing clustering results
  • JSONL files with detailed clustering information

πŸ“Š Supported Tasks and Datasets

Mathematical Reasoning

  • GSM8K: Grade school math word problems
  • AddSub: Addition and subtraction problems
  • MAWPS: Math word problems from various sources
  • MultiArith: Multi-step arithmetic problems
  • SingleEq: Single equation problems
  • SVAMP: Simple variations on arithmetic math word problems

Other Tasks

  • SQuAD: Reading comprehension
  • MMLU: Massive multitask language understanding
  • HotpotQA: Multi-hop reasoning
  • CommonsenseQA: Commonsense reasoning

🧠 Technical Details

Entropy Regularization

The entropy regularization technique calculates activation entropy in MLP layers:

def calculate_entropy(activations):
    probs = torch.abs(activations) / (torch.sum(torch.abs(activations), dim=-1, keepdim=True) + 1e-10)
    entropy = -torch.sum(probs * torch.log(probs + 1e-10), dim=-1)
    return entropy.mean()

The total loss combines standard language modeling loss with entropy regularization:

L_total = L_standard + Ξ» * L_entropy

Activation Resuscitation

The resuscitation method identifies and modifies the lowest-activation neurons:

  1. Find the k% smallest absolute activation values
  2. Replace them with 2Γ— the average activation value
  3. This helps "wake up" dormant neurons during training

Weight Clustering

The weight analysis module:

  1. Extracts down-projection weights from transformer MLP layers
  2. Applies k-means clustering to identify weight patterns
  3. Visualizes clusters using PCA dimensionality reduction
  4. Saves results in JSONL format for further analysis

πŸ“ Project Structure

InfoSteer/
β”œβ”€β”€ train.py                 # Main training script
β”œβ”€β”€ eval.py                  # Evaluation script  
β”œβ”€β”€ compute_metrics.py       # Metrics computation utilities
β”œβ”€β”€ ds_config.json          # DeepSpeed configuration
β”œβ”€β”€ weight_projection/      # Weight analysis module
β”‚   β”œβ”€β”€ fetch_weights.py    # Extract model weights
β”‚   └── clustering.py       # Cluster and visualize weights
β”œβ”€β”€ LICENSE                 # MIT License
└── README.md              # This file

πŸ”§ Configuration

DeepSpeed Configuration (ds_config.json)

  • Zero Stage 2 optimization
  • CPU offloading for optimizer states
  • BF16 precision training
  • Gradient accumulation and clipping

Training Arguments

  • Automatic batch size and learning rate scaling
  • Gradient checkpointing for memory efficiency
  • Evaluation every 100 steps
  • Model checkpointing with best model selection

πŸ“ˆ Results and Metrics

The framework tracks various metrics:

  • Standard Loss: Traditional language modeling loss
  • Entropy Loss: Activation entropy regularization term
  • Accuracy: Task-specific accuracy metrics
  • Resuscitation Stats: Number of modified activations

Results are logged to Weights & Biases and saved locally in JSON format.

🀝 Contributing

We welcome contributions! Please feel free to submit issues, feature requests, or pull requests.

πŸ“„ Citation

If you use InfoSteer in your research, please cite our paper:

@misc{deng2025infosteersteeringinformationutility,
      title={InfoSteer: Steering Information Utility in Language Model Post-Training}, 
      author={Chunyuan Deng and Ruidi Chang and Hanjie Chen},
      year={2025},
      eprint={2507.05158},
      archivePrefix={arXiv},
      primaryClass={cs.CL},
      url={https://arxiv.org/abs/2507.05158}, 
}

πŸ“œ License

This project is licensed under the MIT License - see the LICENSE file for details.

πŸ™ Acknowledgments

  • Built with πŸ€— Transformers and PyTorch
  • DeepSpeed for efficient large model training
  • vLLM for fast inference
  • Weights & Biases for experiment tracking

For questions or issues, please open a GitHub issue or contact the authors.

About

[NeurIPS 2025] Steering Information Utility in Key-Value Memory for Language Model Post-Training

Resources

Stars

4 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages