| name | tracealign-alignment-failure-attribution |
| title | TraceAlign - Tracing Alignment Failures to Training Beliefs |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2508.02063 |
| keywords | ["alignment","attribution","interpretability","training-data"] |
| description | Trace unsafe LLM outputs back to conflicting beliefs in training data using suffix-array matching and Belief Conflict Index. |
TraceAlign: Alignment Failure Attribution
TraceAlign identifies which training data caused unsafe model outputs by measuring semantic conflicts between generated text and alignment policies. The Belief Conflict Index quantifies misalignment, enabling targeted interventions: inference-time filtering, fine-tuning corrections, or safer decoding.
Core Concept
When LLMs produce unsafe outputs, understanding why is crucial. TraceAlign traces back to source: the conflicting beliefs learned during pretraining. It uses suffix-array matching to retrieve training documents similar to unsafe spans, computes a Belief Conflict Index measuring policy violation severity, and enables three interventions: TraceShield (refuse), contrastive deconfliction (fine-tune), and Prov-Decode (constrained beam search).
Architecture Overview
- Belief Conflict Index (BCI): Measures semantic disagreement between generated span and alignment policy using retrieved training documents
- Suffix-Array Matching: Efficiently retrieve training documents matching unsafe spans
- Three-Pronged Intervention: TraceShield (inference filtering), contrastive fine-tuning, Prov-Decode (constrained decoding)
- Theoretical Upper Bound: Formal analysis linking memorization frequency to reactivation risk
- Up to 85% Drift Reduction: Preserves model utility while reducing unsafe behavior
Implementation Steps
Step 1: Implement Belief Conflict Index
import numpy as np
from typing import List, Tuple, Dict
class BeliefConflictIndex:
"""Quantify semantic conflict between model output and alignment policy."""
def __init__(self, policy_descriptions: List[str]):
"""
Args:
policy_descriptions: List of safety policies, e.g., ["Don't help with illegal activities"]
"""
self.policies = policy_descriptions
self.embedder = None
def compute_bci(self, generated_span: str, policy_idx: int) -> float:
"""
Compute Belief Conflict Index for a span against a policy.
High BCI = strong conflict = unsafe.
Returns: BCI score (0-1)
"""
policy = self.policies[policy_idx]
span_embedding = self.embedder.encode(generated_span)
policy_embedding = self.embedder.encode(policy)
cosine_dist = 1 - np.dot(span_embedding, policy_embedding) / (
np.linalg.norm(span_embedding) * np.linalg.norm(policy_embedding) + 1e-6
)
return cosine_dist
def compute_bci_batch(self, spans: [], training_docs: []) -> [, ]:
bci_scores = {}
span spans:
matching_docs = ._find_matching_docs(span, training_docs)
conflicts = []
doc matching_docs:
policy_idx, _ (.policies):
bci = .compute_bci(doc, policy_idx)
conflicts.append(bci)
bci_scores[span] = np.mean(conflicts) conflicts
bci_scores
() -> []:
span_embedding = .embedder.encode(span)
similarities = []
doc training_docs:
doc_embedding = .embedder.encode(doc)
similarity = np.dot(span_embedding, doc_embedding) / (
np.linalg.norm(span_embedding) * np.linalg.norm(doc_embedding) +
)
similarities.append(similarity)
top_indices = np.argsort(similarities)[-top_k:]
[training_docs[i] i top_indices]
Step 2: Implement Suffix-Array Retrieval
class SuffixArrayMatcher:
"""Efficiently retrieve training data matching unsafe spans."""
def __init__(self, training_corpus: List[str]):
self.corpus = training_corpus
self.suffix_array = self._build_suffix_array()
def _build_suffix_array(self) -> List[Tuple[str, int, int]]:
"""Build suffix array for fast substring matching."""
suffix_array = []
for doc_idx, doc in enumerate(self.corpus):
for start_pos in range(len(doc)):
suffix = doc[start_pos:]
suffix_array.append((suffix, doc_idx, start_pos))
suffix_array.sort(key=lambda x: x[0])
return suffix_array
def find_matching_spans(self, query_span: str, k: int = 5) -> List[Tuple[str, int, int]]:
"""
Find top-k training corpus occurrences matching query span.
Uses binary search on suffix array for efficiency.
"""
matches = []
for suffix, doc_idx, start_pos .suffix_array:
suffix.startswith(query_span) query_span suffix[:(query_span) * ]:
matches.append((.corpus[doc_idx], doc_idx, start_pos))
(matches) >= k:
matches[:k]
Step 3: Implement TraceShield (Inference-Time Filtering)
class TraceShield:
"""Refuse outputs containing high-BCI unsafe spans."""
def __init__(self, bci_threshold: float = 0.7):
self.bci_threshold = bci_threshold
self.bci_computer = BeliefConflictIndex(['no harmful content'])
def filter_output(self, generated_text: str, training_docs: List[str]) -> Tuple[str, bool]:
"""
Check if output contains unsafe spans. Refuse if BCI too high.
Returns: (output_text, is_safe)
"""
spans = generated_text.split('. ')
unsafe_spans = []
for span in spans:
bci_scores = self.bci_computer.compute_bci_batch([span], training_docs)
if max(bci_scores.values()) > self.bci_threshold:
unsafe_spans.append(span)
if unsafe_spans:
return "I can't provide that information.", False
return generated_text, True
Step 4: Implement Contrastive Belief Deconfliction
import torch
import torch.nn as nn
class ContrastiveBeliefDeconfliction:
"""Fine-tune to resolve belief conflicts via DPO-style training."""
def __init__(self, model, bci_computer: BeliefConflictIndex):
self.model = model
self.bci_computer = bci_computer
def create_preference_pairs(self, unsafe_completions: List[str],
safe_completions: List[str],
training_docs: List[str]) -> List[Tuple[str, str]]:
"""
Create preference pairs: (preferred_safe, dispreferred_unsafe).
"""
pairs = []
for unsafe, safe in zip(unsafe_completions, safe_completions):
unsafe_bci = self.bci_computer.compute_bci(unsafe, 0)
safe_bci = self.bci_computer.compute_bci(safe, 0)
if unsafe_bci > safe_bci:
pairs.append((safe, unsafe))
return pairs
def train_dpo(self, pairs: List[Tuple[str, str]], learning_rate: float = ):
optimizer = torch.optim.Adam(.model.parameters(), lr=learning_rate)
preferred, dispreferred pairs:
preferred_logp = .model.compute_logp(preferred)
dispreferred_logp = .model.compute_logp(dispreferred)
dpo_loss = -torch.log(torch.sigmoid(preferred_logp - dispreferred_logp))
optimizer.zero_grad()
dpo_loss.backward()
optimizer.step()
dpo_loss.item()
Step 5: Implement Prov-Decode
class ProvDecode:
"""Constrained beam search preventing high-BCI spans."""
def __init__(self, model, bci_computer: BeliefConflictIndex,
bci_threshold: float = 0.7):
self.model = model
self.bci_computer = bci_computer
self.threshold = bci_threshold
def safe_beam_search(self, prompt: str, beam_width: int = 5,
max_length: int = 100,
training_docs: List[str] = None) -> str:
"""
Beam search that avoids high-BCI continuations.
"""
candidates = [(prompt, 0.0)]
for position in range(max_length):
new_candidates = []
for text, log_prob in candidates:
next_logits = self.model.get_next_token_logits(text)
top_k_tokens = torch.topk(next_logits, k=10).indices
for token_idx in top_k_tokens:
token_text = self.model.tokenizer.decode([token_idx])
new_text = text + token_text
training_docs ._contains_unsafe_span(new_text, training_docs):
new_log_prob = log_prob + torch.log_softmax(next_logits, dim=)[token_idx]
new_candidates.append((new_text, new_log_prob.item()))
new_candidates.sort(key= x: x[], reverse=)
candidates = new_candidates[:beam_width]
candidates:
candidates[][] candidates prompt
() -> :
recent_span = text.split()[-:]
recent_text = .join(recent_span)
bci = .bci_computer.compute_bci(recent_text, )
bci > .threshold
Step 6: Theoretical Analysis
def compute_drift_upper_bound(memorization_freq: float, span_length: int) -> float:
"""
Theoretical upper bound on drift likelihood.
Relates memorization frequency and length to reactivation risk.
Higher memorization + longer span = higher reactivation risk.
"""
drift_bound = memorization_freq * np.log(max(span_length, 2))
return min(drift_bound, 1.0)
Practical Guidance
When to Use:
- Debugging unsafe model outputs
- Fine-tuning with alignment objectives
- Scenarios with extensive logged training data
- Compliance/auditing requirements
When NOT to Use:
- Models with unknown/unavailable training data
- Real-time applications requiring <100ms latency
- Scenarios where approximate attribution is insufficient
Hyperparameters:
| Parameter | Default | Impact |
|---|
bci_threshold | 0.7 | Higher = stricter safety, more refusals |
suffix_match_k | 5 | More matches = better BCI estimates, higher computation |
dpo_learning_rate | 1e-5 | Standard LLM fine-tuning rate |
beam_width | 5 | Wider beams = better diversity but slower |
Reference
Paper: TraceAlign: Tracing Alignment Failures to Training-Time Belief Sources (2508.02063)
- Belief Conflict Index quantifies alignment violations
- Up to 85% drift mitigation
- Three interventions: filtering, fine-tuning, constrained decoding