class ProcessReasoner(nn.Module):
"""Generate explicit reasoning about task progress."""
def __init__(self, embed_dim=768, max_reasoning_steps=5):
super().__init__()
self.embed_dim = embed_dim
self.max_reasoning_steps = max_reasoning_steps
self.reasoning_encoder = nn.TransformerEncoderLayer(
d_model=embed_dim,
nhead=8,
dim_feedforward=2048,
batch_first=True
)
self.step_generator = nn.GRUCell(embed_dim, embed_dim)
self.step_head = nn.Linear(embed_dim, embed_dim)
self.progress_head = nn.Sequential(
nn.Linear(embed_dim, 256),
nn.GELU(),
nn.Linear(256, 1),
nn.Sigmoid()
)
def generate_reasoning_chain(self, anchored_context: torch.Tensor,
max_steps: int = None) -> (str, List[float]):
"""
Generate reasoning steps explaining progress.
Args:
anchored_context: [num_frames, embed_dim] structured video
max_steps: max reasoning steps (default: self.max_reasoning_steps)
"""
if max_steps is None:
max_steps = self.max_reasoning_steps
encoded = self.reasoning_encoder(anchored_context.unsqueeze(0))
context_repr = encoded.mean(dim=1)
reasoning_steps = []
progress_scores = []
hidden = context_repr.squeeze(0)
for step_idx in range(max_steps):
progress = self.progress_head(hidden)
progress_scores.append(progress.item())
step_embed = self.step_generator(context_repr, hidden.unsqueeze(0))
step_repr = self.step_head(step_embed)
reasoning_steps.append(step_repr.detach())
hidden = step_embed
return reasoning_steps, progress_scores
def forward(self, anchored_context: torch.Tensor) -> (List[str], List[float]):
"""Full reasoning generation."""
reasoning_steps, progress_scores = self.generate_reasoning_chain(
anchored_context)
return reasoning_steps, progress_scores
import torch.optim as optim
class ProcessSupervisionRL:
"""Train PRIMO R1 with outcome-based reinforcement learning."""
def __init__(self, model, process_reasoner, embed_dim=768):
self.model = model
self.reasoner = process_reasoner
self.optimizer = optim.AdamW(model.parameters(), lr=1e-5)
self.reasoner_optimizer = optim.AdamW(process_reasoner.parameters(),
lr=1e-5)
def compute_outcome_reward(self, video: torch.Tensor,
ground_truth_result: bool) -> float:
"""Evaluate whether final outcome matches ground truth."""
final_frame = video[:, -1, :]
with torch.no_grad():
outcome_pred = self.model.predict_outcome(final_frame)
reward = 1.0 if (outcome_pred > 0.5) == ground_truth_result else 0.0
return reward
def compute_process_reward(self, reasoning_steps: List[torch.Tensor],
progress_scores: List[float],
ground_truth_trajectory: List[bool]) -> float:
"""Evaluate quality of intermediate reasoning."""
if not reasoning_steps or not ground_truth_trajectory:
return 0.0
monotonicity_bonus = 0.0
for i in range(1, len(progress_scores)):
if ground_truth_trajectory[i] and not ground_truth_trajectory[i-1]:
if progress_scores[i] > progress_scores[i-1]:
monotonicity_bonus += 0.1
for i in range(1, len(progress_scores)):
jump = abs(progress_scores[i] - progress_scores[i-1])
if jump > 0.3:
monotonicity_bonus -= 0.05
return monotonicity_bonus
def training_step(self, video: torch.Tensor,
ground_truth_result: bool,
ground_truth_trajectory: List[bool] = None):
"""One RL training step."""
temporal_anchor = TemporalAnchorRepresentation()
initial_frame = video[:, 0, :]
current_frame = video[:, -1, :]
anchored = temporal_anchor.create_anchored_context(
initial_frame, current_frame)
reasoning_steps, progress_scores = self.reasoner(anchored)
outcome_reward = self.compute_outcome_reward(video, ground_truth_result)
process_reward = 0.0
if ground_truth_trajectory:
process_reward = self.compute_process_reward(
reasoning_steps, progress_scores, ground_truth_trajectory)
total_reward = 0.7 * outcome_reward + 0.3 * process_reward
reasoning_logprobs = [torch.log(step.norm()) for step in reasoning_steps]
reasoning_logprob = sum(reasoning_logprobs) / len(reasoning_logprobs)
loss = -total_reward * reasoning_logprob
self.reasoner_optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(self.reasoner.parameters(), 1.0)
self.reasoner_optimizer.step()
return loss.item(), total_reward
def train(self, training_videos: List[torch.Tensor],
ground_truth_results: List[bool],
ground_truth_trajectories: List[List[bool]] = None,
num_epochs: int = 10):
"""Full training loop."""
for epoch in range(num_epochs):
total_loss = 0
total_reward = 0
for i, video in enumerate(training_videos):
gt_result = ground_truth_results[i]
gt_traj = ground_truth_trajectories[i] if ground_truth_trajectories else None
loss, reward = self.training_step(video, gt_result, gt_traj)
total_loss += loss
total_reward += reward
avg_loss = total_loss / len(training_videos)
avg_reward = total_reward / len(training_videos)
print(f"Epoch {epoch+1}: Loss={avg_loss:.4f}, "
f"Reward={avg_reward:.3f}")
def evaluate_primo_r1(model, reasoner, test_videos, ground_truth):
"""Benchmark PRIMO R1 on robotic task evaluation."""
correct = 0
error_reduction = []
for video, ground_truth_result in zip(test_videos, ground_truth):
temporal_anchor = TemporalAnchorRepresentation()
initial = video[:, 0, :]
current = video[:, -1, :]
anchored = temporal_anchor.create_anchored_context(initial, current)
reasoning_steps, progress_scores = reasoner(anchored)
final_frame = video[:, -1, :]
prediction = model.predict_outcome(final_frame)
if (prediction > 0.5) == ground_truth_result:
correct += 1
accuracy = correct / len(test_videos)
print(f"PRIMO R1 Accuracy: {accuracy:.1%}")
print(f"Target: 67% on RoboFail benchmark")
return accuracy