Skip to main content Home Creators adu2021 skillxiv sparse-moe-agentic-intelligence
sparse-moe-agentic-intelligence Deploy frontier-level reasoning with only 11B active parameters using sparse MoE with 288 routed experts plus shared expert. Use Metropolis Independence Sampling-Filtered Policy Optimization (MIS-PO) to stabilize RL training at scale, replacing continuous importance weighting with discrete filtering that ensures trust region stability.
Jump to install Skills Marketplace Discover and explore AI skills built by the community.
Install with Codex or Claude Copy this prompt, paste it into Codex, Claude, or another assistant, and let it review the skill page and install it for you.
Copy promptShow prompt details A direct command skips the review prompt. Inspect the source before running it.
npx skills add https://github.com/ADu2021/skillXiv --skill sparse-moe-agentic-intelligenceThe command stays on one line. Scroll horizontally to inspect it before copying.
Prefer a local copy? Download the files currently available to SkillsMP.
Download Zip Downloading... More from this repository meaningful-kebab-case-name Convert arXiv papers into ready-to-use agent skills using category-aware extraction. First classifies the paper into one or more of 11 research categories, then applies a specialized extraction pipeline for each category — because different types of papers produce different types of usable knowledge. A single paper can yield multiple skills if it spans categories. Use this skill whenever the user wants to turn a paper into a skill, extract practical techniques from research, build a skill library from papers, convert arXiv papers into reusable agent instructions, or batch-process multiple papers into skills. Also trigger when someone asks about extracting actionable knowledge from papers, making research practical for LLM agents, or systematically converting academic contributions into structured agent capabilities.
action-quantization-behavior-cloning Establish regret bounds for behavior cloning with discretized actions combining statistical error and quantization error terms. Prove smoothness requirements for safe quantizer design, show that learning-based quantizers fail these requirements, and propose model-based augmentation to reduce error dependence from H² to H.
adaptive-lora-personalized-ranks Dynamically allocate LoRA ranks per-layer during fine-tuning instead of using fixed uniform ranks. Learn optimal rank for each layer and subject via variational framework with discretized exponential distribution, reducing memory footprint while maintaining fidelity and text-alignment.
Related occupations SOC
Based on SOC occupation classification
name sparse-moe-agentic-intelligence title Step 3.5 Flash: Open Frontier-Level Intelligence with 11B Active Parameters version 0.0.2 engine skillxiv-v0.0.2-claude-opus-4.6 license MIT url https://arxiv.org/abs/2602.10604 keywords ["Sparse Mixture-of-Experts","Efficient Reasoning","Multi-Token Prediction","MIS-PO","Agentic Intelligence"] description Deploy frontier-level reasoning with only 11B active parameters using sparse MoE with 288 routed experts plus shared expert. Use Metropolis Independence Sampling-Filtered Policy Optimization (MIS-PO) to stabilize RL training at scale, replacing continuous importance weighting with discrete filtering that ensures trust region stability.
Step 3.5 Flash: Open Frontier-Level Intelligence with 11B Active Parameters
Problem Context
Frontier reasoning models like o1 use massive compute budgets. Step 3.5 Flash achieves frontier-level performance with minimal active parameters (11B) by combining (1) sparse MoE with 288 experts per layer, (2) multi-token prediction for speculative decoding, (3) MIS-PO for stable RL training at scale. Key challenges: expert collapse, activation explosions, numerical instability in large sparse models.
Core Concept
The architecture uses: (1) 196B total parameters with only 11B active per token via selective expert routing, (2) 3:1 ratio of sliding-window to full attention (local + global reasoning), (3) multi-token prediction (MTP-3) enabling 3x faster decoding, (4) MIS-PO replacing PPO clipping with discrete filtering for stability. This achieves o1-competitive performance on reasoning benchmarks while remaining efficiently deployable.
Architecture Overview
Sparse MoE : 288 routed experts + 1 shared expert, activate 8 per token
Selective routing : Load-balanced expert assignment preventing collapse
Hybrid attention : Sliding-window (efficiency) + full attention (reasoning)
Multi-token prediction : Generate 3 tokens per forward pass
MIS-PO training : Discrete filtering for trust region RL
Monitoring infrastructure : Detect numerical issues, routing pathologies
Implementation
Step 1: Sparse MoE layer with selective routing
import torch
import torch.nn as nn
from typing import Tuple
class SelectiveExpertRouter :
"""Route tokens to subset of experts with load balancing."""
def __init__ (
self,
num_experts: int = 288 ,
active_experts: int = 8 ,
capacity_factor: float = 1.25
):
self .num_experts = num_experts
.active_experts = active_experts
.capacity_factor = capacity_factor
( ) -> [torch.Tensor, torch.Tensor, ]:
batch_size, seq_len, dim = hidden_states.shape
router_probs = torch.softmax(router_logits, dim=- )
top_k_probs, top_k_indices = torch.topk(
router_probs, k= .active_experts, dim=-
)
top_k_probs = top_k_probs / top_k_probs. (dim=- , keepdim= )
expert_load = torch.zeros( .num_experts, device=hidden_states.device)
i ( .active_experts):
expert_load.scatter_add_(
, top_k_indices[:, :, i].reshape(- ),
top_k_probs[:, :, i].reshape(- )
)
load_balancing_loss = ._compute_load_balance_loss(
expert_load, batch_size, seq_len
)
routing_stats = {
: load_balancing_loss.item(),
: expert_load.mean().item(),
: expert_load. ().item()
}
top_k_indices, top_k_probs, routing_stats
( ) -> torch.Tensor:
target_load = (batch_size * seq_len * .active_experts) / .num_experts
load_loss = torch.mean((expert_load - target_load) ** )
load_loss
(nn.Module):
( ):
().__init__()
.dim = dim
.num_experts = num_experts
.active_experts = active_experts
.experts = nn.ModuleList([
nn.Sequential(
nn.Linear(dim, * dim),
nn.GELU(),
nn.Linear( * dim, dim)
)
_ (num_experts)
])
.shared_expert = nn.Sequential(
nn.Linear(dim, * dim),
nn.GELU(),
nn.Linear( * dim, dim)
)
.router = nn.Linear(dim, num_experts)
.selector = SelectiveExpertRouter(num_experts, active_experts)
( ) -> [torch.Tensor, ]:
batch_size, seq_len, dim = hidden_states.shape
shared_output = .shared_expert(hidden_states)
router_logits = .router(hidden_states)
expert_indices, expert_weights, routing_stats = .selector.route_tokens(
hidden_states, router_logits
)
expert_outputs = []
expert .experts:
expert_outputs.append(expert(hidden_states))
expert_outputs = torch.stack(expert_outputs, dim=- )
batch_idx = torch.arange(batch_size, device=hidden_states.device)[:, , ]
seq_idx = torch.arange(seq_len, device=hidden_states.device)[ , :, ]
selected_outputs = expert_outputs[batch_idx, seq_idx, expert_indices]
weighted_outputs = (
selected_outputs * expert_weights.unsqueeze(- )
). (dim= )
output = weighted_outputs + shared_output
output, routing_stats
self
self
def
route_tokens
self,
hidden_states: torch.Tensor,
router_logits: torch.Tensor
Tuple
dict
"""
Route tokens to experts using top-k selection.
Args:
hidden_states: Input hidden states
router_logits: Routing scores from router network
Returns:
(routed_output, routing_mask, routing_stats)
"""
1
self
1
sum
1
True
self
for
in
range
self
0
1
1
self
'load_balancing_loss'
'mean_expert_load'
'max_expert_load'
max
return
def
_compute_load_balance_loss
self,
expert_load: torch.Tensor,
batch_size: int ,
seq_len: int
"""Auxiliary loss to balance expert utilization."""
self
self
2
return
class
SparseExpertLayer
"""Single sparse MoE layer with selective routing."""
def
__init__
self,
dim: int = 4096 ,
num_experts: int = 288 ,
active_experts: int = 8
super
self
self
self
self
4
4
for
in
range
self
4
4
self
self
def
forward
self, hidden_states: torch.Tensor
Tuple
dict
"""Route through selected experts."""
self
self
self
for
in
self
2
None
None
None
None
1
sum
2
return
Step 2: Multi-token prediction for acceleration class MultiTokenPredictor (nn.Module):
"""Predict multiple tokens per forward pass."""
def __init__ (self, vocab_size: int = 100000 , num_predict: int = 3 ):
super ().__init__()
self .num_predict = num_predict
self .vocab_size = vocab_size
self .heads = nn.ModuleList([
nn.Linear(768 , vocab_size) for _ in range (num_predict)
])
def forward (self, hidden_states: torch.Tensor ) -> Tuple [torch.Tensor, torch.Tensor]:
"""
Predict multiple tokens.
Args:
hidden_states: [batch, seq_len, dim]
Returns:
(logits, confidences): List of logit tensors and confidence per head
"""
logits_list = []
confidences = []
for head_idx, head in enumerate (self .heads):
logits = head(hidden_states)
logits_list.append(logits)
probs = torch.softmax(logits, dim=-1 )
max_prob, _ = torch.max (probs, dim=-1 )
confidences.append(max_prob)
return logits_list, confidences
def select_confident_predictions (
self,
logits_list: list ,
confidences: list ,
min_confidence: float = 0.8
) -> torch.Tensor:
"""
Select most confident multi-token predictions.
Returns selected token sequence.
"""
all_confidences = torch.stack(confidences, dim=-1 )
best_head_idx = torch.argmax(all_confidences, dim=-1 )
selected_logits = torch.zeros_like(logits_list[0 ])
for i in range (self .num_predict):
mask = (best_head_idx == i)
selected_logits[mask] = logits_list[i][mask]
return selected_logits
Step 3: MIS-PO (Metropolis Independence Sampling-Filtered Policy Optimization) class MetropolisIndependenceSamplingPO :
"""Discrete filtering policy optimization for RL at scale."""
def __init__ (self, model, optimizer, trust_region_threshold: float = 0.05 ):
self .model = model
self .optimizer = optimizer
self .trust_region_threshold = trust_region_threshold
def compute_mis_po_loss (
self,
log_probs: torch.Tensor,
log_probs_old: torch.Tensor,
rewards: torch.Tensor,
advantages: torch.Tensor
) -> torch.Tensor:
"""
Compute MIS-PO loss using discrete filtering.
Unlike PPO (continuous clipping), MIS-PO uses discrete filtering:
Keep update if log_prob_ratio falls within trust region, reject otherwise.
"""
log_prob_ratio = log_probs.sum (dim=1 ) - log_probs_old.sum (dim=1 )
within_trust_region = torch.abs (log_prob_ratio) < self .trust_region_threshold
ppo_loss = -(log_prob_ratio * advantages)[within_trust_region].mean()
rejection_rate = 1.0 - within_trust_region.float ().mean().item()
return ppo_loss, {'rejection_rate' : rejection_rate}
def training_step (
self,
batch: dict ,
reward_fn
) -> dict :
"""Single MIS-PO training step."""
prompts = batch['prompts' ]
batch_size = len (prompts)
responses = []
log_probs_list = []
for prompt in prompts:
response, log_probs = self .model.generate_with_logprobs(
prompt, max_tokens=500
)
responses.append(response)
log_probs_list.append(log_probs)
log_probs = torch.stack(log_probs_list)
rewards = torch.tensor([reward_fn(r) for r in responses])
group_mean = rewards.mean()
advantages = rewards - group_mean
log_probs_ref = log_probs.detach()
loss, stats = self .compute_mis_po_loss(log_probs, log_probs_ref, rewards, advantages)
self .optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(self .model.parameters(), 1.0 )
self .optimizer.step()
return {
'loss' : loss.item(),
'rejection_rate' : stats['rejection_rate' ],
'avg_reward' : rewards.mean().item()
}
Step 4: Monitoring for training stability class StabilityMonitor :
"""Monitor for expert collapse, activation explosions, numerical issues."""
def __init__ (self ):
self .history = {'expert_load' : [], 'activation_stats' : []}
def check_expert_collapse (self, routing_stats: dict , threshold: float = 0.8 ):
"""Check if some experts are unused."""
max_load = routing_stats['max_expert_load' ]
return max_load > threshold
def check_activation_explosion (
self,
hidden_states: torch.Tensor,
norm_threshold: float = 100.0
):
"""Check for exploding norms in hidden states."""
state_norms = torch.norm(hidden_states, dim=-1 )
max_norm = state_norms.max ().item()
return max_norm > norm_threshold
def check_numerical_stability (
self,
model: nn.Module
):
"""Check for NaN/Inf in parameters."""
for param in model.parameters():
if torch.isnan(param).any () or torch.isinf(param).any ():
return False
return True
def remediate_issues (self, model: nn.Module, issue_type: str ):
"""Apply remediation for detected issues."""
if issue_type == 'activation_explosion' :
for module in model.modules():
if isinstance (module, nn.Linear):
torch.clamp_(module.bias, -10 , 10 )
elif issue_type == 'numerical_instability' :
pass
Step 5: Full training integration def train_step_flash_model (
model,
train_loader,
verifier,
optimizer,
num_steps: int = 100000 ,
device: str = 'cuda'
):
"""
Train Step 3.5 Flash using MIS-PO.
"""
mis_po = MetropolisIndependenceSamplingPO(model, optimizer)
monitor = StabilityMonitor()
for step in range (num_steps):
batch = next (iter (train_loader))
def reward_fn (response ):
return float (verifier(response))
metrics = mis_po.training_step(batch, reward_fn)
with torch.no_grad():
for name, module in model.named_modules():
if isinstance (module, SparseExpertLayer):
if monitor.check_expert_collapse({'max_expert_load' : 0.9 }):
print (f"Step {step} : Expert collapse detected, applying remediation" )
if (step + 1 ) % 1000 == 0 :
print (f"Step {step + 1 } : "
f"Loss={metrics['loss' ]:.4 f} , "
f"Rejection={metrics['rejection_rate' ]:.2 %} , "
f"Reward={metrics['avg_reward' ]:.4 f} " )
return model
Practical Guidance When to use : Frontier reasoning at scale; resource-constrained deployments; reasoning-heavy workloads
num_experts : 128-512 (tradeoff: capacity vs. compute)
active_experts : 4-16 (active per token)
capacity_factor : 1.0-1.5 (load balancing)
trust_region_threshold : 0.03-0.1 (MIS-PO filtering)
multi_token_predict : 2-4 (acceleration vs. quality)
Frontier performance with 11B active parameters
Stable RL training via MIS-PO
Fast inference via multi-token prediction
Efficient hybrid attention (local + global)
Expert collapse without load balancing auxiliary loss
Activation explosions in deep sparse networks (use clipping)
MIS-PO threshold too loose → no actual filtering
Multi-token prediction conflicting token dependencies
Scaling : Linear in number of experts. Distributed routing enables 10K+ expert systems.
Reference Paper: https://arxiv.org/abs/2602.10604
Related work: Sparse MoE, policy optimization, multi-token prediction
Benchmarks: IMO-AnswerBench (85.4%), LiveCodeBench-v6 (86.4%)
Architecture: 196B total, 11B active, 288 experts/layer