Combine diffusion-based parallel drafting with autoregressive sampling in a single forward pass using structured attention masks—achieving 5x+ token throughput while maintaining autoregressive-level output quality through hybrid generation.
Standardmäßig ist der Prompt ausgewählt, der zuerst die Quelle prüft. Sie können zu einem direkten Befehl wechseln oder eine lokale Kopie herunterladen.
Quelldateien prüfen
Lesen Sie SKILL.md und alle von SkillsMP angezeigten Begleitdateien, bevor Sie sich für eine Installation entscheiden.
Mit Codex oder Claude installieren Kopieren Sie diesen Prompt, fügen Sie ihn in Codex, Claude oder einen anderen Assistant ein und lassen Sie die Skill-Seite prüfen und installieren.
Ein direkter Befehl überspringt den Prüf-Prompt. Prüfen Sie die Quelle, bevor Sie ihn ausführen.
Combine diffusion-based parallel drafting with autoregressive sampling in a single forward pass using structured attention masks—achieving 5x+ token throughput while maintaining autoregressive-level output quality through hybrid generation.
Hybrid Diffusion-Autoregressive Generation for High-Throughput Language Models
Autoregressive language models are high-quality but slow (one token per forward pass). Diffusion models generate tokens in parallel but with lower quality. TiDAR merges both paradigms: diffusion generates candidate tokens in parallel (drafting phase), then autoregression selects final outputs sequentially (refinement phase)—all in a single forward pass using structured attention.
The approach achieves 4.71x to 5.91x tokens per second compared to pure autoregression while maintaining comparable quality, solving a fundamental speed-quality tradeoff.
Core Concept
TiDAR operates in two phases within one neural forward pass:
Thinking (Diffusion) - Parallel iterative refinement generates k candidate tokens for each position
Talking (Autoregression) - Sequential sampling selects final tokens using context and diffusion candidates
Structured attention masks enable this hybrid within a single transformer: diffusion layers have all-to-all connectivity (parallel thinking), while autoregressive layers have causal masks (sequential talking). The architecture transitions smoothly between thinking and talking phases.
Architecture Overview
Diffusion Thinking Layers: Parallel token generation with iterative refinement
Structured Attention Masks: All-to-all for diffusion; causal for autoregression
Candidate Representation: Stores k candidate tokens per position for AR selection
Autoregressive Refinement: Sequential sampling from diffusion candidates
Hybrid Router: Decides when to transition from diffusion to autoregressive phase
Efficient Masking: Single forward pass enables gradient flow through both paradigms
Implementation Steps
Step 1: Diffusion-Based Candidate Generation
Generate multiple token candidates per position through iterative diffusion.
self, vocab_size: int, embed_dim: int, num_candidates: int = 8,
num_diffusion_steps: int = 4
"""
Args:
vocab_size: Size of vocabulary
embed_dim: Embedding dimension
num_candidates: Number of candidate tokens per position
num_diffusion_steps: Refinement iterations
"""
super
self
self
self
self
# Learnable noise scheduler
self
1.0
0.0
# Refinement layers
self
for
in
range
def
generate_candidates
self, hidden_states: torch.Tensor
"""
Generate candidate tokens through diffusion.
Args:
hidden_states: Model hidden states [batch_size, seq_len, embed_dim]
Returns:
candidates: Candidate logits [batch_size, seq_len, num_candidates, vocab_size]
"""
# Initialize candidates with noise
# Start from uniform random; refine toward true distribution
self
# Iterative refinement (diffusion steps)
for
in
range
self
self
# Refine candidates with context
# Attend to hidden states to bias candidates toward relevant tokens
2
# Add context bias
self
# Gradually remove noise
1
# Project to logits
# Map embedding space to vocabulary
self
return
Step 2: Structured Attention Masks
Create attention patterns enabling diffusion (all-to-all) and autoregression (causal) in one forward pass.
defcreate_hybrid_attention_mask(batch_size: int, seq_len: int, num_candidates: int,
diffusion_layers: int, ar_layers: int,
device: torch.device) -> Dict[str, torch.Tensor]:
"""
Create structured attention masks for hybrid architecture.
Args:
batch_size: Batch size
seq_len: Sequence length
num_candidates: Number of candidates per position
diffusion_layers: Number of diffusion (all-to-all) layers
ar_layers: Number of autoregressive (causal) layers
device: torch device
Returns:
masks: {diffusion_mask, ar_mask, candidate_mask}
"""# Diffusion mask: all-to-all connectivity (thinking phase)# Every position can attend to every other position
diffusion_mask = torch.ones(
batch_size, seq_len, seq_len,
device=device, dtype=torch.bool
)
# Autoregressive mask: causal (talking phase)# Each position attends to itself and previous positions only
ar_mask = torch.tril(
torch.ones(seq_len, seq_len, device=device, dtype=torch.bool)
).unsqueeze(0).expand(batch_size, -1, -1)
# Candidate mask: connections between AR positions and diffusion candidates# AR layer can attend to all candidate positions from previous step
candidate_mask = torch.ones(
batch_size, seq_len, num_candidates, seq_len,
device=device, dtype=torch.bool
)
return {
'diffusion_mask': diffusion_mask,
'ar_mask': ar_mask,
'candidate_mask': candidate_mask
}
classHybridAttentionLayer(nn.Module):
"""
Single attention layer supporting both diffusion and AR patterns.
"""def__init__(self, embed_dim: int, num_heads: int = 8):
super().__init__()
self.embed_dim = embed_dim
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
self.query = nn.Linear(embed_dim, embed_dim)
self.key = nn.Linear(embed_dim, embed_dim)
self.value = nn.Linear(embed_dim, embed_dim)
self.output = nn.Linear(embed_dim, embed_dim)
defforward(self, hidden_states: torch.Tensor, attention_mask: torch.Tensor,
phase: str = 'diffusion') -> torch.Tensor:
"""
Apply hybrid attention.
Args:
hidden_states: [batch, seq_len, embed_dim]
attention_mask: Attention mask (diffusion or AR)
phase: 'diffusion' or 'autoregressive'
Returns:
output: Attended states [batch, seq_len, embed_dim]
"""
batch_size, seq_len, embed_dim = hidden_states.shape
# Compute Q, K, V
Q = self.query(hidden_states)
K = self.key(hidden_states)
V = self.value(hidden_states)
# Reshape for multi-head attention
Q = Q.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
K = K.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
V = V.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
# Scaled dot-product attention
scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.head_dim ** 0.5)
# Apply maskif attention_mask isnotNone:
# Mask shape: [batch, 1, seq_len, seq_len]
scores = scores.masked_fill(~attention_mask.unsqueeze(1), float('-inf'))
# Softmax and dropout
attn_weights = F.softmax(scores, dim=-1)
attn_output = torch.matmul(attn_weights, V)
# Reshape back
attn_output = attn_output.transpose(1, 2).contiguous()
attn_output = attn_output.view(batch_size, seq_len, embed_dim)
# Output projection
output = self.output(attn_output)
return output
Step 3: Autoregressive Refinement from Candidates
Select final tokens from diffusion candidates using autoregressive sampling.
classAutoregressiveRefinement(nn.Module):
"""
Refines diffusion candidates through autoregressive sampling.
"""def__init__(self, vocab_size: int, embed_dim: int):
super().__init__()
self.vocab_size = vocab_size
self.embed_dim = embed_dim
# Selection network: learns to pick best candidateself.selector = nn.Sequential(
nn.Linear(embed_dim + vocab_size, embed_dim),
nn.ReLU(),
nn.Linear(embed_dim, 1)
)
defselect_tokens(self, candidates_logits: torch.Tensor,
context_hidden: torch.Tensor) -> torch.Tensor:
"""
Select best candidate tokens given context.
Args:
candidates_logits: [batch, seq_len, num_candidates, vocab_size]
context_hidden: [batch, seq_len, embed_dim]
Returns:
selected_tokens: [batch, seq_len, vocab_size]
"""
batch_size, seq_len, num_candidates, vocab_size = candidates_logits.shape
# Compute candidate probabilities
candidate_probs = F.softmax(candidates_logits, dim=-1)
# Score each candidate position
scores = []
for c inrange(num_candidates):
# Get probabilities for this candidate set
cand_probs = candidate_probs[:, :, c, :]
# Combine with context
combined = torch.cat([
context_hidden,
cand_probs
], dim=-1)
# Compute selection score
score = self.selector(combined).squeeze(-1)
scores.append(score)
scores = torch.stack(scores, dim=-1) # [batch, seq_len, num_candidates]# Select highest-scoring candidate per position
selected_idx = torch.argmax(scores, dim=-1) # [batch, seq_len]# Gather selected logits
selected_logits = torch.gather(
candidates_logits,
2,
selected_idx.unsqueeze(-1).unsqueeze(-1).expand(-1, -1, 1, vocab_size)
).squeeze(2)
return selected_logits
Step 4: Unified Forward Pass
Combine diffusion thinking and AR talking into single forward pass.