Using MergeDataset and ConcatDataset for Multi‑Source Data Composition in Stable‑WorldModel
Use MergeDataset to horizontally join columns from synchronous datasets of equal length, and ConcatDataset to vertically concatenate episodes from heterogeneous sources, both exposing the standard Dataset API required by world‑model trainers.
Stable‑WorldModel stores experience data in episode‑based datasets (e.g., image folders, video files, HDF5 archives). To train a world model on heterogeneous sources, you often need to merge columns from several datasets that share the same timeline, or concatenate whole episodes from different datasets. The library provides two composable wrappers—MergeDataset and ConcatDataset—that implement these patterns while maintaining full compatibility with trainers and planners.
MergeDataset: Horizontal Column‑Wise Join
MergeDataset performs a horizontal join on datasets that share identical lengths and episode boundaries. Use this wrapper when you have multiple sensors—such as pixels and audio—recorded synchronously and need to present them as a single unified observation.
Construction and Automatic Deduplication
The wrapper accepts a list of datasets and automatically resolves columns, discarding duplicates by keeping the first occurrence:
from stable_worldmodel.data import MergeDataset
# Auto-deduplicate columns (first occurrence wins)
merged = MergeDataset([dataset_a, dataset_b])
# Explicitly select columns per source to avoid ambiguity
merged = MergeDataset(
[dataset_a, dataset_b],
keys_from_dataset=[
["pixels", "action"], # from dataset_a
["audio"], # from dataset_b
],
)
According to the source code at /stable_worldmodel/data/dataset.py#L122-L124, passing an empty list raises a ValueError during initialization.
Internal Mechanics and Source Mapping
The implementation enforces strict length parity across children. As shown at /stable_worldmodel/data/dataset.py#L125-L129, the wrapper derives its total length from the first child dataset, assuming all others match exactly.
Column resolution follows a first‑seen strategy. If keys_from_dataset is omitted, the constructor iterates through datasets in order and populates a seen set to discard duplicates, as implemented at /stable_worldmodel/data/dataset.py#L130-L136. The final column_names property concatenates these resolved key lists (/stable_worldmodel/data/dataset.py#L138-L143).
Index access delegates to each child individually. The __getitem__ method (/stable_worldmodel/data/dataset.py#L151-L158) retrieves the indexed item from every child, then merges only the selected keys into a single dictionary. For batch loading, load_chunk (/stable_worldmodel/data/dataset.py#L160-L172) loads slices from each child and merges the resulting dictionaries per time step.
Specialized accessors route requests intelligently. The get_col_data method forwards column requests to the first child that owns the column (/stable_worldmodel/data/dataset.py#L174-L178), while get_row_data aggregates rows across children according to their respective key lists (/stable_worldmodel/data/dataset.py#L180-L187).
ConcatDataset: Vertical Episode‑Wise Concatenation
ConcatDataset performs a vertical concatenation, stacking episodes from multiple datasets end‑to‑end. Use this when combining rollouts collected from different environments, simulation seeds, or experimental runs.
Construction and Cumulative Index Mapping
The wrapper maintains running sums of lengths to translate global indices into local dataset indices:
from stable_worldmodel.data import ConcatDataset
concat = ConcatDataset([dataset_a, dataset_c])
The constructor validates the input list (raising ValueError if empty) at /stable_worldmodel/data/dataset.py#L94-L96. It then builds cumulative offset arrays: _cum stores the running sum of step counts (/stable_worldmodel/data/dataset.py#L98-L100), and _ep_cum stores episode counts (/stable_worldmodel/data/dataset.py#L101-L103).
Unified Column Interface
The column_names property returns the union of all child columns while preserving first‑seen order, as implemented at /stable_worldmodel/data/dataset.py#L104-L113. This ensures that even if children have heterogeneous schemas, the concatenated dataset exposes all available columns.
Efficient Batch Handling
Index translation uses np.searchsorted for O(log n) lookup. The _loc method (/stable_worldmodel/data/dataset.py#L118-L124) maps a global index to a (dataset_idx, local_idx) pair, which __getitem__ then uses to fetch the correct item (/stable_worldmodel/data/dataset.py#L125-L127).
For vectorized access, __getitems__ groups indices by target dataset, preserves the original ordering, and delegates to each child’s own __getitems__ implementation when available (/stable_worldmodel/data/dataset.py#L129-L156). Chunk loading follows a similar pattern: it splits episode indices per child, calls each child’s load_chunk, and re‑assembles results in the original order (/stable_worldmodel/data/dataset.py#L158-L176).
Column‑wise data retrieval concatenates arrays across all children that provide the requested key (/stable_worldmodel/data/dataset.py#L178-L185), while get_row_data aggregates rows from the appropriate children into a single dictionary of arrays (/stable_worldmodel/data/dataset.py#L187-L203).
Nesting MergeDataset and ConcatDataset
Because both wrappers implement the identical Dataset interface—including __len__, __getitem__, load_chunk, get_col_data, and get_row_data—they can be nested arbitrarily. A typical workflow for multi‑sensor, multi‑run composition is:
# 1️⃣ Merge synchronous sensor streams (horizontal join)
merged = MergeDataset(
[image_dataset, audio_dataset],
keys_from_dataset=[["pixels", "observation"], ["audio"]],
)
# 2️⃣ Concatenate several runs (vertical join)
full_dataset = ConcatDataset([merged, another_merged_run])
# Pass to trainer
trainer = WorldModelTrainer(dataset=full_dataset, ...)
The integration test at /tests/data/test_datasets.py#L665-L690 verifies that combined length calculations and item access work correctly across nested wrappers.
Practical Implementation Example
The following runnable example demonstrates both wrappers using the MockDataset helper from the test suite:
import torch
import numpy as np
from stable_worldmodel.data import MergeDataset, ConcatDataset
class MockDataset:
def __init__(self, data, length, num_episodes=1):
self._data = data
self._length = length
self.lengths = np.full(num_episodes, length // num_episodes)
@property
def column_names(self):
return list(self._data.keys())
def __len__(self):
return self._length
def __getitem__(self, idx):
return {k: v[idx] for k, v in self._data.items()}
def load_chunk(self, episodes_idx, start, end):
return [{k: v[s:e] for k, v in self._data.items()}
for s, e in zip(start, end)]
def get_col_data(self, col):
return self._data[col]
def get_row_data(self, row_idx):
if isinstance(row_idx, int):
return {k: v[row_idx] for k, v in self._data.items()}
return {k: v[row_idx] for k, v in self._data.items()}
# Create synchronous streams (20 steps each)
pixels_ds = MockDataset(
data={"pixels": torch.randn(20, 3, 64, 64), "action": torch.randn(20, 2)},
length=20,
)
audio_ds = MockDataset(
data={"audio": torch.randn(20, 16000), "action": torch.randn(20, 2)},
length=20,
)
# Horizontal merge: combine pixels + audio (action auto-deduplicated)
merged = MergeDataset([pixels_ds, audio_ds])
print("Merged columns:", merged.column_names) # ['pixels', 'action', 'audio']
# Vertical concat: add 15 steps from another run
other_ds = MockDataset(
data={"pixels": torch.randn(15, 3, 64, 64), "action": torch.randn(15, 2)},
length=15,
)
full = ConcatDataset([merged, other_ds])
print("Total length:", len(full)) # 35
# Access across boundary (index 22 maps to second dataset)
item = full[22]
print("Keys:", item.keys()) # dict_keys(['pixels', 'action'])
# Chunk loading (batch access)
episodes = np.array([0, 0])
start = np.array([0, 5])
end = np.array([5, 10])
chunk = merged.load_chunk(episodes, start, end)
print("Chunk shape:", chunk[0]["pixels"].shape) # torch.Size([5, 3, 64, 64])
Summary
MergeDatasetperforms horizontal, column‑wise joins on datasets with identical lengths, automatically deduplicating columns by keeping the first occurrence (/stable_worldmodel/data/dataset.py#L130-L136).ConcatDatasetperforms vertical, episode‑wise concatenation across datasets of arbitrary lengths, exposing a unified column namespace via cumulative offset tracking (/stable_worldmodel/data/dataset.py#L98-L103).- Both wrappers delegate storage‑specific operations (chunk loading, column access) to their children, making them agnostic to underlying formats (HDF5, images, video).
- Validation is strict:
MergeDatasetrequires equal lengths across sources, whileConcatDatasetmaps indices using binary search on cumulative arrays (/stable_worldmodel/data/dataset.py#L118-L124). - Test coverage for
MergeDatasetresides at/tests/data/test_datasets.py#L777-L894, and forConcatDatasetat/tests/data/test_datasets.py#L891-L1054.
Frequently Asked Questions
What happens if datasets in MergeDataset have different lengths?
The constructor does not explicitly validate equal lengths at initialization, but length is taken from the first child (/stable_worldmodel/data/dataset.py#L125-L129). Accessing indices beyond the length of other children will raise an IndexError during __getitem__ or load_chunk. The test suite assumes equal lengths for correct operation.
Does ConcatDataset preserve episode boundaries for chunk loading?
Yes. The wrapper tracks cumulative episode counts separately from step counts in _ep_cum (/stable_worldmodel/data/dataset.py#L101-L103). When calling load_chunk, it correctly maps global episode indices to local ones for each child and re‑assembles results in the original order (/stable_worldmodel/data/dataset.py#L158-L176).
Can I nest these wrappers more than two levels deep?
Yes. Because both MergeDataset and ConcatDataset implement the full Dataset interface—including column_names, __getitem__, load_chunk, get_col_data, and get_row_data—they can be composed arbitrarily. The library tests confirm that ConcatDataset can wrap MergeDataset and vice versa (/tests/data/test_datasets.py#L665-L690).
How does column deduplication work in MergeDataset without explicit keys?
Without the keys_from_dataset parameter, the constructor builds a seen set and adds columns in order, skipping any already present (/stable_worldmodel/data/dataset.py#L130-L136). This means the first dataset’s version of a column name takes precedence, and subsequent occurrences are silently ignored.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →