diff --git a/robomimic/algo/diffusion_policy.py b/robomimic/algo/diffusion_policy.py index 0215ad71b..98f234d57 100644 --- a/robomimic/algo/diffusion_policy.py +++ b/robomimic/algo/diffusion_policy.py @@ -1,15 +1,13 @@ """ Implementation of Diffusion Policy https://diffusion-policy.cs.columbia.edu/ by Cheng Chi """ -from typing import Callable, Union -import math from collections import OrderedDict, deque +from copy import deepcopy +from typing import Callable from packaging.version import parse as parse_version -import random import torch import torch.nn as nn import torch.nn.functional as F -# requires diffusers==0.11.1 from diffusers.schedulers.scheduling_ddpm import DDPMScheduler from diffusers.schedulers.scheduling_ddim import DDIMScheduler from diffusers.training_utils import EMAModel @@ -22,10 +20,34 @@ from robomimic.algo import register_algo_factory_func, PolicyAlgo -import random -import robomimic.utils.torch_utils as TorchUtils -import robomimic.utils.tensor_utils as TensorUtils -import robomimic.utils.obs_utils as ObsUtils + +class _DiffusionPolicyEMA: + """Own a modern Diffusers EMA tracker and its inference network.""" + + def __init__(self, model: nn.Module, power: float): + self._power = power + self.averaged_model = deepcopy(model).eval().requires_grad_(False) + self._tracker = self._new_tracker() + + def _new_tracker(self) -> EMAModel: + return EMAModel( + parameters=self.averaged_model.parameters(), + power=self._power, + ) + + def step(self, model: nn.Module) -> None: + """Update averaged weights from the training network.""" + self._tracker.step(model.parameters()) + self._tracker.copy_to(self.averaged_model.parameters()) + + def state_dict(self): + """Return the historical checkpoint representation: model weights.""" + return self.averaged_model.state_dict() + + def load_state_dict(self, state_dict) -> None: + """Restore averaged weights and seed a fresh tracker from them.""" + self.averaged_model.load_state_dict(state_dict) + self._tracker = self._new_tracker() @register_algo_factory_func("diffusion_policy") @@ -110,7 +132,7 @@ def _create_networks(self): # setup EMA ema = None if self.algo_config.ema.enabled: - ema = EMAModel(model=nets, power=self.algo_config.ema.power) + ema = _DiffusionPolicyEMA(nets, power=self.algo_config.ema.power) # set attrs self.nets = nets @@ -375,7 +397,7 @@ def serialize(self): "nets": self.nets.state_dict(), "optimizers": { k : self.optimizers[k].state_dict() for k in self.optimizers }, "lr_schedulers": { k : self.lr_schedulers[k].state_dict() if self.lr_schedulers[k] is not None else None for k in self.lr_schedulers }, - "ema": self.ema.averaged_model.state_dict() if self.ema is not None else None, + "ema": self.ema.state_dict() if self.ema is not None else None, } def deserialize(self, model_dict, load_optimizers=False): @@ -397,7 +419,7 @@ def deserialize(self, model_dict, load_optimizers=False): model_dict["lr_schedulers"] = {} if model_dict.get("ema", None) is not None: - self.ema.averaged_model.load_state_dict(model_dict["ema"]) + self.ema.load_state_dict(model_dict["ema"]) if load_optimizers: for k in model_dict["optimizers"]: diff --git a/setup.py b/setup.py index 704db21b3..840c908dd 100644 --- a/setup.py +++ b/setup.py @@ -29,9 +29,9 @@ "egl_probe>=1.0.1", "torch", "torchvision", - "huggingface_hub==0.23.4", + "huggingface_hub==0.36.2", "transformers==4.41.2", - "diffusers==0.11.1", + "diffusers==0.35.2", ], eager_resources=['*'], include_package_data=True, diff --git a/tests/test_diffusion_policy_ema.py b/tests/test_diffusion_policy_ema.py new file mode 100644 index 000000000..ad6479370 --- /dev/null +++ b/tests/test_diffusion_policy_ema.py @@ -0,0 +1,40 @@ +"""Regression tests for the Diffusion Policy EMA dependency boundary.""" + +import torch +import torch.nn as nn + +from robomimic.algo.diffusion_policy import _DiffusionPolicyEMA + + +def _weight(module: nn.Module) -> torch.Tensor: + return next(module.parameters()).detach().clone() + + +def test_diffusion_policy_ema_updates_and_round_trips_weights(): + source = nn.Linear(2, 1, bias=False) + with torch.no_grad(): + source.weight.fill_(1.0) + + ema = _DiffusionPolicyEMA(source, power=0.75) + assert not ema.averaged_model.training + assert all(not parameter.requires_grad for parameter in ema.averaged_model.parameters()) + torch.testing.assert_close(_weight(ema.averaged_model), _weight(source)) + + with torch.no_grad(): + source.weight.fill_(3.0) + ema.step(source) + assert not torch.equal(_weight(ema.averaged_model), torch.ones((1, 2))) + + saved_weights = ema.state_dict() + restored = _DiffusionPolicyEMA(nn.Linear(2, 1, bias=False), power=0.75) + restored.load_state_dict(saved_weights) + torch.testing.assert_close( + _weight(restored.averaged_model), + _weight(ema.averaged_model), + ) + + before_step = _weight(restored.averaged_model) + with torch.no_grad(): + source.weight.fill_(5.0) + restored.step(source) + assert not torch.equal(_weight(restored.averaged_model), before_step)