Spaces:
Sleeping
Sleeping
File size: 8,243 Bytes
377b913 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 | #!/usr/bin/env python
# -*- coding: utf-8 -*-
# File: physics_model.py
#
# Physics-aware network components for the HJ-Communicative RL study.
#
# Contents
# 1. PhysicsEncoder -- small CNN that ingests the extra physics channels.
# 2. MultiHeadCommNet -- the PDF's "Step 8": replaces CommNet's mean/single-head
# pooling with genuine multi-head self-attention across
# agents. (The original repo ALREADY has single-head
# softmax attention via `attention=True`; this is the
# real multi-head upgrade the PDF describes.)
# 3. hj_residual_loss / consistency_loss -- the physics loss terms, written so
# they are internally consistent (both use the geodesic
# field, unlike the PDF which mixes geodesic + Euclidean).
#
# These modules mirror the shapes used in DQNModel.py:
# input : (batch, agents, frame_history, D, W, H) for the image stream
# physics: (batch, agents, C_phys, D, W, H) for the physics stream
# output: (batch, agents, number_actions)
#
# NOTE: this file is torch code and depends on the same torch version as the
# rest of the repo. It has not been executed in this environment (no GPU/torch
# here), but it is structured to be a drop-in sibling of DQNModel.CommNet.
import torch
import torch.nn as nn
# =============================================================================
# 1. Physics encoder branch (PDF Step 2)
# =============================================================================
class PhysicsEncoder(nn.Module):
"""Encodes the (speed / potential / distance / gradient) channels into a
feature vector concatenated with the image-CNN features before communication.
Deliberately lightweight -- the physics maps are smooth and low-frequency."""
def __init__(self, in_channels=3, out_features=128):
super().__init__()
self.net = nn.Sequential(
nn.Conv3d(in_channels, 16, 3, padding=1), nn.PReLU(),
nn.MaxPool3d(2),
nn.Conv3d(16, 32, 3, padding=1), nn.PReLU(),
nn.MaxPool3d(2),
nn.Conv3d(32, 32, 3, padding=0), nn.PReLU(),
nn.AdaptiveAvgPool3d(1),
)
self.fc = nn.Linear(32, out_features)
def forward(self, x): # x: (B*agents, C, D, W, H)
h = self.net(x).flatten(1)
return self.fc(h)
# =============================================================================
# 2. Multi-head attention communication (PDF Step 8, done properly)
# =============================================================================
class MultiHeadCommNet(nn.Module):
"""CommNet with true multi-head self-attention between agents at each stage.
Compared to the repo's CommNet:
* mean pooling -> nn.MultiheadAttention (query/key/value over the agent
axis), with a residual connection and LayerNorm (a transformer block).
* optionally fuses a PhysicsEncoder feature into the shared representation.
"""
def __init__(self, agents, frame_history, number_actions,
phys_channels=0, n_heads=4, xavier=True):
super().__init__()
self.agents = agents
self.frame_history = frame_history
self.use_phys = phys_channels > 0
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# --- shared image CNN (identical to DQNModel.CommNet backbone) --------
def conv(i, o, k, p):
return nn.Conv3d(i, o, k, padding=p)
self.conv0 = conv(frame_history, 32, 5, 1); self.pool0 = nn.MaxPool3d(2); self.act0 = nn.PReLU()
self.conv1 = conv(32, 32, 5, 1); self.pool1 = nn.MaxPool3d(2); self.act1 = nn.PReLU()
self.conv2 = conv(32, 64, 4, 1); self.pool2 = nn.MaxPool3d(2); self.act2 = nn.PReLU()
self.conv3 = conv(64, 64, 3, 0); self.act3 = nn.PReLU()
self.img_dim = 512
# --- physics branch ---------------------------------------------------
if self.use_phys:
self.phys = PhysicsEncoder(phys_channels, out_features=128)
self.fuse = nn.Linear(self.img_dim + 128, self.img_dim)
# --- three transformer communication blocks over the agent axis -------
self.attn = nn.ModuleList([
nn.MultiheadAttention(self.img_dim, n_heads, batch_first=True)
for _ in range(3)])
self.norm = nn.ModuleList([nn.LayerNorm(self.img_dim) for _ in range(3)])
self.ff = nn.ModuleList([
nn.Sequential(nn.Linear(self.img_dim, self.img_dim), nn.PReLU())
for _ in range(3)])
# --- per-agent heads --------------------------------------------------
self.head = nn.ModuleList([
nn.Sequential(nn.Linear(self.img_dim, 256), nn.PReLU(),
nn.Linear(256, number_actions))
for _ in range(agents)])
self.to(self.device)
if xavier:
for m in self.modules():
if isinstance(m, (nn.Conv3d, nn.Linear)):
nn.init.xavier_uniform_(m.weight)
def _cnn(self, x): # x: (B, frame_history, D, W, H)
x = self.act0(self.pool0(self.conv0(x)))
x = self.act1(self.pool1(self.conv1(x)))
x = self.act2(self.pool2(self.conv2(x)))
x = self.act3(self.conv3(x))
return x.reshape(x.size(0), -1)[:, :self.img_dim]
def forward(self, image, physics=None):
"""image: (B, agents, frame_history, D, W, H)
physics: (B, agents, C, D, W, H) or None"""
image = image.to(self.device) / 255.0
B = image.size(0)
feats = []
for i in range(self.agents):
f = self._cnn(image[:, i])
if self.use_phys and physics is not None:
pf = self.phys(physics[:, i].to(self.device))
f = self.fuse(torch.cat([f, pf], dim=-1))
feats.append(f)
x = torch.stack(feats, dim=1) # (B, agents, img_dim)
for attn, norm, ff in zip(self.attn, self.norm, self.ff):
a, _ = attn(x, x, x) # attention across agents
x = norm(x + a) # residual + norm
x = x + ff(x)
out = torch.stack([self.head[i](x[:, i]) for i in range(self.agents)], dim=1)
return out.cpu()
# =============================================================================
# 3. Physics loss terms (PDF Step 7, made internally consistent)
# =============================================================================
def hj_residual_loss(V, F):
"""L_HJ = ( ||grad V|| * F - 1 )^2 on a predicted potential field V.
V, F: (B, D, W, H) tensors. Uses finite differences for the gradient.
Only needed if you actually *learn* V with a PINN head. If V is precomputed
offline (recommended), this term is unnecessary -- V already satisfies the
equation by construction.
"""
gx = V[:, 1:, :, :] - V[:, :-1, :, :]
gy = V[:, :, 1:, :] - V[:, :, :-1, :]
gz = V[:, :, :, 1:] - V[:, :, :, :-1]
# pad back to equal size
gx = torch.nn.functional.pad(gx, (0, 0, 0, 0, 0, 1))
gy = torch.nn.functional.pad(gy, (0, 0, 0, 1, 0, 0))
gz = torch.nn.functional.pad(gz, (0, 1, 0, 0, 0, 0))
grad_norm = torch.sqrt(gx ** 2 + gy ** 2 + gz ** 2 + 1e-8)
return ((grad_norm * F - 1.0) ** 2).mean()
def consistency_loss(V_pred, V_geodesic):
"""L_Consistency = ||V_pred - V_geodesic||^2.
IMPORTANT FIX vs the PDF: the target is the GEODESIC field V_geodesic (the
solution of the Eikonal equation), NOT the Euclidean distance d(x). Using
Euclidean d, as the PDF does, directly contradicts L_HJ whenever the speed
field F varies (which is the entire point of the physics). See EVALUATION.md
for the ~48-voxel discrepancy this causes on real ADNI data.
"""
return ((V_pred - V_geodesic) ** 2).mean()
def communication_loss(agent_features):
"""L_Comm = sum_i ||z_i - mean_j z_j||^2 (encourages a shared code)."""
z_bar = agent_features.mean(dim=1, keepdim=True)
return ((agent_features - z_bar) ** 2).mean()
|