| name | infialign-data-selection |
| title | InfiAlign - Scalable Framework for Aligning LLMs for Reasoning |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2508.05496 |
| keywords | ["data-selection","supervised-fine-tuning","direct-preference-optimization","reasoning"] |
| description | Combines SFT and DPO with robust data selection pipeline using multidimensional quality metrics. Achieves DeepSeek-R1 performance with 12% training data, enabling efficient reasoning model alignment. |
InfiAlign: Scalable Framework for Aligning LLMs for Reasoning
Core Concept
InfiAlign addresses the inefficiency of reasoning model alignment by combining SFT and DPO with an intelligent data selection pipeline. Rather than training on all available data, the framework uses multidimensional quality metrics to automatically curate high-quality alignment data from open-source reasoning datasets. This dramatically reduces data requirements while maintaining or exceeding performance compared to full-data training.
Architecture Overview
- Multidimensional Quality Scoring: Evaluates data across multiple dimensions (correctness, clarity, complexity)
- Automated Data Curation: Selects only high-quality examples for training
- SFT Phase: Initial fine-tuning on curated data
- DPO Phase: Direct Preference Optimization to refine reasoning quality
- Data Efficiency: Achieves strong performance with minimal training data
Implementation Steps
Step 1: Implement Multidimensional Quality Scoring
Create system to evaluate training data quality across multiple dimensions.
import numpy as np
from typing import List, Dict, Tuple
from dataclasses import dataclass
@dataclass
class QualityScore:
"""Quality assessment for training example."""
correctness: float
clarity: float
complexity: float
efficiency: float
diversity: float
overall: float
class QualityEvaluator:
"""
Multi-dimensional quality scoring for reasoning examples.
"""
def __init__(self, verifier_model, embedding_model):
self.verifier = verifier_model
self.embedder = embedding_model
self.example_embeddings = []
def evaluate_example(self, problem: str, solution: str, answer: str) -> QualityScore:
"""
Evaluate quality of training example.
Args:
problem: Problem statement
solution: Reasoning steps
answer: Final answer
Returns:
Multi-dimensional quality score
"""
correctness = ._evaluate_correctness(problem, answer)
clarity = ._evaluate_clarity(solution)
complexity = ._evaluate_complexity(problem)
efficiency = ._evaluate_efficiency(solution)
diversity = ._evaluate_diversity(solution)
overall = (
* correctness +
* clarity +
* complexity +
* efficiency +
* diversity
)
QualityScore(
correctness=correctness,
clarity=clarity,
complexity=complexity,
efficiency=efficiency,
diversity=diversity,
overall=overall
)
() -> :
prompt =
response = .verifier.generate(prompt)
response.upper()
() -> :
token_count = (solution.split())
<= token_count <= :
clarity =
token_count < :
clarity = token_count /
:
clarity = (, - (token_count - ) / )
explanation_keywords = [, , , , ]
keyword_count = ( kw explanation_keywords kw solution.lower())
clarity = (clarity + (keyword_count / , )) /
clarity
() -> :
complexity_signals = []
word_count = (problem.split())
length_score = (word_count, ) /
complexity_signals.append(length_score * )
math_ops = problem.count() + problem.count() + problem.count() + problem.count()
op_score = (math_ops, ) /
complexity_signals.append(op_score * )
constraints = problem.count() + problem.count() + problem.count()
constraint_score = (constraints, ) /
complexity_signals.append(constraint_score * )
complexity = (complexity_signals)
(complexity, )
() -> :
token_count = (solution.split())
token_count < :
efficiency =
token_count < :
efficiency =
:
efficiency = (, - (token_count - ) / )
efficiency
() -> :
.example_embeddings:
new_embedding = .embedder.encode(solution)
similarities = [
np.dot(new_embedding, existing) / (np.linalg.norm(new_embedding) * np.linalg.norm(existing))
existing .example_embeddings
]
max_similarity = (similarities)
diversity = - max_similarity
diversity
() -> [[, QualityScore]]:
scored = []
example examples:
score = .evaluate_example(
example[],
example[],
example[]
)
embedding = .embedder.encode(example[])
.example_embeddings.append(embedding)
scored.append((example, score))
scored.sort(key= x: x[].overall, reverse=)
scored
Step 2: Implement Data Selection Pipeline
Create smart data selection based on quality scores.
class DataSelectionPipeline:
"""
Intelligently select training data based on quality metrics.
"""
def __init__(self, quality_evaluator: QualityEvaluator):
self.evaluator = quality_evaluator
def select_data(
self,
examples: List[Dict],
target_size: int = 5000,
min_quality: float = 0.6,
quality_distribution: str = "balanced"
) -> List[Dict]:
"""
Select data for training.
Args:
examples: Candidate examples
target_size: Target number of examples
min_quality: Minimum quality threshold
quality_distribution: How to distribute quality levels
Returns:
Selected training data
"""
scored = self.evaluator.score_dataset(examples)
filtered = [
(ex, score) for ex, score in scored
if score.overall >= min_quality
]
print(f"After quality filter: {len(filtered)} / {len(scored)} examples")
if quality_distribution == "balanced":
selected = self._select_balanced(filtered, target_size)
elif quality_distribution == "top_k":
selected = ._select_top_k(filtered, target_size)
quality_distribution == :
selected = ._select_stratified(filtered, target_size)
:
selected = ._select_top_k(filtered, target_size)
[ex ex, _ selected]
() -> :
scored[:(k, (scored))]
() -> :
easy = [s s scored s[].complexity < ]
medium = [s s scored <= s[].complexity < ]
hard = [s s scored s[].complexity >= ]
selected = []
ratios = [(easy), (medium), (hard)]
total = (ratios)
total > :
k_easy = (k * ratios[] / total)
k_medium = (k * ratios[] / total)
k_hard = k - k_easy - k_medium
selected.extend(easy[:k_easy])
selected.extend(medium[:k_medium])
selected.extend(hard[:k_hard])
selected
() -> :
buckets = {}
ex, score scored:
correctness_bucket = (score.correctness * )
complexity_bucket = (score.complexity * )
key = (correctness_bucket, complexity_bucket)
key buckets:
buckets[key] = []
buckets[key].append((ex, score))
selected = []
per_bucket = (, k // (buckets))
bucket_examples buckets.values():
selected.extend(bucket_examples[:per_bucket])
selected[:k]
Step 3: Implement SFT + DPO Training
Combine supervised fine-tuning with direct preference optimization.
class InfiAlignTrainer:
"""
SFT + DPO training with data selection.
"""
def __init__(self, model, data_selector):
self.model = model
self.data_selector = data_selector
def train(
self,
training_data: List[Dict],
num_epochs: int = 3,
sft_weight: float = 0.7,
dpo_weight: float = 0.3
):
"""
Train using SFT and DPO.
Args:
training_data: Selected training data
num_epochs: Training epochs
sft_weight: Weight for SFT loss
dpo_weight: Weight for DPO loss
"""
for epoch in range(num_epochs):
total_loss = 0
for batch in create_batches(training_data, batch_size=32):
sft_loss = self._compute_sft_loss(batch)
dpo_loss = self._compute_dpo_loss(batch)
total_loss = sft_weight * sft_loss + dpo_weight * dpo_loss
total_loss.backward()
self.model.optimizer.step()
self.model.optimizer.zero_grad()
print(f"Epoch {epoch}: Loss=")
():
loss =
example batch:
prompt =
target = example[]
loss += .model.compute_language_modeling_loss(prompt, target)
loss / (batch)
():
dpo_loss =
example batch:
preferred = example[]
prompt =
dispreferred = .model.generate(prompt, temperature=)
preferred_logprobs = .model.get_logprobs(prompt, preferred)
dispreferred_logprobs = .model.get_logprobs(prompt, dispreferred)
loss = -torch.log(
torch.sigmoid(preferred_logprobs - dispreferred_logprobs)
)
dpo_loss += loss
dpo_loss / (batch)
Practical Guidance
When to Use InfiAlign
- Efficient reasoning model training: Limited compute/data budgets
- Open-source dataset utilization: Leverage existing reasoning datasets
- Multi-metric data quality: Complex domains requiring multidimensional evaluation
- Production alignment: Data efficiency critical for cost
When NOT to Use InfiAlign
- Abundant high-quality data: Full-data training may be simpler
- Custom domains: Generic quality metrics may not apply
- Real-time data augmentation: Selection pipeline adds latency
- Unknown data distribution: Quality metrics untested on domain
Hyperparameter Recommendations
- Quality thresholds: 0.5-0.7 overall score for selection
- Correctness weight: 0.4 (highest priority)
- Clarity weight: 0.2 (important for reasoning)
- Complexity weight: 0.15 (prefer diverse problems)
- SFT/DPO split: 70% SFT, 30% DPO
Key Insights
The critical innovation is recognizing that data quality matters more than quantity. By carefully evaluating examples across multiple dimensions and selecting only high-quality data, InfiAlign achieves comparable performance to full-dataset training with 10x less data. The multidimensional scoring ensures balanced representation across difficulty and reasoning style.
Reference
InfiAlign: Scalable Framework for Aligning LLMs for Reasoning (arXiv:2508.05496)
Combines SFT and DPO with multidimensional data selection pipeline. Achieves DeepSeek-R1-distill performance using only 12% of training data through intelligent curation of reasoning examples.