Skip to content

Commit 19b9e47

Browse files
committed
Add multimodal observation schema and DQN encoders
1 parent 81902af commit 19b9e47

6 files changed

Lines changed: 327 additions & 0 deletions

File tree

README.md

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -70,3 +70,30 @@ uv run black --check adeptly tests
7070
uv run mypy adeptly
7171
uv run pytest
7272
```
73+
74+
75+
### Multimodal DQN inference (batched)
76+
```python
77+
import torch
78+
from adeptly.observations import MultimodalQNetwork, ObservationBatch
79+
80+
batch_size = 8
81+
num_actions = 4
82+
model = MultimodalQNetwork(
83+
image_channels=3,
84+
telemetry_dim=6,
85+
action_size=num_actions,
86+
sequence_vocab_size=256,
87+
)
88+
89+
# Minimal batched environment loop for inference-only usage.
90+
for _ in range(5):
91+
observations = ObservationBatch(
92+
image_frames=torch.randint(0, 256, (batch_size, 3, 84, 84), dtype=torch.uint8),
93+
scalar_telemetry=torch.randn(batch_size, 6),
94+
events_or_text=torch.randint(0, 256, (batch_size, 12), dtype=torch.int64),
95+
)
96+
q_values = model(observations)
97+
actions = torch.argmax(q_values, dim=1)
98+
# send `actions` back to your vectorized environment
99+
```

adeptly/observations/__init__.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,22 @@
1+
"""Observation schemas and multimodal DQN building blocks."""
2+
3+
from adeptly.observations.dqn_multimodal import MultimodalQNetwork
4+
from adeptly.observations.encoders import EventSequenceEncoder, ModalityFusion, TelemetryEncoder, VisionEncoder
5+
from adeptly.observations.schema import (
6+
ObservationBatch,
7+
normalize_image_frames,
8+
normalize_scalar_telemetry,
9+
validate_observation_batch,
10+
)
11+
12+
__all__ = [
13+
"EventSequenceEncoder",
14+
"ModalityFusion",
15+
"MultimodalQNetwork",
16+
"ObservationBatch",
17+
"TelemetryEncoder",
18+
"VisionEncoder",
19+
"normalize_image_frames",
20+
"normalize_scalar_telemetry",
21+
"validate_observation_batch",
22+
]
Lines changed: 71 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,71 @@
1+
"""Multimodal DQN network composed from observation encoders and fusion."""
2+
3+
from __future__ import annotations
4+
5+
from torch import Tensor, nn
6+
7+
from adeptly.observations.encoders import EventSequenceEncoder, ModalityFusion, TelemetryEncoder, VisionEncoder
8+
from adeptly.observations.schema import (
9+
ObservationBatch,
10+
normalize_image_frames,
11+
normalize_scalar_telemetry,
12+
validate_observation_batch,
13+
)
14+
15+
16+
class MultimodalQNetwork(nn.Module):
17+
"""DQN action-value network for multimodal observations."""
18+
19+
def __init__(
20+
self,
21+
image_channels: int,
22+
telemetry_dim: int,
23+
action_size: int,
24+
*,
25+
vision_embedding_dim: int = 128,
26+
telemetry_embedding_dim: int = 64,
27+
sequence_vocab_size: int | None = None,
28+
sequence_embedding_dim: int = 32,
29+
sequence_output_dim: int = 64,
30+
fused_dim: int = 256,
31+
use_attention_fusion: bool = False,
32+
) -> None:
33+
super().__init__()
34+
35+
if action_size <= 0:
36+
raise ValueError("action_size must be > 0")
37+
38+
self.vision_encoder = VisionEncoder(image_channels, vision_embedding_dim)
39+
self.telemetry_encoder = TelemetryEncoder(telemetry_dim, telemetry_embedding_dim)
40+
41+
self.sequence_encoder: EventSequenceEncoder | None = None
42+
encoder_dims = [vision_embedding_dim, telemetry_embedding_dim]
43+
if sequence_vocab_size is not None:
44+
self.sequence_encoder = EventSequenceEncoder(
45+
vocab_size=sequence_vocab_size,
46+
embedding_dim=sequence_embedding_dim,
47+
output_dim=sequence_output_dim,
48+
)
49+
encoder_dims.append(sequence_output_dim)
50+
51+
self.fusion = ModalityFusion(encoder_dims, fused_dim=fused_dim, use_attention=use_attention_fusion)
52+
self.action_head = nn.Linear(fused_dim, action_size)
53+
54+
def forward(self, observations: ObservationBatch) -> Tensor:
55+
validate_observation_batch(observations)
56+
57+
image_frames = normalize_image_frames(observations.image_frames)
58+
scalar_telemetry = normalize_scalar_telemetry(observations.scalar_telemetry)
59+
60+
embeddings = [
61+
self.vision_encoder(image_frames),
62+
self.telemetry_encoder(scalar_telemetry),
63+
]
64+
65+
if observations.events_or_text is not None:
66+
if self.sequence_encoder is None:
67+
raise ValueError("events_or_text provided but sequence encoder was not configured")
68+
embeddings.append(self.sequence_encoder(observations.events_or_text))
69+
70+
fused = self.fusion(embeddings)
71+
return self.action_head(fused)

adeptly/observations/encoders.py

Lines changed: 88 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,88 @@
1+
"""Composable multimodal encoders for DQN observations."""
2+
3+
from __future__ import annotations
4+
5+
import torch
6+
from torch import Tensor, nn
7+
8+
9+
class VisionEncoder(nn.Module):
10+
"""Small CNN encoder for image frame observations."""
11+
12+
def __init__(self, in_channels: int, embedding_dim: int) -> None:
13+
super().__init__()
14+
self.backbone = nn.Sequential(
15+
nn.Conv2d(in_channels, 32, kernel_size=8, stride=4),
16+
nn.ReLU(),
17+
nn.Conv2d(32, 64, kernel_size=4, stride=2),
18+
nn.ReLU(),
19+
nn.Conv2d(64, 64, kernel_size=3, stride=1),
20+
nn.ReLU(),
21+
nn.AdaptiveAvgPool2d((1, 1)),
22+
nn.Flatten(),
23+
nn.Linear(64, embedding_dim),
24+
nn.ReLU(),
25+
)
26+
27+
def forward(self, image_frames: Tensor) -> Tensor:
28+
return self.backbone(image_frames)
29+
30+
31+
class TelemetryEncoder(nn.Module):
32+
"""MLP encoder for scalar telemetry channels."""
33+
34+
def __init__(self, input_dim: int, embedding_dim: int, hidden_dim: int = 128) -> None:
35+
super().__init__()
36+
self.model = nn.Sequential(
37+
nn.Linear(input_dim, hidden_dim),
38+
nn.ReLU(),
39+
nn.Linear(hidden_dim, embedding_dim),
40+
nn.ReLU(),
41+
)
42+
43+
def forward(self, scalar_telemetry: Tensor) -> Tensor:
44+
return self.model(scalar_telemetry)
45+
46+
47+
class EventSequenceEncoder(nn.Module):
48+
"""Embedding + GRU encoder for event/text token sequences."""
49+
50+
def __init__(self, vocab_size: int, embedding_dim: int, output_dim: int) -> None:
51+
super().__init__()
52+
self.embedding = nn.Embedding(vocab_size, embedding_dim)
53+
self.gru = nn.GRU(embedding_dim, output_dim, batch_first=True)
54+
55+
def forward(self, events_or_text: Tensor) -> Tensor:
56+
embedded = self.embedding(events_or_text.long())
57+
_, hidden = self.gru(embedded)
58+
return hidden.squeeze(0)
59+
60+
61+
class ModalityFusion(nn.Module):
62+
"""Fuse modality embeddings via concatenation projection or attention pooling."""
63+
64+
def __init__(self, input_dims: list[int], fused_dim: int, use_attention: bool = False) -> None:
65+
super().__init__()
66+
self.use_attention = use_attention
67+
self.fused_dim = fused_dim
68+
69+
if use_attention:
70+
if len(set(input_dims)) != 1:
71+
raise ValueError("All modality dims must match when use_attention=True")
72+
self.attention = nn.MultiheadAttention(embed_dim=input_dims[0], num_heads=1, batch_first=True)
73+
self.output_projection = nn.Linear(input_dims[0], fused_dim)
74+
else:
75+
self.output_projection = nn.Linear(sum(input_dims), fused_dim)
76+
77+
def forward(self, modality_embeddings: list[Tensor]) -> Tensor:
78+
if not modality_embeddings:
79+
raise ValueError("modality_embeddings cannot be empty")
80+
81+
if self.use_attention:
82+
stacked = torch.stack(modality_embeddings, dim=1)
83+
attended, _ = self.attention(stacked, stacked, stacked)
84+
pooled = attended.mean(dim=1)
85+
return self.output_projection(pooled)
86+
87+
concatenated = torch.cat(modality_embeddings, dim=1)
88+
return self.output_projection(concatenated)

adeptly/observations/schema.py

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,57 @@
1+
"""Observation schema and validation utilities for multimodal DQN inputs."""
2+
3+
from __future__ import annotations
4+
5+
from dataclasses import dataclass
6+
7+
import torch
8+
from torch import Tensor
9+
10+
11+
@dataclass(slots=True)
12+
class ObservationBatch:
13+
"""Batched multimodal observations consumed by multimodal DQN networks."""
14+
15+
image_frames: Tensor
16+
scalar_telemetry: Tensor
17+
events_or_text: Tensor | None = None
18+
19+
20+
def validate_observation_batch(observations: ObservationBatch) -> None:
21+
"""Validate batch shapes early to prevent downstream runtime shape errors."""
22+
23+
image = observations.image_frames
24+
telemetry = observations.scalar_telemetry
25+
sequence = observations.events_or_text
26+
27+
if image.ndim != 4:
28+
raise ValueError("image_frames must have shape [batch, channels, height, width]")
29+
if telemetry.ndim != 2:
30+
raise ValueError("scalar_telemetry must have shape [batch, features]")
31+
32+
batch_size = image.shape[0]
33+
if telemetry.shape[0] != batch_size:
34+
raise ValueError("image_frames and scalar_telemetry must share the same batch dimension")
35+
36+
if sequence is not None:
37+
if sequence.ndim != 2:
38+
raise ValueError("events_or_text must have shape [batch, sequence_length]")
39+
if sequence.shape[0] != batch_size:
40+
raise ValueError("events_or_text must share the same batch dimension as image_frames")
41+
42+
43+
def normalize_image_frames(image_frames: Tensor) -> Tensor:
44+
"""Convert image frames to float tensors in [0, 1]."""
45+
46+
if not torch.is_floating_point(image_frames):
47+
image_frames = image_frames.to(torch.float32)
48+
return image_frames / 255.0 if image_frames.max().item() > 1.0 else image_frames
49+
50+
51+
def normalize_scalar_telemetry(scalar_telemetry: Tensor, eps: float = 1e-6) -> Tensor:
52+
"""Per-batch z-score normalization for scalar telemetry channels."""
53+
54+
telemetry = scalar_telemetry.to(torch.float32)
55+
mean = telemetry.mean(dim=0, keepdim=True)
56+
std = telemetry.std(dim=0, keepdim=True).clamp_min(eps)
57+
return (telemetry - mean) / std
Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,62 @@
1+
import pytest
2+
import torch
3+
4+
from adeptly.observations import (
5+
MultimodalQNetwork,
6+
ObservationBatch,
7+
normalize_image_frames,
8+
normalize_scalar_telemetry,
9+
validate_observation_batch,
10+
)
11+
12+
13+
def test_validate_observation_batch_rejects_mismatched_batch_dims():
14+
observations = ObservationBatch(
15+
image_frames=torch.zeros(2, 3, 84, 84),
16+
scalar_telemetry=torch.zeros(3, 4),
17+
)
18+
19+
with pytest.raises(ValueError, match="share the same batch dimension"):
20+
validate_observation_batch(observations)
21+
22+
23+
def test_normalization_helpers_return_expected_ranges_and_stats():
24+
image = torch.randint(0, 256, (4, 3, 16, 16), dtype=torch.uint8)
25+
normalized_image = normalize_image_frames(image)
26+
assert normalized_image.dtype == torch.float32
27+
assert float(normalized_image.max()) <= 1.0
28+
29+
telemetry = torch.tensor([[1.0, 2.0], [3.0, 6.0], [5.0, 10.0]])
30+
normalized_telemetry = normalize_scalar_telemetry(telemetry)
31+
assert torch.allclose(normalized_telemetry.mean(dim=0), torch.zeros(2), atol=1e-5)
32+
33+
34+
def test_multimodal_q_network_outputs_action_values_for_batched_input():
35+
model = MultimodalQNetwork(
36+
image_channels=3,
37+
telemetry_dim=5,
38+
action_size=4,
39+
sequence_vocab_size=128,
40+
use_attention_fusion=False,
41+
)
42+
43+
observations = ObservationBatch(
44+
image_frames=torch.randint(0, 256, (6, 3, 84, 84), dtype=torch.uint8),
45+
scalar_telemetry=torch.randn(6, 5),
46+
events_or_text=torch.randint(0, 128, (6, 10), dtype=torch.int64),
47+
)
48+
49+
output = model(observations)
50+
assert output.shape == (6, 4)
51+
52+
53+
def test_multimodal_q_network_raises_if_sequence_not_configured():
54+
model = MultimodalQNetwork(image_channels=3, telemetry_dim=3, action_size=2)
55+
observations = ObservationBatch(
56+
image_frames=torch.randint(0, 256, (2, 3, 84, 84), dtype=torch.uint8),
57+
scalar_telemetry=torch.randn(2, 3),
58+
events_or_text=torch.randint(0, 100, (2, 8), dtype=torch.int64),
59+
)
60+
61+
with pytest.raises(ValueError, match="sequence encoder"):
62+
model(observations)

0 commit comments

Comments
 (0)