# Using MergeDataset and ConcatDataset for Multi‑Source Data Composition in Stable‑WorldModel

> Master multi-source data composition in world models using MergeDataset and ConcatDataset. Join columns horizontally or concatenate episodes vertically for seamless integration with trainers.

- Repository: [GalilAI-group/stable-worldmodel](https://github.com/galilai-group/stable-worldmodel)
- Tags: how-to-guide
- Published: 2026-05-30

---

**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:

```python
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:

```python
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:

```python

# 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:

```python
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

- **`MergeDataset`** performs horizontal, column‑wise joins on datasets with **identical lengths**, automatically deduplicating columns by keeping the first occurrence (`/stable_worldmodel/data/dataset.py#L130-L136`).
- **`ConcatDataset`** performs 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: `MergeDataset` requires equal lengths across sources, while `ConcatDataset` maps indices using binary search on cumulative arrays (`/stable_worldmodel/data/dataset.py#L118-L124`).
- **Test coverage** for `MergeDataset` resides at `/tests/data/test_datasets.py#L777-L894`, and for `ConcatDataset` at `/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.