Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,14 @@ to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
the freeze semantics of Stable-Baselines3's `VecNormalize`.

### Fixed
- `BC` and `GAIL` raised on any machine with a GPU. Both take `device="auto"`, which puts
the network on CUDA when there is one, while `TransitionDataset` defaults to `"cpu"` and
`collect_expert_dataset` never passes anything else — so the documented way of using them
failed on the first batch with "Expected all tensors to be on the same device". CQL, IQL,
TD3BC and DiffusionPolicy already moved each batch with `ReplayBatch.to(self.device)`,
whose docstring describes this exact case; `decisionrl.imitation` was the one place that
did not. Three tests in `test_imitation.py` failed on every GPU machine and passed on
every CI runner, because the runners have no GPU.
- The on-policy rollout buffer shuffled its minibatches with the global NumPy RNG. Two
on-policy agents seeded differently in the same process therefore drew from one shared
shuffle stream and perturbed each other, and the buffer never held the per-instance RNG
Expand Down
15 changes: 12 additions & 3 deletions src/decisionrl/imitation.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,12 @@ def train(self, dataset: TransitionDataset, n_iters: int = 2000, batch_size: int
log_interval: int = 0) -> dict:
losses: deque = deque(maxlen=100)
for it in range(n_iters):
batch = dataset.sample(batch_size)
# .to(self.device), as CQL, IQL, TD3BC and DiffusionPolicy all do.
# `device="auto"` puts the actor on a GPU when there is one, while
# TransitionDataset defaults to "cpu" and collect_expert_dataset
# never passes anything else -- so without this the documented way
# of using BC raises on the first batch of any machine with a GPU.
batch = dataset.sample(batch_size).to(self.device)
dist = self.actor(batch.obs)
if self.discrete:
loss = F.cross_entropy(dist.logits, batch.actions.long())
Expand Down Expand Up @@ -208,8 +213,12 @@ def _update_discriminator(self, pol_obs, pol_act, epochs, batch_size):
np.zeros(len(pol_obs)), device=str(self.device))
losses = []
for _ in range(epochs):
e = self.expert.sample(batch_size)
p = pol.sample(batch_size)
# The expert dataset comes from the caller and is on whatever
# device they built it on, usually the CPU; `pol` is built on
# self.device just above. Both are moved, so the discriminator
# never has to care which of its two inputs came from where.
e = self.expert.sample(batch_size).to(self.device)
p = pol.sample(batch_size).to(self.device)
e_logits = self.discriminator.logits(e.obs, e.actions)
p_logits = self.discriminator.logits(p.obs, p.actions)
loss = F.binary_cross_entropy_with_logits(e_logits, torch.ones_like(e_logits)) + \
Expand Down
68 changes: 68 additions & 0 deletions tests/test_imitation.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import numpy as np
import pytest
import torch

from decisionrl.envs import CartPole
from decisionrl.imitation import BC, GAIL, DAgger, GAILDiscriminator, collect_expert_dataset
Expand Down Expand Up @@ -58,3 +59,70 @@ def test_gail_imitates_expert(quiet_logger):
after, _ = evaluate_policy(gail, CartPole(), n_episodes=10, seed=100)
# GAIL matches the expert from demonstrations alone (no env reward); random ~= 22.
assert after > 200.0


def _spy_on_batch_moves(dataset, recorder):
"""Wrap ``dataset.sample`` so every ``batch.to(device)`` is recorded."""
original_sample = dataset.sample

def sample(batch_size):
batch = original_sample(batch_size)
original_to = batch.to

def to(device):
recorder.append(torch.device(device))
return original_to(device)

batch.to = to
return batch

dataset.sample = sample
return dataset


def test_bc_moves_every_batch_to_the_actors_device(quiet_logger):
"""BC.train must hand the actor tensors that are where the actor is.

`device="auto"` puts the actor on a GPU when the machine has one, while
TransitionDataset defaults to "cpu" and collect_expert_dataset never passes
anything else. Without the move, the documented way of using BC raises
"Expected all tensors to be on the same device" on its very first batch --
on a GPU machine, which CI is not, which is why this went unnoticed.

Checked here as the step rather than as the crash, so it runs everywhere:
on CI the move is a no-op, but its absence is still a failure.
"""
moved: list = []
data = _spy_on_batch_moves(collect_expert_dataset(CartPole(), _expert, 200, seed=0), moved)
bc = BC(CartPole(), seed=0, logger=quiet_logger)
bc.train(data, n_iters=3, batch_size=16)
assert moved == [bc.device] * 3


def test_gail_moves_the_expert_batch_to_the_discriminators_device(quiet_logger):
"""The expert dataset is the caller's, and is usually on the CPU.

GAIL builds its policy dataset on `self.device` already, so the expert side
is the one that can arrive from somewhere else — and does, every time
collect_expert_dataset is used the way the README uses it.
"""
moved: list = []
expert = _spy_on_batch_moves(collect_expert_dataset(CartPole(), _expert, 200, seed=0), moved)
gail = GAIL(CartPole(), expert, n_steps=64, batch_size=16, n_epochs=1, seed=0,
logger=quiet_logger)
gail.learn(iterations=1, steps_per_iter=64, disc_epochs=2, disc_batch=16)
assert moved == [gail.device] * 2


@pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a GPU; CI runners have none")
def test_bc_trains_when_the_dataset_is_on_another_device(quiet_logger):
"""The crash itself, for anyone running the suite on a machine with a GPU.

CI cannot run this, and saying so is the point: three tests in this file
failed on every GPU machine and passed on every runner, for as long as that
difference went unstated.
"""
data = collect_expert_dataset(CartPole(), _expert, 200, seed=0) # cpu
bc = BC(CartPole(), seed=0, logger=quiet_logger) # cuda
assert data.device != bc.device
bc.train(data, n_iters=3, batch_size=16)
Loading