[NeurIPS 2025] Official implementation of "Steering Information Utility in Key-Value Memory for Language Model Post-Training"
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:
- Entropy Regularization: Controlling activation entropy in MLP layers to optimize information flow
- Weight Projection Analysis: Clustering and analyzing down-projection weights to understand model behavior
- Activation Resuscitation: Reviving dormant neurons by modifying low-activation values
- 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.)
- Python 3.8+
- PyTorch 2.0+
- Transformers 4.30+
- DeepSpeed (for large models)
- vLLM (for efficient inference)
- scikit-learn
- matplotlib
- wandb (for experiment tracking)
- Clone the repository:
git clone https://github.com/your-username/InfoSteer.git
cd InfoSteer- Install dependencies:
pip install torch transformers datasets deepspeed vllm scikit-learn matplotlib wandb tqdm numpy- Set up Weights & Biases (optional but recommended):
wandb loginpython 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"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 onlyintermediate: Apply to middle layersfinal: Apply to final layersall: Apply to all MLP layerscontrastive: Maximize final layer entropy while minimizing early layer entropy
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"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"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.
cd weight_projection
python fetch_weights.pyThis will extract down-projection weights from models in the predefined pool and save them locally.
python clustering.pyThis 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
- 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
- SQuAD: Reading comprehension
- MMLU: Massive multitask language understanding
- HotpotQA: Multi-hop reasoning
- CommonsenseQA: Commonsense reasoning
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
The resuscitation method identifies and modifies the lowest-activation neurons:
- Find the k% smallest absolute activation values
- Replace them with 2Γ the average activation value
- This helps "wake up" dormant neurons during training
The weight analysis module:
- Extracts down-projection weights from transformer MLP layers
- Applies k-means clustering to identify weight patterns
- Visualizes clusters using PCA dimensionality reduction
- Saves results in JSONL format for further analysis
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
- Zero Stage 2 optimization
- CPU offloading for optimizer states
- BF16 precision training
- Gradient accumulation and clipping
- Automatic batch size and learning rate scaling
- Gradient checkpointing for memory efficiency
- Evaluation every 100 steps
- Model checkpointing with best model selection
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.
We welcome contributions! Please feel free to submit issues, feature requests, or pull requests.
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},
}This project is licensed under the MIT License - see the LICENSE file for details.
- 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.