class TeacherAgent(nn.Module):
"""Generates Socratic questions to guide prompt refinement."""
def __init__(self, hidden_dim=1024):
super().__init__()
self.question_generator = nn.Sequential(
nn.Linear(hidden_dim * 2, hidden_dim),
nn.GELU(),
nn.Linear(hidden_dim, hidden_dim)
)
def forward(self, task_state: torch.Tensor, current_objective: str) -> torch.Tensor:
"""
Args:
task_state: (B, D) current task state
current_objective: current sub-goal description
Returns:
question_embedding: (B, D) Socratic question embedding
"""
objective_embedding = self._encode_text(current_objective)
combined = torch.cat([task_state, objective_embedding], dim=1)
question = self.question_generator(combined)
return question
def _encode_text(self, text: str) -> torch.Tensor:
"""Simple text encoding; in practice use LLM embeddings."""
return torch.randn(1, 1024)
class CriticAgent(nn.Module):
"""Evaluates question quality and prompt coherence."""
def __init__(self, hidden_dim=1024):
super().__init__()
self.quality_scorer = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, 1),
nn.Sigmoid()
)
self.coherence_scorer = nn.Linear(hidden_dim, 1)
def forward(self, question: torch.Tensor, prompt: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Args:
question: (B, D) question embedding
prompt: (B, D) current prompt embedding
Returns:
quality_score: (B, 1) question quality [0, 1]
coherence_score: (B, 1) prompt coherence [0, 1]
"""
quality = self.quality_scorer(question)
coherence = torch.sigmoid(self.coherence_scorer(prompt))
return quality, coherence
class StudentAgent(nn.Module):
"""Maintains state and generates refined prompts."""
def __init__(self, hidden_dim=1024):
super().__init__()
self.state_update = nn.GRUCell(hidden_dim, hidden_dim)
self.prompt_generator = nn.Linear(hidden_dim, hidden_dim)
self.dialogue_memory = nn.Linear(hidden_dim, hidden_dim)
def forward(
self,
current_state: torch.Tensor,
question: torch.Tensor,
critic_feedback: torch.Tensor,
dialogue_history: List[torch.Tensor]
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Args:
current_state: (B, D) student's current state
question: (B, D) teacher's Socratic question
critic_feedback: (B, 2) [quality, coherence] scores
dialogue_history: list of previous interaction embeddings
Returns:
new_state: (B, D) updated student state
refined_prompt: (B, D) refined prompt embedding
"""
feedback_signal = critic_feedback[:, 0:1]
combined_input = question + feedback_signal * current_state
new_state = self.state_update(combined_input, current_state)
if dialogue_history:
history_embedding = torch.stack(dialogue_history).mean(dim=0)
memory_signal = self.dialogue_memory(history_embedding)
new_state = new_state + 0.3 * memory_signal
refined_prompt = self.prompt_generator(new_state)
return new_state, refined_prompt
class TargetAgent(nn.Module):
"""Evaluates refined prompts on downstream tasks."""
def __init__(self, hidden_dim=1024):
super().__init__()
self.task_evaluator = nn.Sequential(
nn.Linear(hidden_dim * 2, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, 1),
nn.Sigmoid()
)
def forward(self, prompt: torch.Tensor, task_embedding: torch.Tensor) -> torch.Tensor:
"""
Args:
prompt: (B, D) candidate prompt
task_embedding: (B, D) task representation
Returns:
reward: (B, 1) performance score [0, 1]
"""
combined = torch.cat([prompt, task_embedding], dim=1)
reward = self.task_evaluator(combined)
return reward
def socratic_optimization_loop(
initial_prompt: str,
task_embedding: torch.Tensor,
planner: PlannerAgent,
teacher: TeacherAgent,
critic: CriticAgent,
student: StudentAgent,
target: TargetAgent,
max_iterations: int = 10,
early_stopping_delta: float = 0.01
):
"""
Orchestrates the full MARS optimization loop.
Args:
initial_prompt: starting prompt
task_embedding: (B, D) task representation
Other agents: defined above
max_iterations: maximum refinement steps
early_stopping_delta: convergence threshold
Returns:
best_prompt: optimized prompt
optimization_trajectory: history of improvements
"""
student_state = torch.randn_like(task_embedding)
dialogue_history = []
best_reward = 0
optimization_trajectory = []
sub_goals, planning_state = planner(task_embedding)
for iteration in range(max_iterations):
current_objective = sub_goals[min(iteration, len(sub_goals) - 1)]
question = teacher(student_state, current_objective)
prompt_embedding = torch.randn(task_embedding.shape[0], 1024)
quality, coherence = critic(question, prompt_embedding)
feedback = torch.cat([quality, coherence], dim=1)
new_state, refined_prompt = student(
student_state, question, feedback, dialogue_history
)
dialogue_history.append(refined_prompt)
reward = target(refined_prompt, task_embedding)
improvement = reward - best_reward
optimization_trajectory.append({
'iteration': iteration,
'reward': reward.item(),
'improvement': improvement.item()
})
best_reward = max(best_reward, reward)
student_state = new_state
if improvement < early_stopping_delta and iteration > 2:
break
return refined_prompt, optimization_trajectory