Token-adaptive preference optimization framework using min-max formulation to reduce multimodal LLM hallucination. Achieves 50% hallucination reduction using min-max distributional robustness while preserving visual grounding.
TARS: MinMax Token-Adaptive Preference Strategy for Hallucination Reduction
TARS addresses the critical problem of hallucination in multimodal large language models through a novel token-level adaptive preference optimization approach. Rather than treating preference signals as fixed targets, TARS uses min-max optimization to train models that are robust to preference label uncertainty while maintaining visual grounding.
Core Concept
The key insight is that hallucinations arise when language models overfit to spurious patterns in preference data, disconnecting outputs from actual visual content. TARS reformulates preference optimization as a min-max problem:
Outer maximization: Shift the model's token distribution to satisfy preference constraints
Inner minimization: Perturb preferences within a semantic budget to simulate uncertainty
Equilibrium: Train model to be robust against preference variations that don't compromise visual grounding
This approach reduces hallucinations from 26.4% to 13.2% using only 4.8k preference samples, achieving state-of-the-art performance.
Architecture Overview
The framework consists of:
Preference Encoding: Represent human preferences as token-level target distributions
Semantic Constraint Layer: Define allowed perturbations without changing meaning
Min-Max Optimizer: Outer loop for distribution shift, inner loop for robustness
Visual Grounding Preservation: Ensure steering doesn't disconnect from image content
Token-Level Adaptation: Apply different strengths of preference steering per token
Implementation Steps
Step 1: Encode preferences as token-level distributions
Convert pairwise preference labels into token-level target distributions:
"""
Convert preferred vs dispreferred outputs to target distributions.
Args:
preferred_text: The preferred response
dispreferred_text: The dispreferred response
Returns:
(preferred_dist, dispreferred_dist) of shape (seq_len, vocab_size)
"""
self
self
# Create one-hot distributions for each sequence
max
len
len
self
self
for
in
enumerate
1.0
for
in
enumerate
1.0
return
def
create_preference_labels
self, pairs: List[Tuple[str, str]]
"""
Create batch of preference target distributions.
Args:
pairs: List of (preferred, dispreferred) text pairs
Returns:
Tensor of shape (batch_size, seq_len, vocab_size)
"""
for
in
self
# Pad to same length
max
0
for
in
len
self
for
in
enumerate
0
return
This represents preferences as target probability distributions over tokens.
Step 2: Implement min-max optimization with perturbation
The core training loop applies min-max optimization to achieve robustness:
classMinMaxPreferenceOptimizer:
"""Min-max optimization for robust preference learning"""def__init__(self, model, tokenizer, epsilon: float = 0.1):
self.model = model
self.tokenizer = tokenizer
self.epsilon = epsilon # Perturbation budgetself.vocab_size = model.config.vocab_size
defcompute_preference_loss(self, logits: torch.Tensor,
target_dist: torch.Tensor,
positions: torch.Tensor) -> torch.Tensor:
"""
Compute loss between predicted logits and preference targets.
Args:
logits: Model output logits, shape (batch, seq_len, vocab_size)
target_dist: Target probability distribution over tokens
positions: Which positions correspond to preference targets
Returns:
Scalar loss value
"""# Extract logits at preference positions
pred_probs = F.softmax(logits, dim=-1)
# KL divergence from model distribution to target
loss = F.kl_div(
F.log_softmax(logits, dim=-1),
target_dist.clamp(min=1e-8),
reduction='batchmean'
)
return loss
defperturbation_step(self, target_dist: torch.Tensor,
step_size: float = 0.01) -> torch.Tensor:
"""
Inner optimization: find worst-case perturbation within semantic budget.
Args:
target_dist: Current preference target
step_size: Optimization step size
Returns:
Perturbed distribution that remains semantically valid
"""
perturbed = target_dist.clone().detach().requires_grad_(True)
optimizer = torch.optim.Adam([perturbed], lr=step_size)
for _ inrange(5): # Inner loop iterations# Compute model loss on perturbed targets
logits = self.model(...) # Forward pass
loss = self.compute_preference_loss(logits, perturbed, mask=None)
# Maximize loss (adversarial): find worst perturbation
loss_adv = -loss
optimizer.zero_grad()
loss_adv.backward()
optimizer.step()
# Project back to semantic constraint set
perturbed.data = self._project_to_semantic_budget(perturbed.data)
return perturbed.detach()
def_project_to_semantic_budget(self, dist: torch.Tensor) -> torch.Tensor:
"""
Project distribution back to semantic constraint set.
Keep similar tokens with higher weight.
"""# Reproject to valid probability distribution
dist = F.softmax(dist, dim=-1)
# Constraint: perturbed distribution cannot exceed epsilon away from original# (in terms of TV divergence or other semantic distance metric)# Simplified: clamp large changesreturn dist
defmin_max_training_step(self, images: torch.Tensor,
prompts: torch.Tensor,
preference_targets: torch.Tensor) -> Tuple[torch.Tensor, float]:
"""
Single min-max training step: maximize robustness to preference perturbations.
Args:
images: Visual input
prompts: Text prompts
preference_targets: Target distributions for next tokens
Returns:
(updated_model_state, loss_value)
"""# Outer loop: optimize model parameters
optimizer = torch.optim.AdamW(self.model.parameters(), lr=1e-5)
for outer_step inrange(3): # Outer iterations# Inner loop: find worst perturbation
perturbed_targets = self.perturbation_step(preference_targets)
# Forward pass with model
outputs = self.model(images, prompts)
logits = outputs.logits
# Compute loss on perturbed targets
loss = self.compute_preference_loss(logits, perturbed_targets, positions=None)
# Outer optimization: minimize loss on worst-case perturbation
optimizer.zero_grad()
loss.backward()
optimizer.step()
returnself.model.state_dict(), loss.item()
This implements the core min-max algorithm that creates robustness to preference uncertainty.
Step 3: Implement token-level adaptive steering
Apply different preference strengths to different tokens: