summaryrefslogtreecommitdiff
path: root/sdil/babyai_shared.py
diff options
context:
space:
mode:
Diffstat (limited to 'sdil/babyai_shared.py')
-rw-r--r--sdil/babyai_shared.py299
1 files changed, 299 insertions, 0 deletions
diff --git a/sdil/babyai_shared.py b/sdil/babyai_shared.py
new file mode 100644
index 0000000..a4c871a
--- /dev/null
+++ b/sdil/babyai_shared.py
@@ -0,0 +1,299 @@
+"""Shared inference-and-teaching pathway for BabyAI imitation learning.
+
+The visual observation is basal. The language mission is projected into each
+hidden population through a fixed apical path. Manual updates implement exact
+BP, clean reciprocal KP, raw shared-path KP, or SDIL residualization.
+"""
+
+from dataclasses import asdict, dataclass
+import math
+import re
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+
+
+CONDITIONS = ("bp", "clean_kp", "raw_shared", "sdil")
+TOKEN_PATTERN = re.compile(r"[a-z0-9]+")
+
+
+def tokenize_mission(mission):
+ return TOKEN_PATTERN.findall(str(mission).lower())
+
+
+def build_vocabulary(missions):
+ tokens = sorted({token for mission in missions
+ for token in tokenize_mission(mission)})
+ return ("<unk>", *tokens)
+
+
+def missions_to_bow(missions, vocabulary):
+ lookup = {token: index for index, token in enumerate(vocabulary)}
+ result = np.zeros((len(missions), len(vocabulary)), dtype=np.float32)
+ for row, mission in enumerate(missions):
+ tokens = tokenize_mission(mission)
+ for token in tokens:
+ result[row, lookup.get(token, 0)] += 1.0
+ norm = np.linalg.norm(result[row])
+ if norm > 0:
+ result[row] /= norm
+ return result
+
+
+@dataclass(frozen=True)
+class BabyAISharedConfig:
+ input_dim: int
+ mission_dim: int
+ action_dim: int = 7
+ width: int = 256
+ hidden_layers: int = 2
+ learning_rate: float = 0.03
+ reciprocal_learning_rate: float = 0.03
+ momentum: float = 0.9
+ weight_decay: float = 1e-4
+ context_gain: float = 1.0
+
+ def to_dict(self):
+ return asdict(self)
+
+
+class BabyAISharedNet:
+ """Tanh policy with fixed mission projections and local manual updates."""
+
+ def __init__(self, config, seed, device="cpu", dtype=torch.float32):
+ if config.hidden_layers < 1:
+ raise ValueError("BabyAI shared policy needs a hidden population")
+ self.config = config
+ self.device = torch.device(device)
+ self.dtype = dtype
+ sizes = ([config.input_dim]
+ + [config.width] * config.hidden_layers
+ + [config.action_dim])
+ forward_generator = torch.Generator(device="cpu").manual_seed(seed)
+ reciprocal_generator = torch.Generator(device="cpu").manual_seed(seed + 1)
+ context_generator = torch.Generator(device="cpu").manual_seed(seed + 2)
+
+ def normal(shape, scale, generator):
+ value = torch.randn(*shape, generator=generator) * scale
+ return value.to(device=self.device, dtype=dtype)
+
+ self.W = [normal((sizes[i + 1], sizes[i]),
+ 1.0 / math.sqrt(sizes[i]), forward_generator)
+ for i in range(len(sizes) - 1)]
+ self.b = [torch.zeros(size, device=self.device, dtype=dtype)
+ for size in sizes[1:]]
+ # Q[i] has the same storage orientation as W[i]. Q[0] is not needed
+ # because no teaching signal is transported into the visual input.
+ self.Q = [None] + [normal(tuple(self.W[i].shape),
+ 1.0 / math.sqrt(sizes[i]),
+ reciprocal_generator)
+ for i in range(1, len(self.W))]
+ self.C = [normal((config.width, config.mission_dim),
+ config.context_gain
+ / math.sqrt(config.mission_dim),
+ context_generator)
+ for _ in range(config.hidden_layers)]
+
+ self.P = [torch.zeros(config.width, device=self.device, dtype=dtype)
+ for _ in range(config.hidden_layers)]
+ self.P_bias = [torch.zeros_like(value) for value in self.P]
+ self.mW = [torch.zeros_like(value) for value in self.W]
+ self.mb = [torch.zeros_like(value) for value in self.b]
+ self.mQ = [None] + [torch.zeros_like(value) for value in self.Q[1:]]
+
+ def clone(self):
+ copied = BabyAISharedNet(
+ self.config, seed=0, device=self.device, dtype=self.dtype)
+ for name in ("W", "b", "Q", "C", "P", "P_bias", "mW", "mb", "mQ"):
+ values = getattr(self, name)
+ setattr(copied, name, [
+ None if value is None else value.clone() for value in values])
+ return copied
+
+ def context_fields(self, missions, enabled=True):
+ if not enabled:
+ return [torch.zeros(
+ (missions.shape[0], self.config.width),
+ device=self.device, dtype=self.dtype) for _ in self.C]
+ return [missions @ projection.t() for projection in self.C]
+
+ def forward_features(self, features, missions, context_enabled=True):
+ contexts = self.context_fields(missions, enabled=context_enabled)
+ h = [features]
+ u = []
+ for layer in range(self.config.hidden_layers):
+ value = (h[-1] @ self.W[layer].t() + self.b[layer]
+ + contexts[layer])
+ u.append(value)
+ h.append(torch.tanh(value))
+ logits = h[-1] @ self.W[-1].t() + self.b[-1]
+ h.append(logits)
+ return {"h": h, "u": u, "context": contexts, "logits": logits}
+
+ def predictor(self, layer, soma):
+ return self.P[layer] * soma + self.P_bias[layer]
+
+ @torch.no_grad()
+ def fit_neutral_predictor(self, features, missions):
+ """Fit the per-cell neutral relation without actions or teaching."""
+ state = self.forward_features(features, missions)
+ reports = []
+ for layer, target in enumerate(state["context"]):
+ soma = state["h"][layer + 1]
+ soma_centered = soma - soma.mean(0)
+ target_centered = target - target.mean(0)
+ variance = soma_centered.square().mean(0)
+ covariance = (soma_centered * target_centered).mean(0)
+ slope = covariance / variance.clamp_min(1e-8)
+ intercept = target.mean(0) - slope * soma.mean(0)
+ self.P[layer].copy_(slope)
+ self.P_bias[layer].copy_(intercept)
+ prediction = slope * soma + intercept
+ residual = target - prediction
+ target_ss = target_centered.square().sum(0)
+ residual_ss = residual.square().sum(0)
+ valid = target_ss > 1e-12
+ r2 = 1.0 - residual_ss[valid] / target_ss[valid]
+ context_rms = target.square().mean().sqrt()
+ reports.append({
+ "mean_per_cell_r2": float(r2.mean()) if r2.numel() else 0.0,
+ "context_rms": float(context_rms),
+ "residual_context_rms_ratio": float(
+ residual.square().mean().sqrt()
+ / context_rms.clamp_min(1e-12)),
+ "neutral_observations": int(features.shape[0]),
+ "action_observations": 0,
+ "teaching_observations": 0,
+ })
+ return reports
+
+
+def visual_input_dim(object_cardinality, color_cardinality,
+ state_cardinality, view_size=7):
+ return (view_size * view_size
+ * (object_cardinality + color_cardinality + state_cardinality) + 4)
+
+
+def encode_visual(images, directions, object_cardinality,
+ color_cardinality, state_cardinality, device,
+ dtype=torch.float32):
+ """One-hot encode the compact MiniGrid symbolic observation."""
+ images = torch.as_tensor(images, device=device, dtype=torch.long)
+ directions = torch.as_tensor(directions, device=device, dtype=torch.long)
+ object_features = F.one_hot(
+ images[..., 0], num_classes=object_cardinality)
+ color_features = F.one_hot(
+ images[..., 1], num_classes=color_cardinality)
+ state_features = F.one_hot(
+ images[..., 2], num_classes=state_cardinality)
+ grid = torch.cat(
+ (object_features, color_features, state_features), dim=-1)
+ grid = grid.reshape(grid.shape[0], -1).to(dtype)
+ direction = F.one_hot(directions, num_classes=4).to(dtype)
+ return torch.cat((grid, direction), dim=1)
+
+
+def select_teaching_signal(net, layer, instruction, context, soma, condition):
+ raw = instruction + context
+ innovation = raw - net.predictor(layer, soma)
+ if condition == "clean_kp":
+ used = instruction
+ elif condition == "raw_shared":
+ used = raw
+ elif condition == "sdil":
+ used = innovation
+ else:
+ raise ValueError(f"no shared teaching rule for condition {condition}")
+ return used, raw, innovation
+
+
+@torch.no_grad()
+def manual_step(net, features, missions, actions, condition):
+ """Apply one exact-BP or reciprocal local-learning update."""
+ if condition not in CONDITIONS:
+ raise ValueError(f"unknown BabyAI condition: {condition}")
+ state = net.forward_features(features, missions)
+ h, u, contexts = state["h"], state["u"], state["context"]
+ probabilities = torch.softmax(state["logits"], dim=1)
+ output_instruction = (
+ F.one_hot(actions, num_classes=net.config.action_dim).to(net.dtype)
+ - probabilities)
+ batch = features.shape[0]
+
+ hidden_deltas = [None] * net.config.hidden_layers
+ used_fields = [None] * net.config.hidden_layers
+ raw_fields = [None] * net.config.hidden_layers
+ innovation_fields = [None] * net.config.hidden_layers
+ child_delta = output_instruction
+ for layer in reversed(range(net.config.hidden_layers)):
+ if condition == "bp":
+ used = child_delta @ net.W[layer + 1]
+ raw = used + contexts[layer]
+ innovation = raw - net.predictor(layer, h[layer + 1])
+ else:
+ instruction = child_delta @ net.Q[layer + 1]
+ used, raw, innovation = select_teaching_signal(
+ net, layer, instruction, contexts[layer], h[layer + 1],
+ condition)
+ delta = used * (1.0 - torch.tanh(u[layer]).square())
+ hidden_deltas[layer] = delta
+ used_fields[layer] = used
+ raw_fields[layer] = raw
+ innovation_fields[layer] = innovation
+ child_delta = delta
+
+ directions = [hidden_deltas[0].t() @ h[0] / batch]
+ for layer in range(1, net.config.hidden_layers):
+ directions.append(hidden_deltas[layer].t() @ h[layer] / batch)
+ directions.append(output_instruction.t() @ h[-2] / batch)
+ bias_directions = [delta.mean(0) for delta in hidden_deltas]
+ bias_directions.append(output_instruction.mean(0))
+
+ for layer, (direction, bias_direction) in enumerate(
+ zip(directions, bias_directions)):
+ net.mW[layer].mul_(net.config.momentum).add_(
+ direction - net.config.weight_decay * net.W[layer])
+ net.mb[layer].mul_(net.config.momentum).add_(bias_direction)
+ net.W[layer].add_(net.mW[layer], alpha=net.config.learning_rate)
+ net.b[layer].add_(net.mb[layer], alpha=net.config.learning_rate)
+ if condition != "bp" and layer > 0:
+ net.mQ[layer].mul_(net.config.momentum).add_(
+ direction - net.config.weight_decay * net.Q[layer])
+ net.Q[layer].add_(
+ net.mQ[layer], alpha=net.config.reciprocal_learning_rate)
+
+ loss = F.cross_entropy(state["logits"], actions)
+ return float(loss), {
+ "directions": directions,
+ "bias_directions": bias_directions,
+ "used": used_fields,
+ "raw": raw_fields,
+ "innovation": innovation_fields,
+ "context": contexts,
+ }
+
+
+@torch.no_grad()
+def evaluate_actions(net, images, directions, missions, actions,
+ cardinalities, batch_size=1024, context_enabled=True):
+ device = net.device
+ total_loss = 0.0
+ correct = 0
+ total = len(actions)
+ for start in range(0, total, batch_size):
+ stop = min(start + batch_size, total)
+ features = encode_visual(
+ images[start:stop], directions[start:stop], *cardinalities,
+ device=device, dtype=net.dtype)
+ mission_batch = torch.as_tensor(
+ missions[start:stop], device=device, dtype=net.dtype)
+ action_batch = torch.as_tensor(
+ actions[start:stop], device=device, dtype=torch.long)
+ logits = net.forward_features(
+ features, mission_batch,
+ context_enabled=context_enabled)["logits"]
+ total_loss += float(F.cross_entropy(
+ logits, action_batch, reduction="sum"))
+ correct += int((logits.argmax(1) == action_batch).sum())
+ return {"accuracy": correct / total, "loss": total_loss / total}