diff --git a/CHANGELOG.md b/CHANGELOG.md index 76b22f8..8e877f9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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 diff --git a/src/decisionrl/imitation.py b/src/decisionrl/imitation.py index c2ef155..96f66b1 100644 --- a/src/decisionrl/imitation.py +++ b/src/decisionrl/imitation.py @@ -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()) @@ -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)) + \ diff --git a/tests/test_imitation.py b/tests/test_imitation.py index f08c05c..463e127 100644 --- a/tests/test_imitation.py +++ b/tests/test_imitation.py @@ -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 @@ -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)