| name | resa-transparent-reasoning |
| title | Resa: Transparent Reasoning Models via SAEs |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2506.09967 |
| keywords | ["sparse autoencoders","reasoning transfer","interpretability","efficient training","reasoning tuning"] |
| description | Extract and transfer reasoning abilities using sparse autoencoders (SAE-Tuning) on CoT-free data, achieving RL-equivalent performance at 2000x lower cost and 450x faster training. |
Resa: Transparent Reasoning Models via SAEs
Core Concept
Resa demonstrates that reasoning abilities are learnable, transferable features that can be extracted via sparse autoencoders (SAE) without expensive reinforcement learning. The two-stage SAE-Tuning procedure trains an autoencoder on source model activations, then inserts it into a target model with low-rank adapters. This achieves RL-equivalent reasoning performance at approximately $1 cost and 20-minute training time, compared to months and $2000+ for traditional RL approaches.
Architecture Overview
- Two-Stage SAE-Tuning: Stage I trains SAE on source model activations; Stage II inserts frozen SAE into target model with LoRA for implicit reasoning pattern transfer
- CoT-Free Training Data: Uses only verified question-answer pairs (no intermediate reasoning traces), reducing data requirements
- Reasoning as Portable Adapter: Extracted reasoning features function as modular "reasoning adapters" transferable across model families without retraining
- Transparent Feature Extraction: Prompt-only method identifies latent reasoning features; layer-wise distribution correlates with reasoning performance
- Massive Cost Reduction: $1 per model vs $2000 for RL; 20 minutes training vs months for full RL pipeline
Implementation
Step 1: Sparse Autoencoder Training
import torch
import torch.nn as nn
class SparseAutoencoder(nn.Module):
"""
Trains on source model activations to extract reasoning features.
Sparse representation: high dimensionality (65k features) but low k activation.
"""
def __init__(self, input_dim, num_features=65536, k=32):
super().__init__()
self.input_dim = input_dim
self.num_features = num_features
self.k = k
self.encoder = nn.Linear(input_dim, num_features)
self.decoder = nn.Linear(num_features, input_dim)
with torch.no_grad():
self.decoder.weight.copy_(self.encoder.weight.T)
def forward(self, x):
"""
Forward pass with top-k sparsity constraint.
Returns reconstruction and sparse features.
"""
features = self.encoder(x)
k_values, k_indices = torch.topk(torch.abs(features), self.k, dim=1)
sparse_features = torch.zeros_like(features)
sparse_features.scatter_(1, k_indices, features.gather(, k_indices))
reconstruction = .decoder(sparse_features)
reconstruction, sparse_features
():
optimizer = torch.optim.Adam(.parameters(), lr=)
loss_fn = nn.MSELoss()
epoch (num_epochs):
batch trigger_dataset:
question_ids = batch[]
torch.no_grad():
activations = source_model.get_layer_activation(
question_ids,
layer=
)
reconstruction, sparse_features = .forward(activations)
recon_loss = loss_fn(reconstruction, activations)
sparsity_loss = * sparse_features.().(dim=).mean()
total_loss = recon_loss + sparsity_loss
optimizer.zero_grad()
total_loss.backward()
optimizer.step()
batch % == :
()
Step 2: SAE-Guided SFT in Target Model
class LoRAAdapter(nn.Module):
"""
Low-rank adapter for inserting SAE into target model.
Minimizes KL divergence between outputs with/without SAE.
"""
def __init__(self, hidden_dim, lora_rank=32):
super().__init__()
self.lora_rank = lora_rank
self.lora_q = nn.Linear(hidden_dim, lora_rank)
self.lora_q_out = nn.Linear(lora_rank, hidden_dim)
self.lora_k = nn.Linear(hidden_dim, lora_rank)
self.lora_k_out = nn.Linear(lora_rank, hidden_dim)
self.lora_v = nn.Linear(hidden_dim, lora_rank)
self.lora_v_out = nn.Linear(lora_rank, hidden_dim)
with torch.no_grad():
self.lora_q_out.weight.zero_()
self.lora_k_out.weight.zero_()
self.lora_v_out.weight.zero_()
def forward(self, x):
"""Apply LoRA update to activation."""
delta_q = self.lora_q_out(self.lora_q(x))
delta_k = self.lora_k_out(self.lora_k(x))
delta_v = self.lora_v_out(self.lora_v(x))
return {'q': delta_q, 'k': delta_k, 'v': delta_v}
class SAEGuidedTraining:
():
.target_model = target_model
.sae = sae
.layer_idx = layer_idx
.lora = LoRAAdapter(target_model.hidden_size)
param .sae.parameters():
param.requires_grad =
():
optimizer = torch.optim.Adam(.lora.parameters(), lr=learning_rate)
epoch (num_epochs):
batch trigger_dataset:
input_ids = batch[]
target_tokens = batch[]
torch.no_grad():
logits_base = .target_model(input_ids)
log_probs_base = torch.nn.functional.log_softmax(logits_base, dim=-)
activations = .target_model.get_layer_activation(
input_ids,
layer=.layer_idx
)
_, sparse_features = .sae(activations)
lora_delta = .lora(sparse_features)
logits_sae = .target_model(
input_ids,
lora_delta=lora_delta
)
log_probs_sae = torch.nn.functional.log_softmax(logits_sae, dim=-)
kl_loss = torch.nn.functional.kl_div(
log_probs_sae,
torch.exp(log_probs_base),
reduction=
)
optimizer.zero_grad()
kl_loss.backward()
optimizer.step()
()
Step 3: Transparent Reasoning Feature Analysis
class ReasoningFeatureAnalyzer:
"""
Identify and quantify latent reasoning features via prompt-only method.
Reveals which SAE features activate during reasoning tasks.
"""
def __init__(self, sae, model):
self.sae = sae
self.model = model
def identify_reasoning_features(self, test_dataset):
"""
Analyze which SAE features correlate with reasoning performance.
Returns feature importance scores.
"""
feature_activations = torch.zeros(self.sae.num_features)
feature_task_performance = torch.zeros(self.sae.num_features)
for sample in test_dataset:
with torch.no_grad():
activations = self.model.get_layer_activation(
sample['input_ids'],
layer=12
)
_, sparse_features = self.sae(activations)
feature_activations += (sparse_features.abs().sum(dim=0) > 0).float()
correctness = 1.0 if sample['is_correct'] else 0.0
feature_task_performance += sparse_features[0] * correctness
feature_importance = feature_task_performance / (feature_activations + )
feature_importance
():
layer_importance = {}
layer_idx (model.num_layers):
features_this_layer =
sample dataset:
torch.no_grad():
activations = model.get_layer_activation(
sample[],
layer=layer_idx
)
_, sparse_features = .sae(activations)
features_this_layer += sparse_features.().().item()
layer_importance[layer_idx] = features_this_layer / (dataset)
layer_importance
Step 4: Deployment and Cost Analysis
def compute_training_cost(model_size_b, hours_to_train, cost_per_hour_gpu=0.5):
"""Estimate training cost for SAE-tuning."""
num_gpus = max(1, model_size_b // 10)
total_cost = hours_to_train * num_gpus * cost_per_hour_gpu
return total_cost
print(f"SAE-Tuning cost: $1.00")
print(f"RL cost: $2000.00")
print(f"Cost reduction: 2000x")
Practical Guidance
Implementation Steps:
- Select source model with reasoning (R1-Distill, DeepSeek-R1)
- Prepare trigger dataset: verified QAs with
<think></think> markers but no intermediate steps
- Train SAE on source model activations (layer 12 recommended)
- Insert frozen SAE into target model with LoRA adapters
- Fine-tune LoRA on trigger dataset (1 epoch sufficient)
Data Requirements:
- Trigger dataset: 5k-10k verified question-answer pairs
- Source: STILL, DeepScaleR, or similar CoT-free datasets
- No need for intermediate reasoning steps, reducing data requirements by 90%
Generalization:
- Reasoning features transfer across datasets (AIME, MATH, GPQA)
- Works across model families (Qwen, LLaMA variants, DeepSeek)
- Adapters remain small (LoRA rank 32) enabling multi-adapter composition
Performance Targets:
- R1-Distill-1.5B: ~35% AIME pass@1 (RL-equivalent)
- Generalizes to unseen benchmarks without retraining
- Multimodal models: reasoning benefits transfer to vision tasks
Reference
- Sparse Autoencoders: Extract interpretable features from neural networks; top-k sparsity prevents feature collapse
- LoRA: Low-rank adaptation; efficient parameterization for transfer learning
- KL divergence: Measures distribution difference; KL minimization encourages model to adopt SAE-guided behavior