Stereo-CoRE / code /stereo_core /pair_predictive_model.py
B111ue's picture
Add files using upload-large-folder tool
20b0922 verified
Raw
History Blame Contribute Delete
9.07 kB
"""Predictability-Aware Agent Interaction Routing (PAIR).
The policy remains strictly local at deployment. A training-only teacher sees
synchronised action sets and the future *local* consequence of each arm. The
policy is not regressed to an arbitrary privileged role label; instead, its
local interaction representation is contrastively aligned to the part of the
teacher event that is predictable from current wrist RGB-D and qpos.
"""
from __future__ import annotations
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from stereo_decoder_variants import StereoARCA
class StereoPredictabilityPAIR(StereoARCA):
"""Local-ARCA plus a locally inferred interaction-event representation.
``interaction_to_query`` is zero-initialised. Loading a Local-ARCA
checkpoint therefore starts from the exact deployed baseline and lets the
new signal prove useful instead of perturbing the policy at update zero.
"""
def __init__(self, *args, event_dim=128, **kwargs):
super().__init__(*args, **kwargs)
d = self.query.shape[-1]
self.event_dim = event_dim
self.local_event_head = nn.Sequential(
nn.LayerNorm(d),
nn.Linear(d, d),
nn.GELU(),
nn.Linear(d, event_dim),
)
self.interaction_to_query = nn.Linear(event_dim, d, bias=False)
nn.init.zeros_(self.interaction_to_query.weight)
self.last_local_event = None
self.last_gates = None
def _route_with_event(self, state, observation, batch):
query = self.query.expand(batch, -1, -1)
context = self.route_state(state) + self.route_observation(observation.mean(1))
local_event = F.normalize(self.local_event_head(context).float(), dim=-1, eps=1e-6)
event_bias = self.interaction_to_query(local_event.to(context.dtype))
features = self.route_mlp(query + context.unsqueeze(1) + event_bias.unsqueeze(1))
logits = torch.matmul(features, self.role_prototypes.t()) / math.sqrt(features.shape[-1])
values, ids = logits.topk(2, dim=-1)
gates = torch.zeros_like(logits).scatter_(-1, ids, values.softmax(-1).to(logits.dtype))
importance = logits.softmax(-1).mean((0, 1))
load = gates.gt(0).to(logits.dtype).mean((0, 1)) / 2.0
auxiliary = self.roles_n * (importance * load).sum()
return gates, auxiliary, local_event
def forward(self, image, depth_mm, qpos, actions=None):
observation = self._rgbd_tokens(image, depth_mm)
state_vector = self.state(qpos)
gates, auxiliary, local_event = self._route_with_event(
state_vector, observation, image.shape[0]
)
self.last_local_event, self.last_gates = local_event, gates
if actions is not None:
hidden = self.posterior(self.action(actions) + self.pos)
mu, logvar = self.latent(hidden.mean(1)).chunk(2, -1)
logvar = logvar.clamp(-10.0, 5.0)
latent = mu + torch.randn_like(mu) * torch.exp(0.5 * logvar)
else:
mu = logvar = None
latent = torch.zeros(
(image.shape[0], self.z_proj.in_features), device=image.device
)
memory = torch.cat(
(
state_vector.unsqueeze(1),
self.z_proj(latent).unsqueeze(1),
observation,
),
dim=1,
)
decoded = self.decoder(
self.query.expand(image.shape[0], -1, -1),
memory,
observation,
gates,
)
return self.out(decoded), mu, logvar, auxiliary, local_event
class PredictiveInteractionTeacher(nn.Module):
"""Training-only event teacher over unordered synchronised teams.
The teacher is permutation equivariant: each target arm is represented by
its own action/consequence and the mean of the other arms' action features.
No task name, agent identifier, global view, or simulator state is used.
"""
def __init__(self, action_dim, effect_dim, event_dim=128, width=192):
super().__init__()
self.event_dim = event_dim
self.action_step = nn.Sequential(
nn.Linear(action_dim, width),
nn.GELU(),
nn.Linear(width, width),
)
self.action_feature = nn.Sequential(
nn.LayerNorm(2 * width),
nn.Linear(2 * width, width),
nn.GELU(),
nn.Linear(width, width),
)
self.effect_feature = nn.Sequential(
nn.LayerNorm(effect_dim),
nn.Linear(effect_dim, width),
nn.GELU(),
nn.Linear(width, width),
)
self.event = nn.Sequential(
nn.LayerNorm(3 * width),
nn.Linear(3 * width, width),
nn.GELU(),
nn.Linear(width, event_dim),
)
self.own_effect = nn.Sequential(
nn.Linear(width, width),
nn.GELU(),
nn.Linear(width, effect_dim),
)
self.interaction_effect = nn.Sequential(
nn.Linear(event_dim, width),
nn.GELU(),
nn.Linear(width, effect_dim),
)
self.sync_score = nn.Sequential(
nn.LayerNorm(3 * width),
nn.Linear(3 * width, width),
nn.GELU(),
nn.Linear(width, 1),
)
def action_features(self, actions):
delta = actions[:, 1:] - actions[:, :-1]
encoded = self.action_step(delta.float())
return self.action_feature(
torch.cat((encoded.mean(1), encoded.square().mean(1).sqrt()), dim=-1)
)
@staticmethod
def peer_means(features, groups):
peers = torch.zeros_like(features)
for group in groups.unique(sorted=True):
ids = (groups == group).nonzero(as_tuple=False).flatten()
values = features.index_select(0, ids)
if len(ids) > 1:
peer = ((values.sum(0, keepdim=True) - values) / (len(ids) - 1)).to(
features.dtype
)
else:
peer = torch.zeros_like(values)
peers.index_copy_(0, ids, peer)
return peers
@staticmethod
def shuffled_peer_means(peer, groups):
unique = groups.unique(sorted=True)
if len(unique) < 2:
return peer.roll(1, dims=0)
shuffled = peer.clone()
# Shift complete time groups, preserving within-group team structure.
for index, group in enumerate(unique):
source_group = unique[(index + 1) % len(unique)]
target_ids = (groups == group).nonzero(as_tuple=False).flatten()
source_ids = (groups == source_group).nonzero(as_tuple=False).flatten()
source = peer.index_select(0, source_ids)
if len(source) != len(target_ids):
source = source.mean(0, keepdim=True).expand(len(target_ids), -1)
shuffled.index_copy_(0, target_ids, source)
return shuffled
def forward(self, actions, future_effect, groups):
own = self.action_features(actions)
effect = self.effect_feature(future_effect.float())
peer = self.peer_means(own, groups)
teacher_event = F.normalize(
self.event(torch.cat((own, peer, effect), dim=-1)).float(),
dim=-1,
eps=1e-6,
)
own_prediction = self.own_effect(own)
full_prediction = own_prediction + self.interaction_effect(teacher_event)
own_loss = F.mse_loss(own_prediction, future_effect.float())
full_loss = F.mse_loss(full_prediction, future_effect.float())
true_logits = self.sync_score(torch.cat((own, peer, effect), dim=-1))
shuffled_peer = self.shuffled_peer_means(peer, groups)
false_logits = self.sync_score(torch.cat((own, shuffled_peer, effect), dim=-1))
sync_loss = 0.5 * (
F.binary_cross_entropy_with_logits(true_logits, torch.ones_like(true_logits))
+ F.binary_cross_entropy_with_logits(false_logits, torch.zeros_like(false_logits))
)
sync_accuracy = 0.5 * (
(true_logits > 0).float().mean() + (false_logits < 0).float().mean()
)
return teacher_event, own_loss, full_loss, sync_loss, sync_accuracy
def symmetric_contrastive_alignment(local_event, teacher_event, temperature=0.1):
"""Same-sample alignment with all other same-task/time samples as negatives."""
local = F.normalize(local_event.float(), dim=-1, eps=1e-6)
teacher = F.normalize(teacher_event.float(), dim=-1, eps=1e-6)
logits = local @ teacher.detach().t() / temperature
targets = torch.arange(len(local), device=local.device)
student_to_teacher = F.cross_entropy(logits, targets)
# The second direction stabilises the shared subspace without propagating
# gradients into the policy through the teacher branch.
teacher_to_student = F.cross_entropy(logits.t(), targets)
return 0.5 * (student_to_teacher + teacher_to_student)