Skip to content
Open
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
44 changes: 33 additions & 11 deletions robomimic/algo/diffusion_policy.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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")
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand All @@ -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"]:
Expand Down
4 changes: 2 additions & 2 deletions setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
40 changes: 40 additions & 0 deletions tests/test_diffusion_policy_ema.py
Original file line number Diff line number Diff line change
@@ -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)