| from typing import Dict |
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from einops import rearrange, reduce |
| from diffusers.schedulers.scheduling_ddpm import DDPMScheduler |
| import random |
|
|
| from diffusion_policy.model.common.normalizer import LinearNormalizer |
| from diffusion_policy.policy.base_lowdim_policy import BaseLowdimPolicy |
| from diffusion_policy.model.diffusion.conditional_unet1d import ConditionalUnet1D |
| from diffusion_policy.model.diffusion.mask_generator import LowdimMaskGenerator |
|
|
| class DiffusionUnetLowdimPolicy(BaseLowdimPolicy): |
| def __init__(self, |
| model: ConditionalUnet1D, |
| noise_scheduler: DDPMScheduler, |
| horizon, |
| obs_dim, |
| action_dim, |
| n_action_steps, |
| n_obs_steps, |
| num_inference_steps=None, |
| obs_as_local_cond=False, |
| obs_as_global_cond=False, |
| pred_action_steps_only=False, |
| oa_step_convention=False, |
| |
| **kwargs): |
| super().__init__() |
| assert not (obs_as_local_cond and obs_as_global_cond) |
| if pred_action_steps_only: |
| assert obs_as_global_cond |
| self.model = model |
| self.noise_scheduler = noise_scheduler |
| self.mask_generator = LowdimMaskGenerator( |
| action_dim=action_dim, |
| obs_dim=0 if (obs_as_local_cond or obs_as_global_cond) else obs_dim, |
| max_n_obs_steps=n_obs_steps, |
| fix_obs_steps=True, |
| action_visible=False |
| ) |
| self.normalizer = LinearNormalizer() |
| self.horizon = horizon |
| self.obs_dim = obs_dim |
| self.action_dim = action_dim |
| self.n_action_steps = n_action_steps |
| self.n_obs_steps = n_obs_steps |
| self.obs_as_local_cond = obs_as_local_cond |
| self.obs_as_global_cond = obs_as_global_cond |
| self.pred_action_steps_only = pred_action_steps_only |
| self.oa_step_convention = oa_step_convention |
| self.kwargs = kwargs |
|
|
| if num_inference_steps is None: |
| num_inference_steps = noise_scheduler.config.num_train_timesteps |
| self.num_inference_steps = num_inference_steps |
| |
| |
| def conditional_sample( |
| self, |
| condition_data, condition_mask, |
| condition_data_prev=None, condition_mask_prev=None, |
| local_cond=None, global_cond=None, |
| local_cond_prev=None, global_cond_prev=None, |
| generator=None, |
| prior=None, |
| |
| **kwargs |
| ): |
| model = self.model |
| scheduler = self.noise_scheduler |
|
|
| |
| scheduler.set_timesteps(self.num_inference_steps) |
|
|
| if prior is None: |
| trajectory = torch.randn( |
| size=condition_data.shape, |
| dtype=condition_data.dtype, |
| device=condition_data.device, |
| generator=generator) |
| else: |
| trajectory = prior |
|
|
| trajectory = trajectory.to(device=condition_data.device).contiguous() |
| condition_data = condition_data.contiguous() |
| condition_mask = condition_mask.to(dtype=torch.bool).contiguous() |
| if condition_data_prev is not None: |
| condition_data_prev = condition_data_prev.contiguous() |
| condition_mask_prev = condition_mask_prev.to(dtype=torch.bool).contiguous() |
| weight = self.kwargs.get("alpha", 0.0) |
| kwargs.pop('alpha', None) |
| else: |
| weight = 0.0 |
|
|
| for t in scheduler.timesteps: |
| |
| trajectory[condition_mask] = condition_data[condition_mask] |
| if condition_data_prev is not None: |
| trajectory_prev = trajectory.clone() |
| trajectory_prev[condition_mask_prev] = condition_data_prev[condition_mask_prev] |
|
|
| |
| if condition_data_prev is not None: |
| with torch.no_grad(): |
| model_output_current = model(trajectory, t, local_cond=local_cond, global_cond=global_cond) |
| if ( |
| (trajectory == trajectory_prev).all() |
| and ( |
| (global_cond is None and global_cond_prev is None) or |
| ((global_cond is not None and global_cond_prev is not None) and (global_cond == global_cond_prev).all()) |
| ) |
| and ( |
| (local_cond is None and local_cond_prev is None) or |
| ((local_cond is not None and local_cond_prev is not None) and (local_cond == local_cond_prev).all()) |
| ) |
| ): |
| model_output = model_output_current |
| else: |
| model_output_prev = model(trajectory_prev, t, local_cond=local_cond_prev, global_cond=global_cond_prev) |
| model_output = weight * (model_output_current - model_output_prev) + model_output_current |
| else: |
| with torch.no_grad(): |
| model_output = model(trajectory, t, local_cond=local_cond, global_cond=global_cond) |
|
|
| |
| trajectory = scheduler.step( |
| model_output, t, trajectory, |
| generator=generator, |
| **kwargs |
| ).prev_sample |
|
|
| |
| trajectory[condition_mask] = condition_data[condition_mask] |
|
|
| return trajectory |
|
|
|
|
| def predict_action(self, obs_dict: Dict[str, torch.Tensor], previous_obs_dict: Dict[str, torch.Tensor]=None) -> Dict[str, torch.Tensor]: |
| """ |
| obs_dict: must include "obs" key |
| result: must include "action" key |
| """ |
|
|
| assert 'obs' in obs_dict |
| assert 'past_action' not in obs_dict |
| nobs = self.normalizer['obs'].normalize(obs_dict['obs']) |
| if previous_obs_dict is not None: |
| nobs_prev = self.normalizer['obs'].normalize(previous_obs_dict['obs']) |
|
|
| B, _, Do = nobs.shape |
| To = self.n_obs_steps |
| assert Do == self.obs_dim |
| T = self.horizon |
| Da = self.action_dim |
|
|
| if 'prior' in obs_dict: |
| prior = obs_dict['prior'] |
| else: |
| prior = None |
|
|
| |
| device = self.device |
| dtype = self.dtype |
|
|
| |
| local_cond = None |
| global_cond = None |
| local_cond_prev = None |
| global_cond_prev = None |
| cond_data_prev = None |
| cond_mask_prev = None |
| if self.obs_as_local_cond: |
| |
| |
| local_cond = torch.zeros(size=(B,T,Do), device=device, dtype=dtype) |
| local_cond[:,:To] = nobs[:,:To] |
| shape = (B, T, Da) |
| cond_data = torch.zeros(size=shape, device=device, dtype=dtype) |
| cond_mask = torch.zeros_like(cond_data, dtype=torch.bool) |
| if previous_obs_dict is not None: |
| local_cond_prev = torch.zeros(size=(B, T, Do), device=device, dtype=dtype) |
| local_cond_prev[:, :To] = nobs_prev[:, :To] |
| cond_data_prev = torch.zeros(size=shape, device=device, dtype=dtype) |
| cond_mask_prev = torch.zeros_like(cond_data_prev, dtype=torch.bool) |
|
|
| elif self.obs_as_global_cond: |
| |
| global_cond = nobs[:,:To].reshape(nobs.shape[0], -1) |
| shape = (B, T, Da) |
| if self.pred_action_steps_only: |
| shape = (B, self.n_action_steps, Da) |
| cond_data = torch.zeros(size=shape, device=device, dtype=dtype) |
| cond_mask = torch.zeros_like(cond_data, dtype=torch.bool) |
| if previous_obs_dict is not None: |
| global_cond_prev = nobs_prev[:, :To].reshape(nobs_prev.shape[0], -1) |
| cond_data_prev = torch.zeros(size=shape, device=device, dtype=dtype) |
| cond_mask_prev = torch.zeros_like(cond_data_prev, dtype=torch.bool) |
| else: |
| |
| shape = (B, T, Da+Do) |
| cond_data = torch.zeros(size=shape, device=device, dtype=dtype) |
| cond_mask = torch.zeros_like(cond_data, dtype=torch.bool) |
| cond_data[:,:To,Da:] = nobs[:,:To] |
| cond_mask[:,:To,Da:] = True |
| if previous_obs_dict is not None: |
| cond_data_prev = torch.zeros(size=shape, device=device, dtype=dtype) |
| cond_mask_prev = torch.zeros_like(cond_data_prev, dtype=torch.bool) |
| cond_data_prev[:, :To, Da:] = nobs_prev[:, :To] |
| cond_mask_prev[:, :To, Da:] = True |
|
|
| |
| nsample = self.conditional_sample( |
| cond_data, |
| cond_mask, |
| local_cond=local_cond, |
| global_cond=global_cond, |
| condition_data_prev=cond_data_prev, |
| condition_mask_prev=cond_mask_prev, |
| local_cond_prev=local_cond_prev, |
| global_cond_prev=global_cond_prev, |
| prior=prior, |
| **self.kwargs, |
| ) |
| |
| |
| naction_pred = nsample[...,:Da] |
| action_pred = self.normalizer['action'].unnormalize(naction_pred) |
|
|
| |
| if self.pred_action_steps_only: |
| action = action_pred |
| else: |
| start = To |
| if self.oa_step_convention: |
| start = To - 1 |
| end = start + self.n_action_steps |
| action = action_pred[:,start:end] |
|
|
| result = { |
| 'action': action, |
| 'action_pred': action_pred |
| } |
| if not (self.obs_as_local_cond or self.obs_as_global_cond): |
| nobs_pred = nsample[...,Da:] |
| obs_pred = self.normalizer['obs'].unnormalize(nobs_pred) |
| action_obs_pred = obs_pred[:,start:end] |
| result['action_obs_pred'] = action_obs_pred |
| result['obs_pred'] = obs_pred |
| return result |
|
|
| |
| def set_normalizer(self, normalizer: LinearNormalizer): |
| self.normalizer.load_state_dict(normalizer.state_dict()) |
|
|
| def compute_loss(self, batch): |
| |
| assert 'valid_mask' not in batch |
| nbatch = self.normalizer.normalize(batch) |
| obs = nbatch['obs'] |
| action = nbatch['action'] |
|
|
| |
| local_cond = None |
| global_cond = None |
| trajectory = action |
| if self.obs_as_local_cond: |
| |
| local_cond = obs |
| local_cond[:,self.n_obs_steps:,:] = 0 |
| elif self.obs_as_global_cond: |
| global_cond = obs[:,:self.n_obs_steps,:].reshape( |
| obs.shape[0], -1) |
| if self.pred_action_steps_only: |
| To = self.n_obs_steps |
| start = To |
| if self.oa_step_convention: |
| start = To - 1 |
| end = start + self.n_action_steps |
| trajectory = action[:,start:end] |
| else: |
| trajectory = torch.cat([action, obs], dim=-1) |
|
|
| |
| if self.pred_action_steps_only: |
| condition_mask = torch.zeros_like(trajectory, dtype=torch.bool) |
| else: |
| condition_mask = self.mask_generator(trajectory.shape) |
|
|
| |
| noise = torch.randn(trajectory.shape, device=trajectory.device) |
| bsz = trajectory.shape[0] |
| |
| timesteps = torch.randint( |
| 0, self.noise_scheduler.config.num_train_timesteps, |
| (bsz,), device=trajectory.device |
| ).long() |
| |
| |
| noisy_trajectory = self.noise_scheduler.add_noise( |
| trajectory, noise, timesteps) |
| |
| |
| loss_mask = ~condition_mask |
|
|
| |
| noisy_trajectory[condition_mask] = trajectory[condition_mask] |
| |
| |
| pred = self.model(noisy_trajectory, timesteps, |
| local_cond=local_cond, global_cond=global_cond) |
|
|
| pred_type = self.noise_scheduler.config.prediction_type |
| if pred_type == 'epsilon': |
| target = noise |
| elif pred_type == 'sample': |
| target = trajectory |
| else: |
| raise ValueError(f"Unsupported prediction type {pred_type}") |
|
|
| loss = F.mse_loss(pred, target, reduction='none') |
| loss = loss * loss_mask.type(loss.dtype) |
| loss = reduce(loss, 'b ... -> b (...)', 'mean') |
| loss = loss.mean() |
| return loss |
|
|