import torch
import torch.nn as nn
from torch.utils.data import DataLoader, TensorDataset
class ExitPredictor(nn.Module):
"""Predicts whether answer has arrived based on partial trace."""
def __init__(self, embed_dim=768, hidden_dim=1024):
super().__init__()
self.embedding = nn.Embedding(50257, embed_dim)
self.encoder = nn.TransformerEncoderLayer(
d_model=embed_dim,
nhead=8,
dim_feedforward=hidden_dim,
batch_first=True
)
self.confidence_head = nn.Sequential(
nn.Linear(embed_dim, hidden_dim),
nn.GELU(),
nn.Linear(hidden_dim, 256),
nn.GELU(),
nn.Linear(256, 1),
nn.Sigmoid()
)
def forward(self, token_ids, partial_lengths):
"""
Args:
token_ids: [batch, max_seq_len] partial reasoning token IDs
partial_lengths: [batch] actual length of each partial trace
"""
embeddings = self.embedding(token_ids)
mask = torch.arange(token_ids.size(1)).unsqueeze(0) < partial_lengths.unsqueeze(1)
encoded = self.encoder(embeddings, src_key_padding_mask=~mask)
last_token_indices = (partial_lengths - 1).clamp(min=0)
last_embeddings = encoded[torch.arange(encoded.size(0)),
last_token_indices]
exit_confidence = self.confidence_head(last_embeddings)
return exit_confidence.squeeze(-1)
def train_exit_predictor(arrivals, model, tokenizer, epochs=10, batch_size=32):
"""Train predictor on answer-arrival data."""
X = []
y = []
for arrival in arrivals:
full_tokens = tokenizer.encode(arrival['full_trace'])
answer_pos = arrival['answer_token_position']
for completion_pct in [0.2, 0.4, 0.6, 0.8]:
current_pos = int(len(full_tokens) * completion_pct)
partial_tokens = full_tokens[:current_pos]
X.append(torch.tensor(partial_tokens))
y.append(1.0 if current_pos >= answer_pos else 0.0)
max_len = max(len(x) for x in X)
X_padded = torch.zeros(len(X), max_len, dtype=torch.long)
lengths = []
for i, x in enumerate(X):
X_padded[i, :len(x)] = x
lengths.append(len(x))
y_tensor = torch.tensor(y, dtype=torch.float32)
lengths_tensor = torch.tensor(lengths, dtype=torch.long)
dataset = TensorDataset(X_padded, lengths_tensor, y_tensor)
loader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
predictor = ExitPredictor()
optimizer = torch.optim.Adam(predictor.parameters(), lr=1e-4)
criterion = nn.BCELoss()
for epoch in range(epochs):
total_loss = 0
for X_batch, lengths_batch, y_batch in loader:
logits = predictor(X_batch, lengths_batch)
loss = criterion(logits, y_batch)
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
print(f"Epoch {epoch+1}: Loss = {total_loss/len(loader):.4f}")
return predictor
def generate_with_early_exit(model, prompt, predictor, tokenizer,
exit_threshold=0.7, max_tokens=2000):
"""Generate with learned early stopping."""
generated_tokens = []
token_ids = tokenizer.encode(prompt)
for step in range(max_tokens):
with torch.no_grad():
logits = model(torch.tensor([token_ids]))
next_token = torch.argmax(logits[-1, -1, :])
generated_tokens.append(next_token.item())
if step > 50 and step % 10 == 0:
partial_tokens = torch.tensor([token_ids + generated_tokens])
partial_length = torch.tensor([len(token_ids) + len(generated_tokens)])
with torch.no_grad():
exit_conf = predictor(partial_tokens, partial_length)
print(f"Step {step}: Exit confidence = {exit_conf:.2%}")
if exit_conf > exit_threshold:
print(f"Early exit at step {step} with confidence {exit_conf:.2%}")
break
token_ids.append(next_token.item())
if next_token == tokenizer.eos_token_id:
break
return tokenizer.decode(generated_tokens)
def benchmark_early_exit(model, predictor, tokenizer, test_tasks,
thresholds=[0.5, 0.6, 0.7, 0.8]):
"""Measure accuracy vs token reduction across thresholds."""
results = []
for threshold in thresholds:
total_tokens = 0
correct = 0
for task in test_tasks:
output = generate_with_early_exit(model, task['prompt'],
predictor, tokenizer,
exit_threshold=threshold)
total_tokens += len(tokenizer.encode(output))
if task['answer'].lower() in output.lower():
correct += 1
avg_tokens = total_tokens / len(test_tasks)
accuracy = correct / len(test_tasks)
results.append({
'threshold': threshold,
'accuracy': accuracy,
'avg_tokens': avg_tokens
})
print(f"Threshold {threshold}: Accuracy={accuracy:.1%}, "
f"Avg Tokens={avg_tokens:.0f}")
return results