Skip to content

latent_sokoban.dataset

latent_sokoban.dataset

Trajectory generation and the shared dataset format.

Dataset format (one .npz shard + sidecar .json metadata):

frames          uint8  (F, 64, 64, 3)   all frames, episodes concatenated
actions         int8   (F,)             action taken FROM frame i; -1 on
                                        the final frame of each episode
episode_starts  int64  (E,)             index of each episode's first frame
episode_lens    int64  (E,)             number of frames (T+1) per episode
goal_frames     uint8  (E, 64, 64, 3)   goal observation per episode
pushed          bool   (F,)             transition from frame i pushed a box
invalid         bool   (F,)             transition from frame i was a no-op
solved          bool   (E,)             episode ended solved
kind            int8   (E,)             0 random, 1 solver, 2 perturbed
levels          str    (E,)             ascii level definitions

The transition (frames[i], actions[i], frames[i+1]) is a valid training tuple whenever actions[i] != -1.

rollout

rollout(env, actions, theme=None, rng=None)

Execute actions from reset, rendering every frame.

Source code in latent_sokoban/dataset.py
def rollout(
    env: SokobanEnv,
    actions: list[int],
    theme: Theme | None = None,
    rng: np.random.Generator | None = None,
) -> dict:
    """Execute actions from reset, rendering every frame."""
    env.reset()
    frames = [render(env.level, env.state, theme, rng=rng)]
    pushed, invalid, taken = [], [], []
    for a in actions:
        _, done, info = env.step(a)
        frames.append(render(env.level, env.state, theme, rng=rng))
        taken.append(a)
        pushed.append(info.pushed)
        invalid.append(info.invalid)
        if done:
            break
    return {
        "frames": np.stack(frames),
        "actions": np.array(taken, dtype=np.int8),
        "pushed": np.array(pushed, dtype=bool),
        "invalid": np.array(invalid, dtype=bool),
        "solved": env.solved,
    }

perturbed_solution

perturbed_solution(rng, solution, n_perturb=3)

Insert random detour actions into an optimal solution. The episode may or may not still solve, and both outcomes are useful signal.

Source code in latent_sokoban/dataset.py
def perturbed_solution(
    rng: np.random.Generator, solution: list[int], n_perturb: int = 3
) -> list[int]:
    """Insert random detour actions into an optimal solution. The episode
    may or may not still solve, and both outcomes are useful signal."""
    actions = list(solution)
    for _ in range(n_perturb):
        pos = int(rng.integers(0, len(actions) + 1))
        actions.insert(pos, int(rng.integers(0, 4)))
    return actions

generate_shard

generate_shard(rng, n_episodes, size=6, n_boxes=1, max_steps=40, mix=(0.5, 0.3, 0.2), theme=None)

Generate one dataset shard with the recommended composition: 50% random / 30% solver / 20% perturbed-solver trajectories.

Source code in latent_sokoban/dataset.py
def generate_shard(
    rng: np.random.Generator,
    n_episodes: int,
    size: int = 6,
    n_boxes: int = 1,
    max_steps: int = 40,
    mix: tuple[float, float, float] = (0.5, 0.3, 0.2),
    theme: Theme | None = None,
) -> dict:
    """Generate one dataset shard with the recommended composition:
    50% random / 30% solver / 20% perturbed-solver trajectories."""
    from latent_sokoban.levels import generate_level

    episodes = []
    kinds = rng.choice(3, size=n_episodes, p=list(mix))
    for kind in kinds:
        level, solution = generate_level(rng, size=size, n_boxes=n_boxes)
        env = SokobanEnv(level, max_steps=max_steps)
        if kind == KIND_RANDOM:
            actions = random_actions(rng, max_steps)
        elif kind == KIND_SOLVER:
            actions = solution
        else:
            actions = perturbed_solution(rng, solution)
        ep = rollout(env, actions, theme, rng=rng)
        ep["kind"] = int(kind)
        ep["level"] = level.to_ascii()
        ep["goal_frame"] = render_goal(level, theme, rng=rng)
        episodes.append(ep)
    return pack_episodes(episodes)