| name | poss-speculative-decoding |
| title | PosS: Position Specialist Layers for Improved Speculative Decoding |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2506.03566 |
| keywords | ["speculative-decoding","inference-acceleration","draft-models"] |
| description | Improve speculative decoding throughput by employing position-specialized draft layers that handle position-specific error accumulation patterns. |
PosS: Position Specialist Generates Better Draft for Speculative Decoding
Core Concept
Speculative decoding accelerates LLM inference by using a smaller draft model to predict multiple tokens, then verifying them with the larger target model. PosS identifies a critical limitation: draft models suffer from degraded token quality at later positions due to error accumulation. By deploying multiple position-specialized layers—each handling a narrow, predictable range of feature deviation—PosS improves acceptance length and speed-up ratios by up to 5.7%.
Architecture Overview
- Position-Specific Feature Deviation: Each position in a draft has characteristic feature deviation from the target model. Later positions accumulate errors multiplicatively.
- Position Specialists: Separate transformer layers trained to handle specific positions with their expected deviation levels
- Position-Wise Acceptance Rate (Pos-Acc): New metric analyzing draft quality across different positions in the prediction sequence
- Chain Rule Decomposition: Overall acceptance depends on multiplying individual position acceptance rates; fixing weak positions improves end-to-end throughput
- Training with Simulated Error: Specialists learn from features generated by predecessors, simulating inference conditions
Implementation
Step 1: Analyze Position-Specific Acceptance Degradation
import torch
import numpy as np
from typing import List, Dict
from collections import defaultdict
class PositionAcceptanceAnalyzer:
"""Diagnose where draft models fail in speculative decoding"""
def __init__(self, draft_model, target_model):
self.draft_model = draft_model
self.target_model = target_model
def compute_position_wise_acceptance(self,
prompts: List[torch.Tensor],
num_draft_tokens: int = 4) -> Dict:
"""
For each position k in the draft, measure acceptance rate.
Key finding: pos-acc rapidly deteriorates beyond k=1.
"""
position_stats = defaultdict(list)
for prompt in prompts:
draft_tokens = self.draft_model.generate_tokens(
prompt, num_tokens=num_draft_tokens
)
for pos in range(num_draft_tokens):
draft_token_at_pos = draft_tokens[pos]
target_logits = self.target_model.forward(
torch.cat([prompt, draft_tokens[:pos]], dim=0)
)
target_top_token = torch.argmax(target_logits[-, :])
accepted = (draft_token_at_pos == target_top_token)
position_stats[pos].append((accepted))
position_acceptance_rates = {}
pos (num_draft_tokens):
acc_rate = np.mean(position_stats[pos])
position_acceptance_rates[pos] = acc_rate
()
overall_acceptance = np.prod((position_acceptance_rates.values()))
()
position_acceptance_rates
() -> :
position_deviations = defaultdict()
prompt prompts:
draft_hidden = .draft_model.extract_hidden_states(prompt)
target_hidden = .target_model.extract_hidden_states(prompt)
pos ((draft_hidden) - ):
deviation = torch.norm(
draft_hidden[pos] - target_hidden[pos]
)
position_deviations[pos].append(deviation.item())
avg_deviations = {}
pos ((draft_hidden) - ):
avg_dev = np.mean(position_deviations[pos])
avg_deviations[pos] = avg_dev
()
avg_deviations
Step 2: Design Position-Specialist Architecture
class PositionSpecialist(torch.nn.Module):
"""Transformer layer specialized for a specific position"""
def __init__(self, hidden_dim: int, num_layers: int,
position_id: int, expected_deviation: float):
super().__init__()
self.position_id = position_id
self.expected_deviation = expected_deviation
adapted_layers = max(2, int(num_layers * expected_deviation / 0.5))
self.specialist_layers = torch.nn.ModuleList([
self.TransformerBlock(hidden_dim)
for _ in range(adapted_layers)
])
self.token_loss_weight = 0.7
self.feature_loss_weight = 0.2
self.topk_loss_weight = 0.1
class TransformerBlock(torch.nn.Module):
def __init__(self, hidden_dim):
super().__init__()
self.self_attn = torch.nn.MultiheadAttention(
embed_dim=hidden_dim, num_heads=8
)
self.ffn = torch.nn.Sequential(
torch.nn.Linear(hidden_dim, * hidden_dim),
torch.nn.ReLU(),
torch.nn.Linear( * hidden_dim, hidden_dim),
)
.norm1 = torch.nn.LayerNorm(hidden_dim)
.norm2 = torch.nn.LayerNorm(hidden_dim)
():
x = x + .self_attn(x, x, x)[]
x = .norm1(x)
x = x + .ffn(x)
x = .norm2(x)
x
() -> torch.Tensor:
x = hidden_state
layer .specialist_layers:
x = layer(x)
x
(torch.nn.Module):
():
().__init__()
.num_specialists = num_specialists
.specialists = torch.nn.ModuleList()
position_deviations = .estimate_position_deviations(base_model)
pos (num_specialists):
specialist = PositionSpecialist(
hidden_dim=base_model.hidden_dim,
num_layers=base_model.num_layers,
position_id=pos,
expected_deviation=position_deviations.get(pos, )
)
.specialists.append(specialist)
() -> :
deviations = {}
pos (.num_specialists):
deviations[pos] = * ( ** pos)
deviations
() -> torch.Tensor:
specialist_idx = (position, .num_specialists - )
specialist = .specialists[specialist_idx]
output = specialist(hidden_states)
output
Step 3: Training Position Specialists
class PositionSpecialistTrainer:
"""Train specialists using simulated inference conditions"""
def __init__(self, target_model, specialist_ensemble: PositionSpecialistEnsemble):
self.target_model = target_model
self.specialists = specialist_ensemble
self.optimizer = torch.optim.AdamW(
self.specialists.parameters(), lr=2e-4
)
def train_on_batch(self, prompts: torch.Tensor,
completions: torch.Tensor) -> Dict[str, float]:
"""
Train specialists on data with three loss components:
1. Token-level: predict correct next token
2. Feature-level: match target model features
3. Top-K: preserve top-K token distribution
"""
losses = {}
target_hidden = self.target_model.extract_hidden_states(
torch.cat([prompts, completions], dim=1)
)
target_logits = self.target_model.forward(
torch.cat([prompts, completions], dim=1)
)
total_loss = 0
for pos in range(self.specialists.num_specialists):
if pos == 0:
current_hidden = target_hidden[-1, :].unsqueeze(0)
:
prev_output = .specialists.forward(
target_hidden[-, :].unsqueeze(),
position=pos -
)
current_hidden = prev_output
specialist_output = .specialists.forward(
current_hidden,
position=pos
)
token_loss = .compute_token_loss(
specialist_output,
target_logits[pos],
)
feature_loss = .compute_feature_loss(
specialist_output,
target_hidden[pos],
)
topk_loss = .compute_topk_loss(
specialist_output,
target_logits[pos],
k=
)
position_loss = (
* token_loss +
* feature_loss +
* topk_loss
)
total_loss += position_loss
losses[] = position_loss.item()
.optimizer.zero_grad()
total_loss.backward()
torch.nn.utils.clip_grad_norm_(.specialists.parameters(), )
.optimizer.step()
losses[] = total_loss.item()
losses
() -> torch.Tensor:
torch.nn.functional.cross_entropy(
specialist_logits.view(-, specialist_logits.size(-)),
target_logits.argmax(-)
)
() -> torch.Tensor:
torch.nn.functional.mse_loss(
specialist_hidden,
target_hidden.detach()
)
() -> torch.Tensor:
specialist_probs = torch.nn.functional.softmax(specialist_logits, dim=-)
target_probs = torch.nn.functional.softmax(target_logits, dim=-)
_, topk_indices = torch.topk(target_probs, k)
specialist_topk = specialist_probs[topk_indices]
target_topk = target_probs[topk_indices]
torch.nn.functional.kl_div(
torch.log(specialist_topk + ),
target_topk.detach(),
reduction=
)
Step 4: Integration with Speculative Decoding
class SpeculativeDecodingWithPosS:
"""Enhanced speculative decoding using position specialists"""
def __init__(self, target_model, draft_model,
position_specialists: PositionSpecialistEnsemble):
self.target_model = target_model
self.draft_model = draft_model
self.specialists = position_specialists
def generate_with_verification(self, prompt: torch.Tensor,
max_length: int = 100,
num_draft_tokens: int = 4) -> torch.Tensor:
"""
Speculative decoding with position-specialized draft:
1. Use specialist-enhanced draft to predict k tokens
2. Verify with target model
3. Accept accepted tokens, resample rejected ones
"""
generated = prompt.clone()
total_draft_tokens = 0
total_verified_tokens = 0
while generated.shape[0] < max_length:
draft_tokens = []
for pos in range(num_draft_tokens):
draft_logits = self.draft_model.forward(generated)
draft_hidden = self.draft_model.extract_hidden_state(generated)
enhanced_hidden = self.specialists.forward(
draft_hidden, position=pos
)
enhanced_logits = self.draft_model.decode(enhanced_hidden)
blended_logits = * enhanced_logits + * draft_logits
draft_token = torch.argmax(blended_logits)
draft_tokens.append(draft_token)
total_draft_tokens +=
draft_sequence = torch.cat([generated, torch.stack(draft_tokens)])
target_logits = .target_model.forward(draft_sequence)
accepted_count =
pos, draft_token (draft_tokens):
target_token = torch.argmax(target_logits[-((draft_tokens)-pos)])
draft_token == target_token:
generated = torch.cat([generated, draft_token.unsqueeze()])
accepted_count +=
total_verified_tokens +=
:
target_probs = torch.softmax(target_logits[-((draft_tokens)-pos)], dim=-)
new_token = torch.multinomial(target_probs, )
generated = torch.cat([generated, new_token])
total_verified_tokens +=
acceptance_length = total_verified_tokens / total_draft_tokens
speedup_ratio = total_verified_tokens / (total_verified_tokens / num_draft_tokens + total_verified_tokens)
()
()
generated
Practical Guidance
-
Measure Pos-Acc First: Profile your draft model with position-wise acceptance analysis. Identify which positions have lowest acceptance rates—these are bottlenecks.
-
Specialist Count: Start with 4-8 specialists covering 4-8 draft positions. More specialists provide finer granularity but increase computational overhead.
-
Feature Deviation Estimation: Position deviation grows roughly exponentially. Allocate more parameters to later positions which have higher deviation.
-
Training Strategy: Train specialists with simulated inference conditions using previous specialist outputs, not clean ground truth.
-
Three-Part Loss: Use token-level (70%), feature-level (20%), and top-K (10%) losses. This balances correctness with feature alignment.
-
Integration: Enhance draft model predictions by routing through specialists, then blend enhanced logits with base draft logits (60-40 weighting works well).
Reference
- Paper: PosS (2506.03566)
- Key Metric: Position-wise acceptance rate (pos-acc) analysis
- Improvements: 4.5% on acceptance length, 5.7% on speed-up ratio
- Architecture: Multiple position-specialized transformer layers trained with simulated error accumulation