| name | segment-policy-optimization |
| title | Segment Policy Optimization: Effective Segment-Level Credit Assignment in RL for Large Language Models |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2505.23564 |
| keywords | ["reinforcement-learning","credit-assignment","language-models","policy-optimization","reasoning-tasks"] |
| description | Segment-level credit assignment for RL in LLMs using Monte Carlo advantage estimation, enabling precise reward attribution without critic models for improved reasoning task performance. |
Segment Policy Optimization
Core Concept
Segment Policy Optimization (SPO) addresses a fundamental challenge in reinforcement learning for large language models: effective credit assignment at the right granularity. Traditional approaches suffer from either too-coarse trajectory-level estimation or too-noisy token-level signals. SPO operates at an intermediate "segment" granularity—grouping tokens between decision points—enabling more accurate advantage estimation without requiring a separate critic model.
Architecture Overview
- Segmentation Strategy: Partitions sequences by identifying low-probability tokens as natural cutpoints where the policy diverges significantly, creating meaningful reasoning segments
- Monte Carlo Estimation: Computes unbiased segment-level advantages from sampled trajectories without critic dependencies
- Tree-Based Sampling: For long-horizon tasks, organizes samples hierarchically to enable efficient reuse and advantage propagation
- Probability Masking: Selectively applies advantages only to tokens within segments where uncertainty was highest, focusing optimization effort on critical decision points
- Dual Instantiations: Provides chain-based variant for short reasoning tasks and tree-based variant for long-horizon problems
Implementation
The following pseudo-code illustrates the core SPO algorithm:
import numpy as np
from typing import List, Tuple
class SegmentPolicyOptimizer:
def __init__(self, model, policy_lr=1e-5, discount_gamma=0.99):
self.model = model
self.policy_lr = policy_lr
self.gamma = discount_gamma
def identify_segments(self, token_logits: np.ndarray, threshold=0.1) -> List[int]:
probs = np.exp(token_logits)
cutpoints = []
i ((probs)):
probs[i] < threshold:
cutpoints.append(i)
cutpoints.append((probs))
cutpoints
() -> :
start, end = segment_bounds
segment_return =
t (start, end):
segment_return += (.gamma ** (t - start)) * rewards[t]
segment_return += (.gamma ** (end - start)) * next_value
baseline = np.mean(rewards[start:end])
advantage = segment_return - baseline
advantage
() -> :
policy_loss =
seg_idx, (start, end) (segments):
advantage = advantages[seg_idx]
token_idx (start, end):
mask = / ((log_probs[token_idx]) + )
policy_loss -= mask * log_probs[token_idx] * advantage
policy_loss /= (segments)
(policy_loss)