Saturation Mask Integration Guide
Overview
Saturation masks identify neurons in transformer FFN layers that produce near-zero activations during inference, enabling sparse computation and latency reduction.
Key metrics:
- Frozen neurons: Outputs < threshold (typically 0.01)
- Speedup estimate: 1% speedup per 10% frozen neurons (conservative lower bound)
- Serialization overhead: <1MB per model
- Deserialization latency: <1ms per 32 layers
Module Architecture
Encoder (saturation_mask_encoder.py)
Computes saturation masks from calibration data and embeds them in model files.
from saturation_mask_encoder import (
compute_saturation_masks,
encode_masks_to_safetensors,
)
# Load model and compute masks
model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2.5-7B")
tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen2.5-7B")
masks = compute_saturation_masks(
model=model,
tokenizer=tokenizer,
method="variance_gradient", # or "activity_counting", "statistical"
num_samples=512,
device="cuda",
)
# Embed in SafeTensors
encode_masks_to_safetensors(
model_path="model.safetensors",
masks=masks,
output_path="model_with_masks.safetensors"
)Detection methods:
variance_gradient: Combines neuron variance and temporal gradient flow. Recommended for detecting persistent frozen neurons.activity_counting: Counts near-zero activations (>80% near-zero = frozen). Fast, suitable for high-throughput scenarios.statistical: Quantile-based magnitude thresholding. Simple, works well for diverse architectures.
Decoder (saturation_mask_decoder.py)
Loads masks from model files at inference time.
from saturation_mask_decoder import load_saturation_masks, load_cliff_scores
# Load from SafeTensors or GGUF (auto-detected)
masks = load_saturation_masks("model_with_masks.safetensors")
for layer_id, mask in masks.items():
print(f"Layer {layer_id}: {mask.pct_frozen():.1f}% frozen")
# Check if neuron is frozen
if mask.is_frozen(neuron_idx):
# Skip computation for this neuron
pass
# Load cliff scores (saturation detection confidence, 0-1)
cliff_scores = load_cliff_scores("model_with_masks.safetensors")Integration Patterns
1. vLLM Integration
Apply masks in FFN inference loop:
# vLLM custom operator
class SparseFfnLayer(nn.Module):
def __init__(self, ffn_layer, frozen_mask):
super().__init__()
self.ffn = ffn_layer
self.mask = frozen_mask
def forward(self, x):
# Compute full output
output = self.ffn(x)
# Zero out frozen neurons
frozen_indices = self.mask.to_sparse_indices()
output[:, frozen_indices] = 0.0
return output
# In vLLM initialization
masks = load_saturation_masks(model_path)
for layer_id, mask in masks.items():
original_ffn = model.layers[layer_id].mlp
model.layers[layer_id].mlp = SparseFfnLayer(original_ffn, mask)Expected gains:
- Qwen2.5-7B: ~8-12% FFN latency reduction
- Llama-70B: ~5-10% reduction (sparse layers, more coordination needed)
- Phi-3-mini: ~15-20% reduction (smaller, higher saturation)
2. Aether Integration
Load masks into aether compute graph:
// In aether scheduler
use distributed_inference::saturation_masks::load_saturation_masks;
let masks = load_saturation_masks("model.safetensors")?;
for layer_id in 0..num_layers {
let mask = &masks[&(layer_id as u32)];
// Register frozen neurons with scheduler
scheduler.register_frozen_neurons(
layer_id,
mask.to_sparse_indices(),
);
}
// Scheduler uses indices to skip compute
let compute_nodes = scheduler.plan_ffn_forward(layer_id);
// compute_nodes excludes frozen neurons3. Gnosis-Uring Integration
Integrate with uring ring for distributed saturation handling:
// In gnosis-uring coordinator
use saturation_metadata::SaturationMaps;
let saturation = SaturationMaps::load_from_model(model_path)?;
// Register with uring ring
for boundary in &saturation.ffn_frozen_per_layer {
uring_ring.register_boundary(
BoundaryHint {
layer_id: boundary.layer_id,
frozen_mask: boundary.frozen_indices.clone(),
estimated_speedup: boundary.estimated_speedup(),
}
);
}4. llama.cpp Integration
Inject masks into GGML kernels:
// In llama.cpp model loading
struct saturation_info {
uint32_t num_layers;
uint32_t* frozen_neurons;
uint32_t* frozen_counts;
};
saturation_info sat = load_saturation_from_gguf(model_path);
// In matmul kernel (ggml_mul_mat)
for (int neuron = 0; neuron < num_out; neuron++) {
if (is_neuron_frozen(sat, layer_id, neuron)) {
// Skip output computation for this neuron
output[neuron] = 0.0;
continue;
}
output[neuron] = compute_neuron(input, weights, neuron);
}5. TGI (Text Generation Inference) Integration
Load masks and apply in FFN pipeline:
# In TGI custom forward pass
class SaturatedFFN(nn.Module):
def __init__(self, ffn, masks):
super().__init__()
self.ffn = ffn
self.masks = masks # Dict[layer_id] -> BitmaskArray
self.layer_id = None
def forward(self, x):
out = self.ffn(x)
# Apply mask if available
if self.layer_id in self.masks:
mask = self.masks[self.layer_id]
# Dense mask: shape (hidden_dim,)
out = out * mask.to_dense_mask()
return out
# Apply to model at init
masks = load_saturation_masks(model_path)
for i, layer in enumerate(model.transformer.h):
layer.mlp = SaturatedFFN(layer.mlp, masks)
layer.mlp.layer_id = iPerformance Expectations
Serialization Overhead
| Model | Layers | Avg Frozen % | Blob Size | Overhead |
|---|---|---|---|---|
| Phi-3-mini | 32 | 12% | 48 KB | 0.05% |
| Qwen2.5-7B | 28 | 10% | 35 KB | 0.02% |
| Gemma-9b | 42 | 8% | 42 KB | 0.01% |
| Llama-70b | 80 | 6% | 60 KB | <0.01% |
Metadata (.saturation.json): 5-15 KB per model
Deserialization Performance
Blob size: 0.04 MB (32 layers)
Deserialization: 45 µs/op (per-blob)
Per-model load: 1.4 ms (32 layers)In context: 1.4ms amortized over model lifecycle negligible.
Inference Latency Reduction
Conservative estimates (actual depends on kernel implementation):
- 10% frozen neurons: ~1% latency reduction
- 15% frozen neurons: ~1.5% latency reduction
- 20% frozen neurons: ~2% latency reduction
Maximum observed reduction with well-optimized kernels: 3-5% per 10% frozen.
Supported Models
Pre-computed masks available for:
Phi-3-mini (32 layers, 3072 hidden)
- Variance-gradient method
- 12.5% avg frozen neurons
- Speedup estimate: 1.125x
Qwen2.5-7B (28 layers, 4096 hidden)
- Variance-gradient method
- 10.2% avg frozen neurons
- Speedup estimate: 1.102x
Gemma-9b (42 layers, 3584 hidden)
- Statistical method
- 8.1% avg frozen neurons
- Speedup estimate: 1.081x
Llama-70b (80 layers, 8192 hidden)
- Variance-gradient method
- 6.3% avg frozen neurons
- Speedup estimate: 1.063x
Reproducibility & Consistency
Saturation masks are reproducible across runs with same calibration data:
# Run 1
masks1 = compute_saturation_masks(model, tokenizer, num_samples=512)
# Run 2 (same model, tokenizer, samples)
masks2 = compute_saturation_masks(model, tokenizer, num_samples=512)
# masks1 == masks2 (bit-identical)Consistency guarantee: Masks computed with same method on same model converge to stable set within 2-3 runs.
Cliff Scores
Per-layer cliff detection scores (0-1) indicate saturation sharpness:
- Score 0.1-0.3: Gradual saturation (neurons freeze smoothly over sequence)
- Score 0.3-0.6: Moderate cliff (sharp freezing at certain positions)
- Score 0.6-1.0: Sharp cliff (nearly all frozen neurons activate early)
Use for:
- Layer skipping heuristics: Skip layers with score > 0.8
- Adaptive compute: Vary batch processing based on cliff patterns
- Quality estimation: High cliff scores = more stable quantization candidates
Fallback Behavior
When masks are missing or corrupted:
masks = load_saturation_masks(model_path, fallback_to_empty=True)
# Returns: {} (empty dict, no-op application)
# Model inference continues normally, no masks applied
masks = load_saturation_masks(model_path, fallback_to_empty=False)
# Returns: None (indicates explicit failure)
# Caller can decide on error handlingValidation Checklist
Before deploying saturation masks:
- Masks computed on representative calibration corpus
- Consistency verified across ≥2 runs
- Deserialization tested on target platform
- Baseline accuracy measured (pre-mask)
- Inference latency benchmarked
- Masked inference accuracy ≥99.5% of baseline
- Serialized file < 1MB
Troubleshooting
Masks not loading
# Check if file exists and format is correct
masks = load_saturation_masks(path, fallback_to_empty=False)
if masks is None:
# Debug: Check file format
from saturation_mask_decoder import load_saturation_metadata
meta = load_saturation_metadata(path)
print(meta) # Should show method, timestamp, etc.Deserialization errors
# Check blob integrity
from saturation_mask_decoder import _deserialize_masks_blob
blob = ... # extracted from model
masks = _deserialize_masks_blob(blob)
if not masks:
print("Blob corrupted or empty")Accuracy loss
- Try different calibration corpus (larger, more diverse)
- Increase
num_samplesto 1024 or 2048 - Switch to
statisticalmethod (less aggressive freezing) - Reduce
variance_quantilefrom 0.2 to 0.1
References
- Saturation detection: COMPLETE_TRAINING_SATURATION_INSIGHTS.md
- RKNOT format: rknot/saturation_metadata.rs
- Benchmark results: BENCHMARK_WINS_COMPREHENSIVE.md