Add 4DHuman

This commit is contained in:
Hacker 17082006
2024-03-31 15:46:19 +07:00
parent e4f50bbc63
commit cd340c210e
249 changed files with 29981 additions and 12 deletions
+56
View File
@@ -0,0 +1,56 @@
import warnings
import custom_mmpkg.custom_mmcv as mmcv
from packaging.version import parse
from .version import __version__
def digit_version(version_str: str, length: int = 4):
"""Convert a version string into a tuple of integers.
This method is usually used for comparing two versions. For pre-release
versions: alpha < beta < rc.
Args:
version_str (str): The version string.
length (int): The maximum number of version levels. Default: 4.
Returns:
tuple[int]: The version info in digits (integers).
"""
version = parse(version_str)
assert version.release, f'failed to parse version {version_str}'
release = list(version.release)
release = release[:length]
if len(release) < length:
release = release + [0] * (length - len(release))
if version.is_prerelease:
mapping = {'a': -3, 'b': -2, 'rc': -1}
val = -4
# version.pre can be None
if version.pre:
if version.pre[0] not in mapping:
warnings.warn(f'unknown prerelease version {version.pre[0]}, '
'version checking may go wrong')
else:
val = mapping[version.pre[0]]
release.extend([val, version.pre[-1]])
else:
release.extend([val, 0])
elif version.is_postrelease:
release.extend([1, version.post])
else:
release.extend([0, 0])
return tuple(release)
mmcv_minimum_version = '1.3.17' #Dirty hacking for custom_mmcv
mmcv_maximum_version = '1.9.0'
mmcv_version = digit_version(mmcv.__version__)
assert (mmcv_version >= digit_version(mmcv_minimum_version)
and mmcv_version <= digit_version(mmcv_maximum_version)), \
f'MMCV=={mmcv.__version__} is used but incompatible. ' \
f'Please install mmcv>={mmcv_minimum_version}, <={mmcv_maximum_version}.'
__all__ = ['__version__', 'digit_version']
+13
View File
@@ -0,0 +1,13 @@
from motiondiff_modules.mogen.apis import test, train
from motiondiff_modules.mogen.apis.test import (
collect_results_cpu,
collect_results_gpu,
multi_gpu_test,
single_gpu_test,
)
from motiondiff_modules.mogen.apis.train import set_random_seed, train_model
__all__ = [
'collect_results_cpu', 'collect_results_gpu', 'multi_gpu_test',
'single_gpu_test', 'set_random_seed', 'train_model'
]
+160
View File
@@ -0,0 +1,160 @@
import os.path as osp
import pickle
import shutil
import tempfile
import time
import custom_mmpkg.custom_mmcv as mmcv
import torch
import torch.distributed as dist
from custom_mmpkg.custom_mmcv.runner import get_dist_info
def single_gpu_test(model, data_loader):
"""Test with single gpu."""
model.eval()
results = []
dataset = data_loader.dataset
prog_bar = mmcv.ProgressBar(len(dataset))
for i, data in enumerate(data_loader):
with torch.no_grad():
result = model(return_loss=False, **data)
batch_size = len(result)
if isinstance(result, list):
results.extend(result)
else:
results.append(result)
batch_size = data['motion'].size(0)
for _ in range(batch_size):
prog_bar.update()
return results
def multi_gpu_test(model, data_loader, tmpdir=None, gpu_collect=False):
"""Test model with multiple gpus.
This method tests model with multiple gpus and collects the results
under two different modes: gpu and cpu modes. By setting 'gpu_collect=True'
it encodes results to gpu tensors and use gpu communication for results
collection. On cpu mode it saves the results on different gpus to 'tmpdir'
and collects them by the rank 0 worker.
Args:
model (nn.Module): Model to be tested.
data_loader (nn.Dataloader): Pytorch data loader.
tmpdir (str): Path of directory to save the temporary results from
different gpus under cpu mode.
gpu_collect (bool): Option to use either gpu or cpu to collect results.
Returns:
list: The prediction results.
"""
model.eval()
results = []
dataset = data_loader.dataset
rank, world_size = get_dist_info()
if rank == 0:
# Check if tmpdir is valid for cpu_collect
if (not gpu_collect) and (tmpdir is not None and osp.exists(tmpdir)):
raise OSError((f'The tmpdir {tmpdir} already exists.',
' Since tmpdir will be deleted after testing,',
' please make sure you specify an empty one.'))
prog_bar = mmcv.ProgressBar(len(dataset))
time.sleep(2) # This line can prevent deadlock problem in some cases.
for i, data in enumerate(data_loader):
with torch.no_grad():
result = model(return_loss=False, **data)
if isinstance(result, list):
results.extend(result)
else:
results.append(result)
if rank == 0:
batch_size = data['motion'].size(0)
for _ in range(batch_size * world_size):
prog_bar.update()
# collect results from all ranks
if gpu_collect:
results = collect_results_gpu(results, len(dataset))
else:
results = collect_results_cpu(results, len(dataset), tmpdir)
return results
def collect_results_cpu(result_part, size, tmpdir=None):
"""Collect results in cpu."""
rank, world_size = get_dist_info()
# create a tmp dir if it is not specified
if tmpdir is None:
MAX_LEN = 512
# 32 is whitespace
dir_tensor = torch.full((MAX_LEN, ),
32,
dtype=torch.uint8,
device='cuda')
if rank == 0:
mmcv.mkdir_or_exist('.dist_test')
tmpdir = tempfile.mkdtemp(dir='.dist_test')
tmpdir = torch.tensor(
bytearray(tmpdir.encode()), dtype=torch.uint8, device='cuda')
dir_tensor[:len(tmpdir)] = tmpdir
dist.broadcast(dir_tensor, 0)
tmpdir = dir_tensor.cpu().numpy().tobytes().decode().rstrip()
else:
mmcv.mkdir_or_exist(tmpdir)
# dump the part result to the dir
mmcv.dump(result_part, osp.join(tmpdir, f'part_{rank}.pkl'))
dist.barrier()
# collect all parts
if rank != 0:
return None
else:
# load results of all parts from tmp dir
part_list = []
for i in range(world_size):
part_file = osp.join(tmpdir, f'part_{i}.pkl')
part_result = mmcv.load(part_file)
part_list.append(part_result)
# sort the results
ordered_results = []
for res in zip(*part_list):
ordered_results.extend(list(res))
# the dataloader may pad some samples
ordered_results = ordered_results[:size]
# remove tmp dir
shutil.rmtree(tmpdir)
return ordered_results
def collect_results_gpu(result_part, size):
"""Collect results in gpu."""
rank, world_size = get_dist_info()
# dump result part to tensor with pickle
part_tensor = torch.tensor(
bytearray(pickle.dumps(result_part)), dtype=torch.uint8, device='cuda')
# gather all result part tensor shape
shape_tensor = torch.tensor(part_tensor.shape, device='cuda')
shape_list = [shape_tensor.clone() for _ in range(world_size)]
dist.all_gather(shape_list, shape_tensor)
# padding result part tensor to max length
shape_max = torch.tensor(shape_list).max()
part_send = torch.zeros(shape_max, dtype=torch.uint8, device='cuda')
part_send[:shape_tensor[0]] = part_tensor
part_recv_list = [
part_tensor.new_zeros(shape_max) for _ in range(world_size)
]
# gather all result part
dist.all_gather(part_recv_list, part_send)
if rank == 0:
part_list = []
for recv, shape in zip(part_recv_list, shape_list):
part_result = pickle.loads(recv[:shape[0]].cpu().numpy().tobytes())
part_list.append(part_result)
# sort the results
ordered_results = []
for res in zip(*part_list):
ordered_results.extend(list(res))
# the dataloader may pad some samples
ordered_results = ordered_results[:size]
return ordered_results
+165
View File
@@ -0,0 +1,165 @@
import random
import warnings
import numpy as np
import torch
from custom_mmpkg.custom_mmcv.parallel import MMDataParallel, MMDistributedDataParallel
from custom_mmpkg.custom_mmcv.runner import (
DistSamplerSeedHook,
Fp16OptimizerHook,
OptimizerHook,
build_runner,
)
from motiondiff_modules.mogen.core.distributed_wrapper import DistributedDataParallelWrapper
from motiondiff_modules.mogen.core.evaluation import DistEvalHook, EvalHook
from motiondiff_modules.mogen.core.optimizer import build_optimizers
from motiondiff_modules.mogen.datasets import build_dataloader, build_dataset
from motiondiff_modules.mogen.utils import get_root_logger
def set_random_seed(seed, deterministic=False):
"""Set random seed.
Args:
seed (int): Seed to be used.
deterministic (bool): Whether to set the deterministic option for
CUDNN backend, i.e., set `torch.backends.cudnn.deterministic`
to True and `torch.backends.cudnn.benchmark` to False.
Default: False.
"""
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
if deterministic:
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
def train_model(model,
dataset,
cfg,
distributed=False,
validate=False,
timestamp=None,
device='cuda',
meta=None):
"""Main api for training model."""
logger = get_root_logger(cfg.log_level)
# prepare data loaders
dataset = dataset if isinstance(dataset, (list, tuple)) else [dataset]
data_loaders = [
build_dataloader(
ds,
cfg.data.samples_per_gpu,
cfg.data.workers_per_gpu,
# cfg.gpus will be ignored if distributed
num_gpus=len(cfg.gpu_ids),
dist=distributed,
round_up=True,
seed=cfg.seed) for ds in dataset
]
# determine whether use adversarial training precess or not
use_adverserial_train = cfg.get('use_adversarial_train', False)
# put model on gpus
if distributed:
find_unused_parameters = cfg.get('find_unused_parameters', True)
# Sets the `find_unused_parameters` parameter in
# torch.nn.parallel.DistributedDataParallel
if use_adverserial_train:
# Use DistributedDataParallelWrapper for adversarial training
model = DistributedDataParallelWrapper(
model,
device_ids=[torch.cuda.current_device()],
broadcast_buffers=False,
find_unused_parameters=find_unused_parameters)
else:
model = MMDistributedDataParallel(
model.cuda(),
device_ids=[torch.cuda.current_device()],
broadcast_buffers=False,
find_unused_parameters=find_unused_parameters)
else:
if device == 'cuda':
model = MMDataParallel(
model.cuda(cfg.gpu_ids[0]), device_ids=cfg.gpu_ids)
elif device == 'cpu':
model = model.cpu()
else:
raise ValueError(F'unsupported device name {device}.')
# build runner
optimizer = build_optimizers(model, cfg.optimizer)
if cfg.get('runner') is None:
cfg.runner = {
'type': 'EpochBasedRunner',
'max_epochs': cfg.total_epochs
}
warnings.warn(
'config is now expected to have a `runner` section, '
'please set `runner` in your config.', UserWarning)
runner = build_runner(
cfg.runner,
default_args=dict(
model=model,
batch_processor=None,
optimizer=optimizer,
work_dir=cfg.work_dir,
logger=logger,
meta=meta))
# an ugly walkaround to make the .log and .log.json filenames the same
runner.timestamp = timestamp
if use_adverserial_train:
# The optimizer step process is included in the train_step function
# of the model, so the runner should NOT include optimizer hook.
optimizer_config = None
else:
# fp16 setting
fp16_cfg = cfg.get('fp16', None)
if fp16_cfg is not None:
optimizer_config = Fp16OptimizerHook(
**cfg.optimizer_config, **fp16_cfg, distributed=distributed)
elif distributed and 'type' not in cfg.optimizer_config:
optimizer_config = OptimizerHook(**cfg.optimizer_config)
else:
optimizer_config = cfg.optimizer_config
# register hooks
runner.register_training_hooks(
cfg.lr_config,
optimizer_config,
cfg.checkpoint_config,
cfg.log_config,
cfg.get('momentum_config', None),
custom_hooks_config=cfg.get('custom_hooks', None))
if distributed:
runner.register_hook(DistSamplerSeedHook())
# register eval hooks
if validate:
val_dataset = build_dataset(cfg.data.val, dict(test_mode=True))
val_dataloader = build_dataloader(
val_dataset,
samples_per_gpu=cfg.data.samples_per_gpu,
workers_per_gpu=cfg.data.workers_per_gpu,
dist=distributed,
shuffle=False,
round_up=True)
eval_cfg = cfg.get('evaluation', {})
eval_cfg['by_epoch'] = cfg.runner['type'] != 'IterBasedRunner'
eval_hook = DistEvalHook if distributed else EvalHook
runner.register_hook(eval_hook(val_dataloader, **eval_cfg))
if cfg.resume_from:
runner.resume(cfg.resume_from)
elif cfg.load_from:
runner.load_checkpoint(cfg.load_from)
runner.run(data_loaders, cfg.workflow)
@@ -0,0 +1,136 @@
# Copyright (c) OpenMMLab. All rights reserved.
import torch
import torch.nn as nn
from custom_mmpkg.custom_mmcv.parallel import MODULE_WRAPPERS, MMDistributedDataParallel
from custom_mmpkg.custom_mmcv.parallel.scatter_gather import scatter_kwargs
from torch.cuda._utils import _get_device_index
@MODULE_WRAPPERS.register_module()
class DistributedDataParallelWrapper(nn.Module):
"""A DistributedDataParallel wrapper for models in 3D mesh estimation task.
In 3D mesh estimation task, there is a need to wrap different modules in
the models with separate DistributedDataParallel. Otherwise, it will cause
errors for GAN training.
More specific, the GAN model, usually has two sub-modules:
generator and discriminator. If we wrap both of them in one
standard DistributedDataParallel, it will cause errors during training,
because when we update the parameters of the generator (or discriminator),
the parameters of the discriminator (or generator) is not updated, which is
not allowed for DistributedDataParallel.
So we design this wrapper to separately wrap DistributedDataParallel
for generator and discriminator.
In this wrapper, we perform two operations:
1. Wrap the modules in the models with separate MMDistributedDataParallel.
Note that only modules with parameters will be wrapped.
2. Do scatter operation for 'forward', 'train_step' and 'val_step'.
Note that the arguments of this wrapper is the same as those in
`torch.nn.parallel.distributed.DistributedDataParallel`.
Args:
module (nn.Module): Module that needs to be wrapped.
device_ids (list[int | `torch.device`]): Same as that in
`torch.nn.parallel.distributed.DistributedDataParallel`.
dim (int, optional): Same as that in the official scatter function in
pytorch. Defaults to 0.
broadcast_buffers (bool): Same as that in
`torch.nn.parallel.distributed.DistributedDataParallel`.
Defaults to False.
find_unused_parameters (bool, optional): Same as that in
`torch.nn.parallel.distributed.DistributedDataParallel`.
Traverse the autograd graph of all tensors contained in returned
value of the wrapped module’s forward function. Defaults to False.
kwargs (dict): Other arguments used in
`torch.nn.parallel.distributed.DistributedDataParallel`.
"""
def __init__(self,
module,
device_ids,
dim=0,
broadcast_buffers=False,
find_unused_parameters=False,
**kwargs):
super().__init__()
assert len(device_ids) == 1, (
'Currently, DistributedDataParallelWrapper only supports one'
'single CUDA device for each process.'
f'The length of device_ids must be 1, but got {len(device_ids)}.')
self.module = module
self.dim = dim
self.to_ddp(
device_ids=device_ids,
dim=dim,
broadcast_buffers=broadcast_buffers,
find_unused_parameters=find_unused_parameters,
**kwargs)
self.output_device = _get_device_index(device_ids[0], True)
def to_ddp(self, device_ids, dim, broadcast_buffers,
find_unused_parameters, **kwargs):
"""Wrap models with separate MMDistributedDataParallel.
It only wraps the modules with parameters.
"""
for name, module in self.module._modules.items():
if next(module.parameters(), None) is None:
module = module.cuda()
elif all(not p.requires_grad for p in module.parameters()):
module = module.cuda()
else:
module = MMDistributedDataParallel(
module.cuda(),
device_ids=device_ids,
dim=dim,
broadcast_buffers=broadcast_buffers,
find_unused_parameters=find_unused_parameters,
**kwargs)
self.module._modules[name] = module
def scatter(self, inputs, kwargs, device_ids):
"""Scatter function.
Args:
inputs (Tensor): Input Tensor.
kwargs (dict): Args for
``mmcv.parallel.scatter_gather.scatter_kwargs``.
device_ids (int): Device id.
"""
return scatter_kwargs(inputs, kwargs, device_ids, dim=self.dim)
def forward(self, *inputs, **kwargs):
"""Forward function.
Args:
inputs (tuple): Input data.
kwargs (dict): Args for
``mmcv.parallel.scatter_gather.scatter_kwargs``.
"""
inputs, kwargs = self.scatter(inputs, kwargs,
[torch.cuda.current_device()])
return self.module(*inputs[0], **kwargs[0])
def train_step(self, *inputs, **kwargs):
"""Train step function.
Args:
inputs (Tensor): Input Tensor.
kwargs (dict): Args for
``mmcv.parallel.scatter_gather.scatter_kwargs``.
"""
inputs, kwargs = self.scatter(inputs, kwargs,
[torch.cuda.current_device()])
output = self.module.train_step(*inputs[0], **kwargs[0])
return output
def val_step(self, *inputs, **kwargs):
"""Validation step function.
Args:
inputs (tuple): Input data.
kwargs (dict): Args for ``scatter_kwargs``.
"""
inputs, kwargs = self.scatter(inputs, kwargs,
[torch.cuda.current_device()])
output = self.module.val_step(*inputs[0], **kwargs[0])
return output
@@ -0,0 +1,4 @@
from motiondiff_modules.mogen.core.evaluation.eval_hooks import DistEvalHook, EvalHook
from motiondiff_modules.mogen.core.evaluation.builder import build_evaluator
__all__ = ["DistEvalHook", "EvalHook", "build_evaluator"]
@@ -0,0 +1,29 @@
import copy
import numpy as np
from custom_mmpkg.custom_mmcv.utils import Registry
from .evaluators.precision_evaluator import PrecisionEvaluator
from .evaluators.matching_score_evaluator import MatchingScoreEvaluator
from .evaluators.fid_evaluator import FIDEvaluator
from .evaluators.diversity_evaluator import DiversityEvaluator
from .evaluators.multimodality_evaluator import MultiModalityEvaluator
EVALUATORS = Registry('evaluators')
EVALUATORS.register_module(name='R Precision', module=PrecisionEvaluator)
EVALUATORS.register_module(name='Matching Score', module=MatchingScoreEvaluator)
EVALUATORS.register_module(name='FID', module=FIDEvaluator)
EVALUATORS.register_module(name='Diversity', module=DiversityEvaluator)
EVALUATORS.register_module(name='MultiModality', module=MultiModalityEvaluator)
def build_evaluator(metric, eval_cfg, data_len, eval_indexes):
cfg = copy.deepcopy(eval_cfg)
cfg.update(metric)
cfg.pop('metrics')
cfg['data_len'] = data_len
cfg['eval_indexes'] = eval_indexes
evaluator = EVALUATORS.build(cfg)
if evaluator.append_indexes is not None:
for i in range(eval_cfg['replication_times']):
eval_indexes[i] = np.concatenate((eval_indexes[i], evaluator.append_indexes[i]), axis=0)
return evaluator, eval_indexes
@@ -0,0 +1,138 @@
# Copyright (c) OpenMMLab. All rights reserved.
import tempfile
import warnings
from custom_mmpkg.custom_mmcv.runner import DistEvalHook as BaseDistEvalHook
from custom_mmpkg.custom_mmcv.runner import EvalHook as BaseEvalHook
mogen_GREATER_KEYS = []
mogen_LESS_KEYS = []
class EvalHook(BaseEvalHook):
def __init__(self,
dataloader,
start=None,
interval=1,
by_epoch=True,
save_best=None,
rule=None,
test_fn=None,
greater_keys=mogen_GREATER_KEYS,
less_keys=mogen_LESS_KEYS,
**eval_kwargs):
if test_fn is None:
from motiondiff_modules.mogen.apis import single_gpu_test
test_fn = single_gpu_test
# remove "gpu_collect" from eval_kwargs
if 'gpu_collect' in eval_kwargs:
warnings.warn(
'"gpu_collect" will be deprecated in EvalHook.'
'Please remove it from the config.', DeprecationWarning)
_ = eval_kwargs.pop('gpu_collect')
# update "save_best" according to "key_indicator" and remove the
# latter from eval_kwargs
if 'key_indicator' in eval_kwargs or isinstance(save_best, bool):
warnings.warn(
'"key_indicator" will be deprecated in EvalHook.'
'Please use "save_best" to specify the metric key,'
'e.g., save_best="pa-mpjpe".', DeprecationWarning)
key_indicator = eval_kwargs.pop('key_indicator', None)
if save_best is True and key_indicator is None:
raise ValueError('key_indicator should not be None, when '
'save_best is set to True.')
save_best = key_indicator
super().__init__(dataloader, start, interval, by_epoch, save_best,
rule, test_fn, greater_keys, less_keys, **eval_kwargs)
def evaluate(self, runner, results):
with tempfile.TemporaryDirectory() as tmp_dir:
eval_res = self.dataloader.dataset.evaluate(
results,
work_dir=tmp_dir,
logger=runner.logger,
**self.eval_kwargs)
for name, val in eval_res.items():
runner.log_buffer.output[name] = val
runner.log_buffer.ready = True
if self.save_best is not None:
if self.key_indicator == 'auto':
self._init_rule(self.rule, list(eval_res.keys())[0])
return eval_res[self.key_indicator]
return None
class DistEvalHook(BaseDistEvalHook):
def __init__(self,
dataloader,
start=None,
interval=1,
by_epoch=True,
save_best=None,
rule=None,
test_fn=None,
greater_keys=mogen_GREATER_KEYS,
less_keys=mogen_LESS_KEYS,
broadcast_bn_buffer=True,
tmpdir=None,
gpu_collect=False,
**eval_kwargs):
if test_fn is None:
from motiondiff_modules.mogen.apis import multi_gpu_test
test_fn = multi_gpu_test
# update "save_best" according to "key_indicator" and remove the
# latter from eval_kwargs
if 'key_indicator' in eval_kwargs or isinstance(save_best, bool):
warnings.warn(
'"key_indicator" will be deprecated in EvalHook.'
'Please use "save_best" to specify the metric key,'
'e.g., save_best="pa-mpjpe".', DeprecationWarning)
key_indicator = eval_kwargs.pop('key_indicator', None)
if save_best is True and key_indicator is None:
raise ValueError('key_indicator should not be None, when '
'save_best is set to True.')
save_best = key_indicator
super().__init__(dataloader, start, interval, by_epoch, save_best,
rule, test_fn, greater_keys, less_keys,
broadcast_bn_buffer, tmpdir, gpu_collect,
**eval_kwargs)
def evaluate(self, runner, results):
"""Evaluate the results.
Args:
runner (:obj:`mmcv.Runner`): The underlined training runner.
results (list): Output results.
"""
with tempfile.TemporaryDirectory() as tmp_dir:
eval_res = self.dataloader.dataset.evaluate(
results,
work_dir=tmp_dir,
logger=runner.logger,
**self.eval_kwargs)
for name, val in eval_res.items():
runner.log_buffer.output[name] = val
runner.log_buffer.ready = True
if self.save_best is not None:
if self.key_indicator == 'auto':
# infer from eval_results
self._init_rule(self.rule, list(eval_res.keys())[0])
return eval_res[self.key_indicator]
return None
@@ -0,0 +1,144 @@
import torch
import numpy as np
from ..utils import get_metric_statistics
class BaseEvaluator(object):
def __init__(self,
batch_size=None,
drop_last=False,
replication_times=1,
replication_reduction='statistics',
eval_begin_idx=None,
eval_end_idx=None):
self.batch_size = batch_size
self.drop_last = drop_last
self.replication_times = replication_times
self.replication_reduction = replication_reduction
assert replication_reduction in ['statistics', 'mean', 'concat']
self.eval_begin_idx = eval_begin_idx
self.eval_end_idx = eval_end_idx
def evaluate(self, results):
total_len = len(results)
partial_len = total_len // self.replication_times
all_metrics = []
for replication_idx in range(self.replication_times):
partial_results = results[
replication_idx * partial_len: (replication_idx + 1) * partial_len]
if self.batch_size is not None:
batch_metrics = []
for batch_start in range(self.eval_begin_idx, self.eval_end_idx, self.batch_size):
batch_results = partial_results[batch_start: batch_start + self.batch_size]
if len(batch_results) < self.batch_size and self.drop_last:
continue
batch_metrics.append(self.single_evaluate(batch_results))
all_metrics.append(self.concat_batch_metrics(batch_metrics))
else:
batch_results = partial_results[self.eval_begin_idx: self.eval_end_idx]
all_metrics.append(self.single_evaluate(batch_results))
all_metrics = np.stack(all_metrics, axis=0)
if self.replication_reduction == 'statistics':
values = get_metric_statistics(all_metrics, self.replication_times)
elif self.replication_reduction == 'mean':
values = np.mean(all_metrics, axis=0)
elif self.replication_reduction == 'concat':
values = all_metrics
return self.parse_values(values)
def prepare_results(self, results):
text = []
pred_motion = []
pred_motion_length = []
pred_motion_mask = []
motion = []
motion_length = []
motion_mask = []
token = []
# count the maximum motion length
T = max([result['motion'].shape[0] for result in results])
for result in results:
cur_motion = result['motion']
if cur_motion.shape[0] < T:
padding_values = torch.zeros((T - cur_motion.shape[0], cur_motion.shape[1]))
padding_values = padding_values.type_as(pred_motion)
cur_motion = torch.cat([cur_motion, padding_values], dim=0)
motion.append(cur_motion)
cur_pred_motion = result['pred_motion']
if cur_pred_motion.shape[0] < T:
padding_values = torch.zeros((T - cur_pred_motion.shape[0], cur_pred_motion.shape[1]))
padding_values = padding_values.type_as(cur_pred_motion)
cur_pred_motion = torch.cat([cur_pred_motion, padding_values], dim=0)
pred_motion.append(cur_pred_motion)
cur_motion_mask = result['motion_mask']
if cur_motion_mask.shape[0] < T:
padding_values = torch.zeros((T - cur_motion_mask.shape[0]))
padding_values = padding_values.type_as(cur_motion_mask)
cur_motion_mask= torch.cat([cur_motion_mask, padding_values], dim=0)
motion_mask.append(cur_motion_mask)
cur_pred_motion_mask = result['pred_motion_mask']
if cur_pred_motion_mask.shape[0] < T:
padding_values = torch.zeros((T - cur_pred_motion_mask.shape[0]))
padding_values = padding_values.type_as(cur_pred_motion_mask)
cur_pred_motion_mask= torch.cat([cur_pred_motion_mask, padding_values], dim=0)
pred_motion_mask.append(cur_pred_motion_mask)
motion_length.append(result['motion_length'].item())
pred_motion_length.append(result['pred_motion_length'].item())
if 'text' in result.keys():
text.append(result['text'])
if 'token' in result.keys():
token.append(result['token'])
motion = torch.stack(motion, dim=0)
pred_motion = torch.stack(pred_motion, dim=0)
motion_mask = torch.stack(motion_mask, dim=0)
pred_motion_mask = torch.stack(pred_motion_mask, dim=0)
motion_length = torch.Tensor(motion_length).to(motion.device).long()
pred_motion_length = torch.Tensor(pred_motion_length).to(motion.device).long()
output = {
'pred_motion': pred_motion,
'pred_motion_mask': pred_motion_mask,
'pred_motion_length': pred_motion_length,
'motion': motion,
'motion_mask': motion_mask,
'motion_length': motion_length,
'text': text,
'token': token
}
return output
def to_device(self, device):
for model in self.model_list:
model.to(device)
def motion_encode(self, motion, motion_length, motion_mask, device):
N = motion.shape[0]
motion_emb = []
batch_size = 32
cur_idx = 0
with torch.no_grad():
while cur_idx < N:
cur_motion = motion[cur_idx: cur_idx + batch_size].to(device)
cur_motion_length = motion_length[cur_idx: cur_idx + batch_size].to(device)
cur_motion_mask = motion_mask[cur_idx: cur_idx + batch_size].to(device)
cur_motion_emb = self.motion_encoder(cur_motion, cur_motion_length, cur_motion_mask)
motion_emb.append(cur_motion_emb)
cur_idx += batch_size
motion_emb = torch.cat(motion_emb, dim=0)
return motion_emb
def text_encode(self, text, token, device):
N = len(text)
text_emb = []
batch_size = 32
cur_idx = 0
with torch.no_grad():
while cur_idx < N:
cur_text = text[cur_idx: cur_idx + batch_size]
cur_token = token[cur_idx: cur_idx + batch_size]
cur_text_emb = self.text_encoder(cur_text, cur_token, device)
text_emb.append(cur_text_emb)
cur_idx += batch_size
text_emb = torch.cat(text_emb, dim=0)
return text_emb
@@ -0,0 +1,52 @@
import numpy as np
import torch
from ..get_model import get_motion_model
from .base_evaluator import BaseEvaluator
from ..utils import calculate_diversity
class DiversityEvaluator(BaseEvaluator):
def __init__(self,
data_len=0,
motion_encoder_name=None,
motion_encoder_path=None,
num_samples=300,
batch_size=None,
drop_last=False,
replication_times=1,
replication_reduction='statistics',
**kwargs):
super().__init__(
replication_times=replication_times,
replication_reduction=replication_reduction,
batch_size=batch_size,
drop_last=drop_last,
eval_begin_idx=0,
eval_end_idx=data_len
)
self.num_samples = num_samples
self.append_indexes = None
self.motion_encoder = get_motion_model(motion_encoder_name, motion_encoder_path)
self.model_list = [self.motion_encoder]
def single_evaluate(self, results):
results = self.prepare_results(results)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
motion = results['motion']
pred_motion = results['pred_motion']
pred_motion_length = results['pred_motion_length']
pred_motion_mask = results['pred_motion_mask']
self.motion_encoder.to(device)
self.motion_encoder.eval()
with torch.no_grad():
pred_motion_emb = self.motion_encode(pred_motion, pred_motion_length, pred_motion_mask, device).cpu().detach().numpy()
diversity = calculate_diversity(pred_motion_emb, self.num_samples)
return diversity
def parse_values(self, values):
metrics = {}
metrics['Diversity (mean)'] = values[0]
metrics['Diversity (conf)'] = values[1]
return metrics
@@ -0,0 +1,58 @@
import numpy as np
import torch
from ..get_model import get_motion_model
from .base_evaluator import BaseEvaluator
from ..utils import (
calculate_activation_statistics,
calculate_frechet_distance)
class FIDEvaluator(BaseEvaluator):
def __init__(self,
data_len=0,
motion_encoder_name=None,
motion_encoder_path=None,
batch_size=None,
drop_last=False,
replication_times=1,
replication_reduction='statistics',
**kwargs):
super().__init__(
replication_times=replication_times,
replication_reduction=replication_reduction,
batch_size=batch_size,
drop_last=drop_last,
eval_begin_idx=0,
eval_end_idx=data_len
)
self.append_indexes = None
self.motion_encoder = get_motion_model(motion_encoder_name, motion_encoder_path)
self.model_list = [self.motion_encoder]
def single_evaluate(self, results):
results = self.prepare_results(results)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
pred_motion = results['pred_motion']
pred_motion_length = results['pred_motion_length']
pred_motion_mask = results['pred_motion_mask']
motion = results['motion']
motion_length = results['motion_length']
motion_mask = results['motion_mask']
self.motion_encoder.to(device)
self.motion_encoder.eval()
with torch.no_grad():
pred_motion_emb = self.motion_encode(pred_motion, pred_motion_length, pred_motion_mask, device).cpu().detach().numpy()
gt_motion_emb = self.motion_encode(motion, motion_length, motion_mask, device).cpu().detach().numpy()
gt_mu, gt_cov = calculate_activation_statistics(gt_motion_emb)
pred_mu, pred_cov = calculate_activation_statistics(pred_motion_emb)
fid = calculate_frechet_distance(gt_mu, gt_cov, pred_mu, pred_cov)
return fid
def parse_values(self, values):
metrics = {}
metrics['FID (mean)'] = values[0]
metrics['FID (conf)'] = values[1]
return metrics
@@ -0,0 +1,71 @@
import numpy as np
import torch
from ..get_model import get_motion_model, get_text_model
from .base_evaluator import BaseEvaluator
from ..utils import calculate_top_k, euclidean_distance_matrix
class MatchingScoreEvaluator(BaseEvaluator):
def __init__(self,
data_len=0,
text_encoder_name=None,
text_encoder_path=None,
motion_encoder_name=None,
motion_encoder_path=None,
top_k=3,
batch_size=32,
drop_last=False,
replication_times=1,
replication_reduction='statistics',
**kwargs):
super().__init__(
replication_times=replication_times,
replication_reduction=replication_reduction,
batch_size=batch_size,
drop_last=drop_last,
eval_begin_idx=0,
eval_end_idx=data_len
)
self.append_indexes = None
self.text_encoder = get_text_model(text_encoder_name, text_encoder_path)
self.motion_encoder = get_motion_model(motion_encoder_name, motion_encoder_path)
self.top_k = top_k
self.model_list = [self.text_encoder, self.motion_encoder]
def single_evaluate(self, results):
results = self.prepare_results(results)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
motion = results['motion']
pred_motion = results['pred_motion']
pred_motion_length = results['pred_motion_length']
pred_motion_mask = results['pred_motion_mask']
text = results['text']
token = results['token']
self.text_encoder.to(device)
self.motion_encoder.to(device)
self.text_encoder.eval()
self.motion_encoder.eval()
with torch.no_grad():
word_emb = self.text_encode(text, token, device=device).cpu().detach().numpy()
motion_emb = self.motion_encode(pred_motion, pred_motion_length, pred_motion_mask, device).cpu().detach().numpy()
dist_mat = euclidean_distance_matrix(word_emb, motion_emb)
matching_score = dist_mat.trace()
all_size = word_emb.shape[0]
return matching_score, all_size
def concat_batch_metrics(self, batch_metrics):
matching_score_sum = 0
all_size = 0
for batch_matching_score, batch_all_size in batch_metrics:
matching_score_sum += batch_matching_score
all_size += batch_all_size
matching_score = matching_score_sum / all_size
return matching_score
def parse_values(self, values):
metrics = {}
metrics['Matching Score (mean)'] = values[0]
metrics['Matching Score (conf)'] = values[1]
return metrics
@@ -0,0 +1,63 @@
import numpy as np
import torch
from ..get_model import get_motion_model
from .base_evaluator import BaseEvaluator
from ..utils import calculate_multimodality
class MultiModalityEvaluator(BaseEvaluator):
def __init__(self,
data_len=0,
motion_encoder_name=None,
motion_encoder_path=None,
num_samples=100,
num_repeats=30,
num_picks=10,
batch_size=None,
drop_last=False,
replication_times=1,
replication_reduction='statistics',
**kwargs):
super().__init__(
replication_times=replication_times,
replication_reduction=replication_reduction,
batch_size=batch_size,
drop_last=drop_last,
eval_begin_idx=data_len,
eval_end_idx=data_len + num_samples * num_repeats
)
self.num_samples = num_samples
self.num_repeats = num_repeats
self.num_picks = num_picks
self.append_indexes = []
for i in range(replication_times):
append_indexes = []
selected_indexs = np.random.choice(data_len, self.num_samples)
for index in selected_indexs:
append_indexes = append_indexes + [index] * self.num_repeats
self.append_indexes.append(np.array(append_indexes))
self.motion_encoder = get_motion_model(motion_encoder_name, motion_encoder_path)
self.model_list = [self.motion_encoder]
def single_evaluate(self, results):
results = self.prepare_results(results)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
motion = results['motion']
pred_motion = results['pred_motion']
pred_motion_length = results['pred_motion_length']
pred_motion_mask = results['pred_motion_mask']
self.motion_encoder.to(device)
self.motion_encoder.eval()
with torch.no_grad():
pred_motion_emb = self.motion_encode(pred_motion, pred_motion_length, pred_motion_mask, device).cpu().detach().numpy()
pred_motion_emb = pred_motion_emb.reshape((self.num_samples, self.num_repeats, -1))
multimodality = calculate_multimodality(pred_motion_emb, self.num_picks)
return multimodality
def parse_values(self, values):
metrics = {}
metrics['MultiModality (mean)'] = values[0]
metrics['MultiModality (conf)'] = values[1]
return metrics
@@ -0,0 +1,74 @@
import numpy as np
import torch
from ..get_model import get_motion_model, get_text_model
from .base_evaluator import BaseEvaluator
from ..utils import calculate_top_k, euclidean_distance_matrix
class PrecisionEvaluator(BaseEvaluator):
def __init__(self,
data_len=0,
text_encoder_name=None,
text_encoder_path=None,
motion_encoder_name=None,
motion_encoder_path=None,
top_k=3,
batch_size=32,
drop_last=False,
replication_times=1,
replication_reduction='statistics',
**kwargs):
super().__init__(
replication_times=replication_times,
replication_reduction=replication_reduction,
batch_size=batch_size,
drop_last=drop_last,
eval_begin_idx=0,
eval_end_idx=data_len
)
self.append_indexes = None
self.text_encoder = get_text_model(text_encoder_name, text_encoder_path)
self.motion_encoder = get_motion_model(motion_encoder_name, motion_encoder_path)
self.top_k = top_k
self.model_list = [self.text_encoder, self.motion_encoder]
def single_evaluate(self, results):
results = self.prepare_results(results)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
motion = results['motion']
pred_motion = results['pred_motion']
pred_motion_length = results['pred_motion_length']
pred_motion_mask = results['pred_motion_mask']
text = results['text']
token = results['token']
self.text_encoder.to(device)
self.motion_encoder.to(device)
self.text_encoder.eval()
self.motion_encoder.eval()
with torch.no_grad():
word_emb = self.text_encode(text, token, device=device).cpu().detach().numpy()
motion_emb = self.motion_encode(pred_motion, pred_motion_length, pred_motion_mask, device).cpu().detach().numpy()
dist_mat = euclidean_distance_matrix(word_emb, motion_emb)
argsmax = np.argsort(dist_mat, axis=1)
top_k_mat = calculate_top_k(argsmax, top_k=self.top_k)
top_k_count = top_k_mat.sum(axis=0)
all_size = word_emb.shape[0]
return top_k_count, all_size
def concat_batch_metrics(self, batch_metrics):
top_k_count = 0
all_size = 0
for batch_top_k_count, batch_all_size in batch_metrics:
top_k_count += batch_top_k_count
all_size += batch_all_size
R_precision = top_k_count / all_size
return R_precision
def parse_values(self, values):
metrics = {}
for top_k in range(self.top_k):
metrics['R_precision Top %d (mean)' % (top_k + 1)] = values[0][top_k]
metrics['R_precision Top %d (conf)' % (top_k + 1)] = values[1][top_k]
return metrics
@@ -0,0 +1,46 @@
from motiondiff_modules.mogen.models import build_submodule
def get_motion_model(name, ckpt_path):
if name == 'kit_ml':
model = build_submodule(dict(
type='T2MMotionEncoder',
input_size=251,
movement_hidden_size=512,
movement_latent_size=512,
motion_hidden_size=1024,
motion_latent_size=512,
))
else:
model = build_submodule(dict(
type='T2MMotionEncoder',
input_size=263,
movement_hidden_size=512,
movement_latent_size=512,
motion_hidden_size=1024,
motion_latent_size=512,
))
model.load_pretrained(ckpt_path)
return model
def get_text_model(name, ckpt_path):
if name == 'kit_ml':
model = build_submodule(dict(
type='T2MTextEncoder',
word_size=300,
pos_size=15,
hidden_size=512,
output_size=512,
max_text_len=20
))
else:
model = build_submodule(dict(
type='T2MTextEncoder',
word_size=300,
pos_size=15,
hidden_size=512,
output_size=512,
max_text_len=20
))
model.load_pretrained(ckpt_path)
return model
@@ -0,0 +1,130 @@
import numpy as np
from scipy import linalg
def get_metric_statistics(values, replication_times):
mean = np.mean(values, axis=0)
std = np.std(values, axis=0)
conf_interval = 1.96 * std / np.sqrt(replication_times)
return mean, conf_interval
# (X - X_train)*(X - X_train) = -2X*X_train + X*X + X_train*X_train
def euclidean_distance_matrix(matrix1, matrix2):
"""
Params:
-- matrix1: N1 x D
-- matrix2: N2 x D
Returns:
-- dist: N1 x N2
dist[i, j] == distance(matrix1[i], matrix2[j])
"""
assert matrix1.shape[1] == matrix2.shape[1]
d1 = -2 * np.dot(matrix1, matrix2.T) # shape (num_test, num_train)
d2 = np.sum(np.square(matrix1), axis=1, keepdims=True) # shape (num_test, 1)
d3 = np.sum(np.square(matrix2), axis=1) # shape (num_train, )
dists = np.sqrt(d1 + d2 + d3) # broadcasting
return dists
def calculate_top_k(mat, top_k):
size = mat.shape[0]
gt_mat = np.expand_dims(np.arange(size), 1).repeat(size, 1)
bool_mat = (mat == gt_mat)
correct_vec = False
top_k_list = []
for i in range(top_k):
# print(correct_vec, bool_mat[:, i])
correct_vec = (correct_vec | bool_mat[:, i])
# print(correct_vec)
top_k_list.append(correct_vec[:, None])
top_k_mat = np.concatenate(top_k_list, axis=1)
return top_k_mat
def calculate_activation_statistics(activations):
"""
Params:
-- activation: num_samples x dim_feat
Returns:
-- mu: dim_feat
-- sigma: dim_feat x dim_feat
"""
mu = np.mean(activations, axis=0)
cov = np.cov(activations, rowvar=False)
return mu, cov
def calculate_frechet_distance(mu1, sigma1, mu2, sigma2, eps=1e-6):
"""Numpy implementation of the Frechet Distance.
The Frechet distance between two multivariate Gaussians X_1 ~ N(mu_1, C_1)
and X_2 ~ N(mu_2, C_2) is
d^2 = ||mu_1 - mu_2||^2 + Tr(C_1 + C_2 - 2*sqrt(C_1*C_2)).
Stable version by Dougal J. Sutherland.
Params:
-- mu1 : Numpy array containing the activations of a layer of the
inception net (like returned by the function 'get_predictions')
for generated samples.
-- mu2 : The sample mean over activations, precalculated on an
representative data set.
-- sigma1: The covariance matrix over activations for generated samples.
-- sigma2: The covariance matrix over activations, precalculated on an
representative data set.
Returns:
-- : The Frechet Distance.
"""
mu1 = np.atleast_1d(mu1)
mu2 = np.atleast_1d(mu2)
sigma1 = np.atleast_2d(sigma1)
sigma2 = np.atleast_2d(sigma2)
assert mu1.shape == mu2.shape, \
'Training and test mean vectors have different lengths'
assert sigma1.shape == sigma2.shape, \
'Training and test covariances have different dimensions'
diff = mu1 - mu2
# Product might be almost singular
covmean, _ = linalg.sqrtm(sigma1.dot(sigma2), disp=False)
if not np.isfinite(covmean).all():
msg = ('fid calculation produces singular product; '
'adding %s to diagonal of cov estimates') % eps
print(msg)
offset = np.eye(sigma1.shape[0]) * eps
covmean = linalg.sqrtm((sigma1 + offset).dot(sigma2 + offset))
# Numerical error might give slight imaginary component
if np.iscomplexobj(covmean):
if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
m = np.max(np.abs(covmean.imag))
raise ValueError('Imaginary component {}'.format(m))
covmean = covmean.real
tr_covmean = np.trace(covmean)
return (diff.dot(diff) + np.trace(sigma1) +
np.trace(sigma2) - 2 * tr_covmean)
def calculate_diversity(activation, diversity_times):
assert len(activation.shape) == 2
assert activation.shape[0] > diversity_times
num_samples = activation.shape[0]
first_indices = np.random.choice(num_samples, diversity_times, replace=False)
second_indices = np.random.choice(num_samples, diversity_times, replace=False)
dist = linalg.norm(activation[first_indices] - activation[second_indices], axis=1)
return dist.mean()
def calculate_multimodality(activation, multimodality_times):
assert len(activation.shape) == 3
assert activation.shape[1] > multimodality_times
num_per_sent = activation.shape[1]
first_dices = np.random.choice(num_per_sent, multimodality_times, replace=False)
second_dices = np.random.choice(num_per_sent, multimodality_times, replace=False)
dist = linalg.norm(activation[:, first_dices] - activation[:, second_dices], axis=2)
return dist.mean()
@@ -0,0 +1,3 @@
from .builder import OPTIMIZERS, build_optimizers
__all__ = ['build_optimizers', 'OPTIMIZERS']
@@ -0,0 +1,52 @@
# Copyright (c) OpenMMLab. All rights reserved.
from custom_mmpkg.custom_mmcv.runner import build_optimizer
from custom_mmpkg.custom_mmcv.utils import Registry
OPTIMIZERS = Registry('optimizers')
def build_optimizers(model, cfgs):
"""Build multiple optimizers from configs. If `cfgs` contains several dicts
for optimizers, then a dict for each constructed optimizers will be
returned. If `cfgs` only contains one optimizer config, the constructed
optimizer itself will be returned. For example,
1) Multiple optimizer configs:
.. code-block:: python
optimizer_cfg = dict(
model1=dict(type='SGD', lr=lr),
model2=dict(type='SGD', lr=lr))
The return dict is
``dict('model1': torch.optim.Optimizer, 'model2': torch.optim.Optimizer)``
2) Single optimizer config:
.. code-block:: python
optimizer_cfg = dict(type='SGD', lr=lr)
The return is ``torch.optim.Optimizer``.
Args:
model (:obj:`nn.Module`): The model with parameters to be optimized.
cfgs (dict): The config dict of the optimizer.
Returns:
dict[:obj:`torch.optim.Optimizer`] | :obj:`torch.optim.Optimizer`:
The initialized optimizers.
"""
optimizers = {}
if hasattr(model, 'module'):
model = model.module
# determine whether 'cfgs' has several dicts for optimizers
if all(isinstance(v, dict) for v in cfgs.values()):
for key, cfg in cfgs.items():
cfg_ = cfg.copy()
module = getattr(model, key)
optimizers[key] = build_optimizer(module, cfg_)
return optimizers
return build_optimizer(model, cfgs)
@@ -0,0 +1,11 @@
from .base_dataset import BaseMotionDataset
from .text_motion_dataset import TextMotionDataset
from .builder import DATASETS, PIPELINES, build_dataloader, build_dataset
from .pipelines import Compose
from .samplers import DistributedSampler
__all__ = [
'BaseMotionDataset', 'TextMotionDataset', 'DATASETS', 'PIPELINES', 'build_dataloader',
'build_dataset', 'Compose', 'DistributedSampler'
]
@@ -0,0 +1,117 @@
import os
import copy
from typing import Optional, Union
import numpy as np
from torch.utils.data import Dataset
from .pipelines import Compose
from .builder import DATASETS
from motiondiff_modules.mogen.core.evaluation import build_evaluator
@DATASETS.register_module()
class BaseMotionDataset(Dataset):
"""Base motion dataset.
Args:
data_prefix (str): the prefix of data path.
pipeline (list): a list of dict, where each element represents
a operation defined in `mogen.datasets.pipelines`.
ann_file (str | None, optional): the annotation file. When ann_file is
str, the subclass is expected to read from the ann_file. When
ann_file is None, the subclass is expected to read according
to data_prefix.
test_mode (bool): in train mode or test mode. Default: None.
dataset_name (str | None, optional): the name of dataset. It is used
to identify the type of evaluation metric. Default: None.
"""
def __init__(self,
data_prefix: str,
pipeline: list,
dataset_name: Optional[Union[str, None]] = None,
fixed_length: Optional[Union[int, None]] = None,
ann_file: Optional[Union[str, None]] = None,
motion_dir: Optional[Union[str, None]] = None,
eval_cfg: Optional[Union[dict, None]] = None,
test_mode: Optional[bool] = False):
super(BaseMotionDataset, self).__init__()
self.data_prefix = data_prefix
self.pipeline = Compose(pipeline)
self.dataset_name = dataset_name
self.fixed_length = fixed_length
self.ann_file = os.path.join(data_prefix, 'datasets', dataset_name, ann_file)
self.motion_dir = os.path.join(data_prefix, 'datasets', dataset_name, motion_dir)
self.eval_cfg = copy.deepcopy(eval_cfg)
self.test_mode = test_mode
self.load_annotations()
if self.test_mode:
self.prepare_evaluation()
def load_anno(self, name):
motion_path = os.path.join(self.motion_dir, name + '.npy')
motion_data = np.load(motion_path)
return {'motion': motion_data}
def load_annotations(self):
"""Load annotations from ``ann_file`` to ``data_infos``"""
self.data_infos = []
for line in open(self.ann_file, 'r').readlines():
line = line.strip()
self.data_infos.append(self.load_anno(line))
def prepare_data(self, idx: int):
""""Prepare raw data for the f'{idx'}-th data."""
results = copy.deepcopy(self.data_infos[idx])
results['dataset_name'] = self.dataset_name
results['sample_idx'] = idx
return self.pipeline(results)
def __len__(self):
"""Return the length of current dataset."""
if self.test_mode:
return len(self.eval_indexes)
elif self.fixed_length is not None:
return self.fixed_length
return len(self.data_infos)
def __getitem__(self, idx: int):
"""Prepare data for the ``idx``-th data.
As for video dataset, we can first parse raw data for each frame. Then
we combine annotations from all frames. This interface is used to
simplify the logic of video dataset and other special datasets.
"""
if self.test_mode:
idx = self.eval_indexes[idx]
elif self.fixed_length is not None:
idx = idx % len(self.data_infos)
return self.prepare_data(idx)
def prepare_evaluation(self):
self.evaluators = []
self.eval_indexes = []
for _ in range(self.eval_cfg['replication_times']):
eval_indexes = np.arange(len(self.data_infos))
if self.eval_cfg.get('shuffle_indexes', False):
np.random.shuffle(eval_indexes)
self.eval_indexes.append(eval_indexes)
for metric in self.eval_cfg['metrics']:
evaluator, self.eval_indexes = build_evaluator(
metric, self.eval_cfg, len(self.data_infos), self.eval_indexes)
self.evaluators.append(evaluator)
self.eval_indexes = np.concatenate(self.eval_indexes)
def evaluate(self, results, work_dir, logger=None):
metrics = {}
device = results[0]['motion'].device
for evaluator in self.evaluators:
evaluator.to_device(device)
metrics.update(evaluator.evaluate(results))
if logger is not None:
logger.info(metrics)
return metrics
@@ -0,0 +1,113 @@
import platform
import random
from functools import partial
from typing import Optional, Union
import numpy as np
from custom_mmpkg.custom_mmcv.parallel import collate
from custom_mmpkg.custom_mmcv.runner import get_dist_info
from custom_mmpkg.custom_mmcv.utils import Registry, build_from_cfg
from torch.utils.data import DataLoader
from torch.utils.data.dataset import Dataset
from .samplers import DistributedSampler
if platform.system() != 'Windows':
# https://github.com/pytorch/pytorch/issues/973
import resource
rlimit = resource.getrlimit(resource.RLIMIT_NOFILE)
base_soft_limit = rlimit[0]
hard_limit = rlimit[1]
soft_limit = min(max(4096, base_soft_limit), hard_limit)
resource.setrlimit(resource.RLIMIT_NOFILE, (soft_limit, hard_limit))
DATASETS = Registry('dataset')
PIPELINES = Registry('pipeline')
def build_dataset(cfg: Union[dict, list, tuple],
default_args: Optional[Union[dict, None]] = None):
""""Build dataset by the given config."""
from .dataset_wrappers import (
ConcatDataset,
RepeatDataset,
)
if isinstance(cfg, (list, tuple)):
dataset = ConcatDataset([build_dataset(c, default_args) for c in cfg])
elif cfg['type'] == 'RepeatDataset':
dataset = RepeatDataset(
build_dataset(cfg['dataset'], default_args), cfg['times'])
else:
dataset = build_from_cfg(cfg, DATASETS, default_args)
return dataset
def build_dataloader(dataset: Dataset,
samples_per_gpu: int,
workers_per_gpu: int,
num_gpus: Optional[int] = 1,
dist: Optional[bool] = True,
shuffle: Optional[bool] = True,
round_up: Optional[bool] = True,
seed: Optional[Union[int, None]] = None,
persistent_workers: Optional[bool] = True,
**kwargs):
"""Build PyTorch DataLoader.
In distributed training, each GPU/process has a dataloader.
In non-distributed training, there is only one dataloader for all GPUs.
Args:
dataset (:obj:`Dataset`): A PyTorch dataset.
samples_per_gpu (int): Number of training samples on each GPU, i.e.,
batch size of each GPU.
workers_per_gpu (int): How many subprocesses to use for data loading
for each GPU.
num_gpus (int, optional): Number of GPUs. Only used in non-distributed
training.
dist (bool, optional): Distributed training/test or not. Default: True.
shuffle (bool, optional): Whether to shuffle the data at every epoch.
Default: True.
round_up (bool, optional): Whether to round up the length of dataset by
adding extra samples to make it evenly divisible. Default: True.
kwargs: any keyword argument to be used to initialize DataLoader
Returns:
DataLoader: A PyTorch dataloader.
"""
rank, world_size = get_dist_info()
if dist:
sampler = DistributedSampler(
dataset, world_size, rank, shuffle=shuffle, round_up=round_up)
shuffle = False
batch_size = samples_per_gpu
num_workers = workers_per_gpu
else:
sampler = None
batch_size = num_gpus * samples_per_gpu
num_workers = num_gpus * workers_per_gpu
init_fn = partial(
worker_init_fn, num_workers=num_workers, rank=rank,
seed=seed) if seed is not None else None
data_loader = DataLoader(
dataset,
batch_size=batch_size,
sampler=sampler,
num_workers=num_workers,
collate_fn=partial(collate, samples_per_gpu=samples_per_gpu),
pin_memory=False,
shuffle=shuffle,
worker_init_fn=init_fn,
persistent_workers=persistent_workers,
**kwargs)
return data_loader
def worker_init_fn(worker_id: int, num_workers: int, rank: int, seed: int):
"""Init random seed for each worker."""
# The seed of each worker equals to
# num_worker * rank + worker_id + user_seed
worker_seed = num_workers * rank + worker_id + seed
np.random.seed(worker_seed)
random.seed(worker_seed)
@@ -0,0 +1,42 @@
from torch.utils.data.dataset import ConcatDataset as _ConcatDataset
from torch.utils.data.dataset import Dataset
from .builder import DATASETS
@DATASETS.register_module()
class ConcatDataset(_ConcatDataset):
"""A wrapper of concatenated dataset.
Same as :obj:`torch.utils.data.dataset.ConcatDataset`, but
add `get_cat_ids` function.
Args:
datasets (list[:obj:`Dataset`]): A list of datasets.
"""
def __init__(self, datasets: list):
super(ConcatDataset, self).__init__(datasets)
@DATASETS.register_module()
class RepeatDataset(object):
"""A wrapper of repeated dataset.
The length of repeated dataset will be `times` larger than the original
dataset. This is useful when the data loading time is long but the dataset
is small. Using RepeatDataset can reduce the data loading time between
epochs.
Args:
dataset (:obj:`Dataset`): The dataset to be repeated.
times (int): Repeat times.
"""
def __init__(self, dataset: Dataset, times: int):
self.dataset = dataset
self.times = times
self._ori_len = len(self.dataset)
def __getitem__(self, idx: int):
return self.dataset[idx % self._ori_len]
def __len__(self):
return self.times * self._ori_len
@@ -0,0 +1,18 @@
from .compose import Compose
from .formatting import (
to_tensor,
ToTensor,
Transpose,
Collect,
WrapFieldsToLists
)
from .transforms import (
Crop,
RandomCrop,
Normalize
)
__all__ = [
'Compose', 'to_tensor', 'Transpose', 'Collect', 'WrapFieldsToLists', 'ToTensor',
'Crop', 'RandomCrop', 'Normalize'
]
@@ -0,0 +1,42 @@
from collections.abc import Sequence
from custom_mmpkg.custom_mmcv.utils import build_from_cfg
from ..builder import PIPELINES
@PIPELINES.register_module()
class Compose(object):
"""Compose a data pipeline with a sequence of transforms.
Args:
transforms (list[dict | callable]):
Either config dicts of transforms or transform objects.
"""
def __init__(self, transforms):
assert isinstance(transforms, Sequence)
self.transforms = []
for transform in transforms:
if isinstance(transform, dict):
transform = build_from_cfg(transform, PIPELINES)
self.transforms.append(transform)
elif callable(transform):
self.transforms.append(transform)
else:
raise TypeError('transform must be callable or a dict, but got'
f' {type(transform)}')
def __call__(self, data):
for t in self.transforms:
data = t(data)
if data is None:
return None
return data
def __repr__(self):
format_string = self.__class__.__name__ + '('
for t in self.transforms:
format_string += f'\n {t}'
format_string += '\n)'
return format_string
@@ -0,0 +1,134 @@
from collections.abc import Sequence
import custom_mmpkg.custom_mmcv as mmcv
import numpy as np
import torch
from custom_mmpkg.custom_mmcv.parallel import DataContainer as DC
from PIL import Image
from ..builder import PIPELINES
def to_tensor(data):
"""Convert objects of various python types to :obj:`torch.Tensor`.
Supported types are: :class:`numpy.ndarray`, :class:`torch.Tensor`,
:class:`Sequence`, :class:`int` and :class:`float`.
"""
if isinstance(data, torch.Tensor):
return data
elif isinstance(data, np.ndarray):
return torch.from_numpy(data)
elif isinstance(data, Sequence) and not mmcv.is_str(data):
return torch.tensor(data)
elif isinstance(data, int):
return torch.LongTensor([data])
elif isinstance(data, float):
return torch.FloatTensor([data])
else:
raise TypeError(
f'Type {type(data)} cannot be converted to tensor.'
'Supported types are: `numpy.ndarray`, `torch.Tensor`, '
'`Sequence`, `int` and `float`')
@PIPELINES.register_module()
class ToTensor(object):
def __init__(self, keys):
self.keys = keys
def __call__(self, results):
for key in self.keys:
results[key] = to_tensor(results[key])
return results
def __repr__(self):
return self.__class__.__name__ + f'(keys={self.keys})'
@PIPELINES.register_module()
class Transpose(object):
def __init__(self, keys, order):
self.keys = keys
self.order = order
def __call__(self, results):
for key in self.keys:
results[key] = results[key].transpose(self.order)
return results
def __repr__(self):
return self.__class__.__name__ + \
f'(keys={self.keys}, order={self.order})'
@PIPELINES.register_module()
class Collect(object):
"""Collect data from the loader relevant to the specific task.
This is usually the last stage of the data loader pipeline.
Args:
keys (Sequence[str]): Keys of results to be collected in ``data``.
meta_keys (Sequence[str], optional): Meta keys to be converted to
``mmcv.DataContainer`` and collected in ``data[motion_metas]``.
Default: ``('filename', 'ori_filename', 'ori_shape', 'motion_shape', 'motion_mask')``
Returns:
dict: The result dict contains the following keys
- keys in``self.keys``
- ``motion_metas`` if available
"""
def __init__(self,
keys,
meta_keys=('filename', 'ori_filename', 'ori_shape', 'motion_shape', 'motion_mask')):
self.keys = keys
self.meta_keys = meta_keys
def __call__(self, results):
data = {}
motion_meta = {}
for key in self.meta_keys:
if key in results:
motion_meta[key] = results[key]
data['motion_metas'] = DC(motion_meta, cpu_only=True)
for key in self.keys:
data[key] = results[key]
return data
def __repr__(self):
return self.__class__.__name__ + \
f'(keys={self.keys}, meta_keys={self.meta_keys})'
@PIPELINES.register_module()
class WrapFieldsToLists(object):
"""Wrap fields of the data dictionary into lists for evaluation.
This class can be used as a last step of a test or validation
pipeline for single image evaluation or inference.
Example:
>>> test_pipeline = [
>>> dict(type='LoadImageFromFile'),
>>> dict(type='Normalize',
mean=[123.675, 116.28, 103.53],
std=[58.395, 57.12, 57.375],
to_rgb=True),
>>> dict(type='ImageToTensor', keys=['img']),
>>> dict(type='Collect', keys=['img']),
>>> dict(type='WrapIntoLists')
>>> ]
"""
def __call__(self, results):
# Wrap dict fields into lists
for key, val in results.items():
results[key] = [val]
return results
def __repr__(self):
return f'{self.__class__.__name__}()'
@@ -0,0 +1,120 @@
import math
import random
import custom_mmpkg.custom_mmcv as mmcv
import numpy as np
from ..builder import PIPELINES
import torch
from typing import Optional, Tuple, Union
@PIPELINES.register_module()
class Crop(object):
r"""Crop motion sequences.
Args:
crop_size (int): The size of the cropped motion sequence.
"""
def __init__(self,
crop_size: Optional[Union[int, None]] = None):
self.crop_size = crop_size
assert self.crop_size is not None
def __call__(self, results):
motion = results['motion']
length = len(motion)
if length >= self.crop_size:
idx = random.randint(0, length - self.crop_size)
motion = motion[idx: idx + self.crop_size]
results['motion_length'] = self.crop_size
else:
padding_length = self.crop_size - length
D = motion.shape[1:]
padding_zeros = np.zeros((padding_length, *D), dtype=np.float32)
motion = np.concatenate([motion, padding_zeros], axis=0)
results['motion_length'] = length
assert len(motion) == self.crop_size
results['motion'] = motion
results['motion_shape'] = motion.shape
if length >= self.crop_size:
results['motion_mask'] = torch.ones(self.crop_size).numpy()
else:
results['motion_mask'] = torch.cat(
(torch.ones(length), torch.zeros(self.crop_size - length))).numpy()
return results
def __repr__(self):
repr_str = self.__class__.__name__ + f'(crop_size={self.crop_size})'
return repr_str
@PIPELINES.register_module()
class RandomCrop(object):
r"""Random crop motion sequences. Each sequence will be padded with zeros to the maximum length.
Args:
min_size (int or None): The minimum size of the cropped motion sequence (inclusive).
max_size (int or None): The maximum size of the cropped motion sequence (inclusive).
"""
def __init__(self,
min_size: Optional[Union[int, None]] = None,
max_size: Optional[Union[int, None]] = None):
self.min_size = min_size
self.max_size = max_size
assert self.min_size is not None
assert self.max_size is not None
def __call__(self, results):
motion = results['motion']
length = len(motion)
crop_size = random.randint(self.min_size, self.max_size)
if length > crop_size:
idx = random.randint(0, length - crop_size)
motion = motion[idx: idx + crop_size]
results['motion_length'] = crop_size
else:
results['motion_length'] = length
padding_length = self.max_size - min(crop_size, length)
if padding_length > 0:
D = motion.shape[1:]
padding_zeros = np.zeros((padding_length, *D), dtype=np.float32)
motion = np.concatenate([motion, padding_zeros], axis=0)
results['motion'] = motion
results['motion_shape'] = motion.shape
if length >= self.max_size and crop_size == self.max_size:
results['motion_mask'] = torch.ones(self.max_size).numpy()
else:
results['motion_mask'] = torch.cat((
torch.ones(min(length, crop_size)),
torch.zeros(self.max_size - min(length, crop_size))), dim=0).numpy()
assert len(motion) == self.max_size
return results
def __repr__(self):
repr_str = self.__class__.__name__ + f'(min_size={self.min_size}'
repr_str += f', max_size={self.max_size})'
return repr_str
@PIPELINES.register_module()
class Normalize(object):
"""Normalize motion sequences.
Args:
mean_path (str): Path of mean file.
std_path (str): Path of std file.
"""
def __init__(self, mean_path, std_path, eps=1e-9):
self.mean = np.load(mean_path)
self.std = np.load(std_path)
self.eps = eps
def __call__(self, results):
motion = results['motion']
motion = (motion - self.mean) / (self.std + self.eps)
results['motion'] = motion
results['motion_norm_mean'] = self.mean
results['motion_norm_std'] = self.std
return results
@@ -0,0 +1,3 @@
from .distributed_sampler import DistributedSampler
__all__ = ['DistributedSampler']
@@ -0,0 +1,42 @@
import torch
from torch.utils.data import DistributedSampler as _DistributedSampler
class DistributedSampler(_DistributedSampler):
def __init__(self,
dataset,
num_replicas=None,
rank=None,
shuffle=True,
round_up=True):
super().__init__(dataset, num_replicas=num_replicas, rank=rank)
self.shuffle = shuffle
self.round_up = round_up
if self.round_up:
self.total_size = self.num_samples * self.num_replicas
else:
self.total_size = len(self.dataset)
def __iter__(self):
# deterministically shuffle based on epoch
if self.shuffle:
g = torch.Generator()
g.manual_seed(self.epoch)
indices = torch.randperm(len(self.dataset), generator=g).tolist()
else:
indices = torch.arange(len(self.dataset)).tolist()
# add extra samples to make it evenly divisible
if self.round_up:
indices = (
indices *
int(self.total_size / len(indices) + 1))[:self.total_size]
assert len(indices) == self.total_size
# subsample
indices = indices[self.rank:self.total_size:self.num_replicas]
if self.round_up:
assert len(indices) == self.num_samples
return iter(indices)
@@ -0,0 +1,93 @@
import json
import os
import os.path
from abc import ABCMeta
from collections import OrderedDict
from typing import Any, List, Optional, Union
import custom_mmpkg.custom_mmcv as mmcv
import copy
import numpy as np
import torch
import torch.distributed as dist
from custom_mmpkg.custom_mmcv.runner import get_dist_info
from .base_dataset import BaseMotionDataset
from .builder import DATASETS
@DATASETS.register_module()
class TextMotionDataset(BaseMotionDataset):
"""TextMotion dataset.
Args:
text_dir (str): Path to the directory containing the text files.
"""
def __init__(self,
data_prefix: str,
pipeline: list,
dataset_name: Optional[Union[str, None]] = None,
fixed_length: Optional[Union[int, None]] = None,
ann_file: Optional[Union[str, None]] = None,
motion_dir: Optional[Union[str, None]] = None,
text_dir: Optional[Union[str, None]] = None,
token_dir: Optional[Union[str, None]] = None,
clip_feat_dir: Optional[Union[str, None]] = None,
eval_cfg: Optional[Union[dict, None]] = None,
fine_mode: Optional[bool] = False,
test_mode: Optional[bool] = False):
self.text_dir = os.path.join(data_prefix, 'datasets', dataset_name, text_dir)
if token_dir is not None:
self.token_dir = os.path.join(data_prefix, 'datasets', dataset_name, token_dir)
else:
self.token_dir = None
if clip_feat_dir is not None:
self.clip_feat_dir = os.path.join(data_prefix, 'datasets', dataset_name, clip_feat_dir)
else:
self.clip_feat_dir = None
self.fine_mode = fine_mode
super(TextMotionDataset, self).__init__(
data_prefix=data_prefix,
pipeline=pipeline,
dataset_name=dataset_name,
fixed_length=fixed_length,
ann_file=ann_file,
motion_dir=motion_dir,
eval_cfg=eval_cfg,
test_mode=test_mode)
def load_anno(self, name):
results = super().load_anno(name)
text_path = os.path.join(self.text_dir, name + '.txt')
text_data = []
for line in open(text_path, 'r'):
text_data.append(line.strip())
results['text'] = text_data
if self.token_dir is not None:
token_path = os.path.join(self.token_dir, name + '.txt')
token_data = []
for line in open(token_path, 'r'):
token_data.append(line.strip())
results['token'] = token_data
if self.clip_feat_dir is not None:
clip_feat_path = os.path.join(self.clip_feat_dir, name + '.npy')
clip_feat = torch.from_numpy(np.load(clip_feat_path))
results['clip_feat'] = clip_feat
return results
def prepare_data(self, idx: int):
""""Prepare raw data for the f'{idx'}-th data."""
results = copy.deepcopy(self.data_infos[idx])
text_list = results['text']
idx = np.random.randint(0, len(text_list))
if self.fine_mode:
results['text'] = json.loads(text_list[idx])
else:
results['text'] = text_list[idx]
if 'clip_feat' in results.keys():
results['clip_feat'] = results['clip_feat'][idx]
if 'token' in results.keys():
results['token'] = results['token'][idx]
results['dataset_name'] = self.dataset_name
results['sample_idx'] = idx
return self.pipeline(results)
@@ -0,0 +1,7 @@
from .architectures import *
from .losses import *
from .rnns import *
from .transformers import *
from .attentions import *
from .builder import *
from .utils import *
@@ -0,0 +1,6 @@
from .vae_architecture import MotionVAE
from .diffusion_architecture import MotionDiffusion
__all__ = [
'MotionVAE', 'MotionDiffusion'
]
@@ -0,0 +1,135 @@
from abc import ABCMeta, abstractmethod
from collections import OrderedDict
import torch
import torch.distributed as dist
from custom_mmpkg.custom_mmcv.runner import BaseModule
def to_cpu(x):
if isinstance(x, torch.Tensor):
return x.detach().cpu()
return x
class BaseArchitecture(BaseModule):
"""Base class for mogen architecture."""
def __init__(self, init_cfg=None):
super(BaseArchitecture, self).__init__(init_cfg)
def forward_train(self, **kwargs):
pass
def forward_test(self, **kwargs):
pass
def _parse_losses(self, losses):
"""Parse the raw outputs (losses) of the network.
Args:
losses (dict): Raw output of the network, which usually contain
losses and other necessary information.
Returns:
tuple[Tensor, dict]: (loss, log_vars), loss is the loss tensor \
which may be a weighted sum of all losses, log_vars contains \
all the variables to be sent to the logger.
"""
log_vars = OrderedDict()
for loss_name, loss_value in losses.items():
if isinstance(loss_value, torch.Tensor):
log_vars[loss_name] = loss_value.mean()
elif isinstance(loss_value, list):
log_vars[loss_name] = sum(_loss.mean() for _loss in loss_value)
else:
raise TypeError(
f'{loss_name} is not a tensor or list of tensors')
loss = sum(_value for _key, _value in log_vars.items()
if 'loss' in _key)
log_vars['loss'] = loss
for loss_name, loss_value in log_vars.items():
# reduce loss when distributed training
if dist.is_available() and dist.is_initialized():
loss_value = loss_value.data.clone()
dist.all_reduce(loss_value.div_(dist.get_world_size()))
log_vars[loss_name] = loss_value.item()
return loss, log_vars
def train_step(self, data, optimizer):
"""The iteration step during training.
This method defines an iteration step during training, except for the
back propagation and optimizer updating, which are done in an optimizer
hook. Note that in some complicated cases or models, the whole process
including back propagation and optimizer updating is also defined in
this method, such as GAN.
Args:
data (dict): The output of dataloader.
optimizer (:obj:`torch.optim.Optimizer` | dict): The optimizer of
runner is passed to ``train_step()``. This argument is unused
and reserved.
Returns:
dict: It should contain at least 3 keys: ``loss``, ``log_vars``, \
``num_samples``.
- ``loss`` is a tensor for back propagation, which can be a
weighted sum of multiple losses.
- ``log_vars`` contains all the variables to be sent to the
logger.
- ``num_samples`` indicates the batch size (when the model is
DDP, it means the batch size on each GPU), which is used for
averaging the logs.
"""
losses = self(**data)
loss, log_vars = self._parse_losses(losses)
outputs = dict(
loss=loss, log_vars=log_vars, num_samples=len(data['motion']))
return outputs
def val_step(self, data, optimizer=None):
"""The iteration step during validation.
This method shares the same signature as :func:`train_step`, but used
during val epochs. Note that the evaluation after training epochs is
not implemented with this method, but an evaluation hook.
"""
losses = self(**data)
loss, log_vars = self._parse_losses(losses)
outputs = dict(
loss=loss, log_vars=log_vars, num_samples=len(data['motion']))
return outputs
def forward(self, **kwargs):
if self.training:
return self.forward_train(**kwargs)
else:
return self.forward_test(**kwargs)
def split_results(self, results):
B = results['motion'].shape[0]
output = []
for i in range(B):
batch_output = dict()
batch_output['motion'] = to_cpu(results['motion'][i])
batch_output['pred_motion'] = to_cpu(results['pred_motion'][i])
batch_output['motion_length'] = to_cpu(results['motion_length'][i])
batch_output['motion_mask'] = to_cpu(results['motion_mask'][i])
if 'pred_motion_length' in results.keys():
batch_output['pred_motion_length'] = to_cpu(results['pred_motion_length'][i])
else:
batch_output['pred_motion_length'] = to_cpu(results['motion_length'][i])
if 'pred_motion_mask' in results:
batch_output['pred_motion_mask'] = to_cpu(results['pred_motion_mask'][i])
else:
batch_output['pred_motion_mask'] = to_cpu(results['motion_mask'][i])
if 'motion_metas' in results.keys():
motion_metas = results['motion_metas'][i]
if 'text' in motion_metas.keys():
batch_output['text'] = motion_metas['text']
if 'token' in motion_metas.keys():
batch_output['token'] = motion_metas['token']
output.append(batch_output)
return output
@@ -0,0 +1,127 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from .base_architecture import BaseArchitecture
from ..builder import (
ARCHITECTURES,
build_architecture,
build_submodule,
build_loss
)
from ..utils.gaussian_diffusion import (
GaussianDiffusion, get_named_beta_schedule, create_named_schedule_sampler,
ModelMeanType, ModelVarType, LossType, space_timesteps, SpacedDiffusion
)
def build_diffusion(cfg):
beta_scheduler = cfg['beta_scheduler']
diffusion_steps = cfg['diffusion_steps']
betas = get_named_beta_schedule(beta_scheduler, diffusion_steps)
model_mean_type = {
'start_x': ModelMeanType.START_X,
'previous_x': ModelMeanType.PREVIOUS_X,
'epsilon': ModelMeanType.EPSILON
}[cfg['model_mean_type']]
model_var_type = {
'learned': ModelVarType.LEARNED,
'fixed_small': ModelVarType.FIXED_SMALL,
'fixed_large': ModelVarType.FIXED_LARGE,
'learned_range': ModelVarType.LEARNED_RANGE
}[cfg['model_var_type']]
if cfg.get('respace', None) is not None:
diffusion = SpacedDiffusion(
use_timesteps=space_timesteps(diffusion_steps, cfg['respace']),
betas=betas,
model_mean_type=model_mean_type,
model_var_type=model_var_type,
loss_type=LossType.MSE
)
else:
diffusion = GaussianDiffusion(
betas=betas,
model_mean_type=model_mean_type,
model_var_type=model_var_type,
loss_type=LossType.MSE)
return diffusion
@ARCHITECTURES.register_module()
class MotionDiffusion(BaseArchitecture):
def __init__(self,
model=None,
loss_recon=None,
diffusion_train=None,
diffusion_test=None,
init_cfg=None,
inference_type='ddpm',
**kwargs):
super().__init__(init_cfg=init_cfg, **kwargs)
self.model = build_submodule(model)
self.loss_recon = build_loss(loss_recon)
self.diffusion_train = build_diffusion(diffusion_train)
self.diffusion_test = build_diffusion(diffusion_test)
self.sampler = create_named_schedule_sampler('uniform', self.diffusion_train)
self.inference_type = inference_type
def forward(self, **kwargs):
motion, motion_mask = kwargs['motion'].float(), kwargs['motion_mask'].float()
sample_idx = kwargs.get('sample_idx', None)
clip_feat = kwargs.get('clip_feat', None)
B, T = motion.shape[:2]
text = []
for i in range(B):
text.append(kwargs['motion_metas'][i]['text'])
if self.training:
t, _ = self.sampler.sample(B, motion.device)
output = self.diffusion_train.training_losses(
model=self.model,
x_start=motion,
t=t,
model_kwargs={
'motion_mask': motion_mask,
'motion_length': kwargs['motion_length'],
'text': text,
'clip_feat': clip_feat,
'sample_idx': sample_idx}
)
pred, target = output['pred'], output['target']
recon_loss = self.loss_recon(pred, target, reduction_override='none')
recon_loss = (recon_loss.mean(dim=-1) * motion_mask).sum() / motion_mask.sum()
loss = {'recon_loss': recon_loss}
return loss
else:
dim_pose = kwargs['motion'].shape[-1]
model_kwargs = self.model.get_precompute_condition(device=motion.device, text=text, **kwargs)
model_kwargs['motion_mask'] = motion_mask
model_kwargs['sample_idx'] = sample_idx
inference_kwargs = kwargs.get('inference_kwargs', {})
if self.inference_type == 'ddpm':
output = self.diffusion_test.p_sample_loop(
self.model,
(B, T, dim_pose),
clip_denoised=False,
progress=False,
model_kwargs=model_kwargs,
**inference_kwargs
)
else:
output = self.diffusion_test.ddim_sample_loop(
self.model,
(B, T, dim_pose),
clip_denoised=False,
progress=False,
model_kwargs=model_kwargs,
eta=0,
**inference_kwargs
)
if getattr(self.model, "post_process") is not None:
output = self.model.post_process(output)
results = kwargs
results['pred_motion'] = output
results = self.split_results(results)
return results
@@ -0,0 +1,118 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from .base_architecture import BaseArchitecture
from ..builder import (
ARCHITECTURES,
build_architecture,
build_submodule,
build_loss
)
@ARCHITECTURES.register_module()
class PoseVAE(BaseArchitecture):
def __init__(self,
encoder=None,
decoder=None,
loss_recon=None,
kl_div_loss_weight=None,
init_cfg=None,
**kwargs):
super().__init__(init_cfg=init_cfg, **kwargs)
self.encoder = build_submodule(encoder)
self.decoder = build_submodule(decoder)
self.loss_recon = build_loss(loss_recon)
self.kl_div_loss_weight = kl_div_loss_weight
def reparameterize(self, mu, logvar):
std = torch.exp(logvar / 2)
eps = std.data.new(std.size()).normal_()
latent_code = eps.mul(std).add_(mu)
return latent_code
def encode(self, pose):
mu, logvar = self.encoder(pose)
return mu
def forward(self, **kwargs):
motion = kwargs['motion'].float()
B, T = motion.shape[:2]
pose = motion.reshape(B * T, -1)
pose = pose[:, :-4]
mu, logvar = self.encoder(pose)
z = self.reparameterize(mu, logvar)
pred = self.decoder(z)
loss = dict()
recon_loss = self.loss_recon(pred, pose, reduction_override='none')
loss['recon_loss'] = recon_loss
if self.kl_div_loss_weight is not None:
loss_kl = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())
loss['kl_div_loss'] = (loss_kl * self.kl_div_loss_weight)
return loss
@ARCHITECTURES.register_module()
class MotionVAE(BaseArchitecture):
def __init__(self,
encoder=None,
decoder=None,
loss_recon=None,
kl_div_loss_weight=None,
init_cfg=None,
**kwargs):
super().__init__(init_cfg=init_cfg, **kwargs)
self.encoder = build_submodule(encoder)
self.decoder = build_submodule(decoder)
self.loss_recon = build_loss(loss_recon)
self.kl_div_loss_weight = kl_div_loss_weight
def sample(self, std=1, latent_code=None):
if latent_code is not None:
z = latent_code
else:
z = torch.randn(1, 7, self.decoder.latent_dim).cuda() * std
output = self.decoder(z)
if self.use_normalization:
output = output * self.motion_std
output = output + self.motion_mean
return output
def reparameterize(self, mu, logvar):
std = torch.exp(logvar / 2)
eps = std.data.new(std.size()).normal_()
latent_code = eps.mul(std).add_(mu)
return latent_code
def encode(self, motion, motion_mask):
mu, logvar = self.encoder(motion, motion_mask)
return self.reparameterize(mu, logvar)
def decode(self, z, motion_mask):
return self.decoder(z, motion_mask)
def forward(self, **kwargs):
motion, motion_mask = kwargs['motion'].float(), kwargs['motion_mask']
B, T = motion.shape[:2]
mu, logvar = self.encoder(motion, motion_mask)
z = self.reparameterize(mu, logvar)
pred = self.decoder(z, motion_mask)
loss = dict()
recon_loss = self.loss_recon(pred, motion, reduction_override='none')
recon_loss = (recon_loss.mean(dim=-1) * motion_mask).sum() / motion_mask.sum()
loss['recon_loss'] = recon_loss
if self.kl_div_loss_weight is not None:
loss_kl = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())
loss['kl_div_loss'] = (loss_kl * self.kl_div_loss_weight)
return loss
@@ -0,0 +1,6 @@
from .efficient_attention import (
EfficientSelfAttention,
EfficientCrossAttention
)
from .semantics_modulated import SemanticsModulatedAttention
from .base_attention import BaseMixedAttention
@@ -0,0 +1,146 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from ..utils.stylization_block import StylizationBlock
from ..builder import ATTENTIONS
@ATTENTIONS.register_module()
class BaseMixedAttention(nn.Module):
def __init__(self, latent_dim,
text_latent_dim,
num_heads,
dropout,
time_embed_dim):
super().__init__()
self.num_heads = num_heads
self.norm = nn.LayerNorm(latent_dim)
self.text_norm = nn.LayerNorm(text_latent_dim)
self.query = nn.Linear(latent_dim, latent_dim)
self.key_text = nn.Linear(text_latent_dim, latent_dim)
self.value_text = nn.Linear(text_latent_dim, latent_dim)
self.key_motion = nn.Linear(latent_dim, latent_dim)
self.value_motion = nn.Linear(latent_dim, latent_dim)
self.dropout = nn.Dropout(dropout)
self.proj_out = StylizationBlock(latent_dim, time_embed_dim, dropout)
def forward(self, x, xf, emb, src_mask, cond_type, **kwargs):
"""
x: B, T, D
xf: B, N, L
"""
B, T, D = x.shape
N = xf.shape[1] + x.shape[1]
H = self.num_heads
# B, T, D
query = self.query(self.norm(x)).view(B, T, H, -1)
# B, N, D
text_cond_type = ((cond_type % 10) > 0).float().view(B, 1, 1).repeat(1, xf.shape[1], 1)
key = torch.cat((
self.key_text(self.text_norm(xf)),
self.key_motion(self.norm(x))
), dim=1).view(B, N, H, -1)
attention = torch.einsum('bnhl,bmhl->bnmh', query, key)
motion_mask = src_mask.view(B, 1, T, 1)
text_mask = text_cond_type.view(B, 1, -1, 1)
mask = torch.cat((text_mask, motion_mask), dim=2)
attention = attention + (1 - mask) * -1000000
attention = F.softmax(attention, dim=2)
value = torch.cat((
self.value_text(self.text_norm(xf)) * text_cond_type,
self.value_motion(self.norm(x)) * src_mask,
), dim=1).view(B, N, H, -1)
y = torch.einsum('bnmh,bmhl->bnhl', attention, value).reshape(B, T, D)
y = x + self.proj_out(y, emb)
return y
@ATTENTIONS.register_module()
class BaseSelfAttention(nn.Module):
def __init__(self, latent_dim,
num_heads,
dropout,
time_embed_dim):
super().__init__()
self.num_heads = num_heads
self.norm = nn.LayerNorm(latent_dim)
self.query = nn.Linear(latent_dim, latent_dim)
self.key = nn.Linear(latent_dim, latent_dim)
self.value = nn.Linear(latent_dim, latent_dim)
self.dropout = nn.Dropout(dropout)
self.proj_out = StylizationBlock(latent_dim, time_embed_dim, dropout)
def forward(self, x, emb, src_mask, **kwargs):
"""
x: B, T, D
"""
B, T, D = x.shape
H = self.num_heads
# B, T, D
query = self.query(self.norm(x)).view(B, T, H, -1)
# B, N, D
key = self.key(self.norm(x)).view(B, T, H, -1)
attention = torch.einsum('bnhl,bmhl->bnmh', query, key)
mask = src_mask.view(B, 1, T, 1)
attention = attention + (1 - mask) * -1000000
attention = F.softmax(attention, dim=2)
value = (self.value(self.norm(x)) * src_mask).view(B, T, H, -1)
y = torch.einsum('bnmh,bmhl->bnhl', attention, value).reshape(B, T, D)
y = x + self.proj_out(y, emb)
return y
@ATTENTIONS.register_module()
class BaseCrossAttention(nn.Module):
def __init__(self, latent_dim,
text_latent_dim,
num_heads,
dropout,
time_embed_dim):
super().__init__()
self.num_heads = num_heads
self.norm = nn.LayerNorm(latent_dim)
self.text_norm = nn.LayerNorm(text_latent_dim)
self.query = nn.Linear(latent_dim, latent_dim)
self.key = nn.Linear(text_latent_dim, latent_dim)
self.value = nn.Linear(text_latent_dim, latent_dim)
self.dropout = nn.Dropout(dropout)
self.proj_out = StylizationBlock(latent_dim, time_embed_dim, dropout)
def forward(self, x, xf, emb, src_mask, cond_type, **kwargs):
"""
x: B, T, D
xf: B, N, L
"""
B, T, D = x.shape
N = xf.shape[1]
H = self.num_heads
# B, T, D
query = self.query(self.norm(x)).view(B, T, H, -1)
# B, N, D
text_cond_type = ((cond_type % 10) > 0).float().view(B, 1, 1).repeat(1, xf.shape[1], 1)
key = self.key(self.text_norm(xf)).view(B, N, H, -1)
attention = torch.einsum('bnhl,bmhl->bnmh', query, key)
mask = text_cond_type.view(B, 1, -1, 1)
attention = attention + (1 - mask) * -1000000
attention = F.softmax(attention, dim=2)
value = (self.value(self.text_norm(xf)) * text_cond_type).view(B, N, H, -1)
y = torch.einsum('bnmh,bmhl->bnhl', attention, value).reshape(B, T, D)
y = x + self.proj_out(y, emb)
return y
@@ -0,0 +1,87 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from ..utils.stylization_block import StylizationBlock
from ..builder import ATTENTIONS
@ATTENTIONS.register_module()
class EfficientSelfAttention(nn.Module):
def __init__(self, latent_dim, num_heads, dropout, time_embed_dim=None):
super().__init__()
self.num_heads = num_heads
self.norm = nn.LayerNorm(latent_dim)
self.query = nn.Linear(latent_dim, latent_dim)
self.key = nn.Linear(latent_dim, latent_dim)
self.value = nn.Linear(latent_dim, latent_dim)
self.dropout = nn.Dropout(dropout)
self.time_embed_dim = time_embed_dim
if time_embed_dim is not None:
self.proj_out = StylizationBlock(latent_dim, time_embed_dim, dropout)
def forward(self, x, src_mask, emb=None, **kwargs):
"""
x: B, T, D
"""
B, T, D = x.shape
H = self.num_heads
# B, T, D
query = self.query(self.norm(x))
# B, T, D
key = (self.key(self.norm(x)) + (1 - src_mask) * -1000000)
query = F.softmax(query.view(B, T, H, -1), dim=-1)
key = F.softmax(key.view(B, T, H, -1), dim=1)
# B, T, H, HD
value = (self.value(self.norm(x)) * src_mask).view(B, T, H, -1)
# B, H, HD, HD
attention = torch.einsum('bnhd,bnhl->bhdl', key, value)
y = torch.einsum('bnhd,bhdl->bnhl', query, attention).reshape(B, T, D)
if self.time_embed_dim is None:
y = x + y
else:
y = x + self.proj_out(y, emb)
return y
@ATTENTIONS.register_module()
class EfficientCrossAttention(nn.Module):
def __init__(self, latent_dim, text_latent_dim, num_heads, dropout, time_embed_dim):
super().__init__()
self.num_heads = num_heads
self.norm = nn.LayerNorm(latent_dim)
self.text_norm = nn.LayerNorm(text_latent_dim)
self.query = nn.Linear(latent_dim, latent_dim)
self.key = nn.Linear(text_latent_dim, latent_dim)
self.value = nn.Linear(text_latent_dim, latent_dim)
self.dropout = nn.Dropout(dropout)
self.proj_out = StylizationBlock(latent_dim, time_embed_dim, dropout)
def forward(self, x, xf, emb, cond_type=None, **kwargs):
"""
x: B, T, D
xf: B, N, L
"""
B, T, D = x.shape
N = xf.shape[1]
H = self.num_heads
# B, T, D
query = self.query(self.norm(x))
# B, N, D
key = self.key(self.text_norm(xf))
query = F.softmax(query.view(B, T, H, -1), dim=-1)
if cond_type is None:
key = F.softmax(key.view(B, N, H, -1), dim=1)
# B, N, H, HD
value = self.value(self.text_norm(xf)).view(B, N, H, -1)
else:
text_cond_type = ((cond_type % 10) > 0).float().view(B, 1, 1).repeat(1, xf.shape[1], 1)
key = key + (1 - text_cond_type) * -1000000
key = F.softmax(key.view(B, N, H, -1), dim=1)
value = self.value(self.text_norm(xf) * text_cond_type).view(B, N, H, -1)
# B, H, HD, HD
attention = torch.einsum('bnhd,bnhl->bhdl', key, value)
y = torch.einsum('bnhd,bhdl->bnhl', query, attention).reshape(B, T, D)
y = x + self.proj_out(y, emb)
return y
@@ -0,0 +1,82 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from ..utils.stylization_block import StylizationBlock
from ..builder import ATTENTIONS
def zero_module(module):
"""
Zero out the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().zero_()
return module
@ATTENTIONS.register_module()
class SemanticsModulatedAttention(nn.Module):
def __init__(self, latent_dim,
text_latent_dim,
num_heads,
dropout,
time_embed_dim):
super().__init__()
self.num_heads = num_heads
self.norm = nn.LayerNorm(latent_dim)
self.text_norm = nn.LayerNorm(text_latent_dim)
self.query = nn.Linear(latent_dim, latent_dim)
self.key_text = nn.Linear(text_latent_dim, latent_dim)
self.value_text = nn.Linear(text_latent_dim, latent_dim)
self.key_motion = nn.Linear(latent_dim, latent_dim)
self.value_motion = nn.Linear(latent_dim, latent_dim)
self.retr_norm1 = nn.LayerNorm(2 * latent_dim)
self.retr_norm2 = nn.LayerNorm(latent_dim)
self.key_retr = nn.Linear(2 * latent_dim, latent_dim)
self.value_retr = zero_module(nn.Linear(latent_dim, latent_dim))
self.dropout = nn.Dropout(dropout)
self.proj_out = StylizationBlock(latent_dim, time_embed_dim, dropout)
def forward(self, x, xf, emb, src_mask, cond_type, re_dict=None):
"""
x: B, T, D
xf: B, N, L
"""
B, T, D = x.shape
re_motion = re_dict['re_motion']
re_text = re_dict['re_text']
re_mask = re_dict['re_mask']
re_mask = re_mask.reshape(B, -1, 1)
N = xf.shape[1] + x.shape[1] + re_motion.shape[1] * re_motion.shape[2]
H = self.num_heads
# B, T, D
query = self.query(self.norm(x))
# B, N, D
text_cond_type = (cond_type % 10 > 0).float()
retr_cond_type = (cond_type // 10 > 0).float()
re_text = re_text.repeat(1, 1, re_motion.shape[2], 1)
re_feat_key = torch.cat((re_motion, re_text), dim=-1).reshape(B, -1, 2 * D)
key = torch.cat((
self.key_text(self.text_norm(xf)) + (1 - text_cond_type) * -1000000,
self.key_retr(self.retr_norm1(re_feat_key)) + (1 - retr_cond_type) * -1000000 + (1 - re_mask) * -1000000,
self.key_motion(self.norm(x)) + (1 - src_mask) * -1000000
), dim=1)
query = F.softmax(query.view(B, T, H, -1), dim=-1)
key = F.softmax(key.view(B, N, H, -1), dim=1)
# B, N, H, HD
re_feat_value = re_motion.reshape(B, -1, D)
value = torch.cat((
self.value_text(self.text_norm(xf)) * text_cond_type,
self.value_retr(self.retr_norm2(re_feat_value)) * retr_cond_type * re_mask,
self.value_motion(self.norm(x)) * src_mask,
), dim=1).view(B, N, H, -1)
# B, H, HD, HD
attention = torch.einsum('bnhd,bnhl->bhdl', key, value)
y = torch.einsum('bnhd,bhdl->bnhl', query, attention).reshape(B, T, D)
y = x + self.proj_out(y, emb)
return y
@@ -0,0 +1,32 @@
from custom_mmpkg.custom_mmcv.cnn import MODELS as MMCV_MODELS
from custom_mmpkg.custom_mmcv.utils import Registry
def build_from_cfg(cfg, registry, default_args=None):
if cfg is None:
return None
return MMCV_MODELS.build_func(cfg, registry, default_args)
MODELS = Registry('models', parent=MMCV_MODELS, build_func=build_from_cfg)
LOSSES = MODELS
ARCHITECTURES = MODELS
SUBMODULES = MODELS
ATTENTIONS = MODELS
def build_loss(cfg):
"""Build loss."""
return LOSSES.build(cfg)
def build_architecture(cfg):
"""Build framework."""
return ARCHITECTURES.build(cfg)
def build_submodule(cfg):
"""Build submodule."""
return SUBMODULES.build(cfg)
def build_attention(cfg):
"""Build attention."""
return ATTENTIONS.build(cfg)
@@ -0,0 +1,13 @@
from .mse_loss import MSELoss
from .utils import (
convert_to_one_hot,
reduce_loss,
weight_reduce_loss,
weighted_loss,
)
__all__ = [
'convert_to_one_hot', 'reduce_loss', 'weight_reduce_loss', 'weighted_loss',
'MSELoss'
]
@@ -0,0 +1,70 @@
import torch.nn as nn
import torch.nn.functional as F
from ..builder import LOSSES
from .utils import weighted_loss
def gmof(x, sigma):
"""Geman-McClure error function."""
x_squared = x**2
sigma_squared = sigma**2
return (sigma_squared * x_squared) / (sigma_squared + x_squared)
@weighted_loss
def mse_loss(pred, target):
"""Warpper of mse loss."""
return F.mse_loss(pred, target, reduction='none')
@weighted_loss
def mse_loss_with_gmof(pred, target, sigma):
"""Extended MSE Loss with GMOF."""
loss = F.mse_loss(pred, target, reduction='none')
loss = gmof(loss, sigma)
return loss
@LOSSES.register_module()
class MSELoss(nn.Module):
"""MSELoss.
Args:
reduction (str, optional): The method that reduces the loss to a
scalar. Options are "none", "mean" and "sum".
loss_weight (float, optional): The weight of the loss. Defaults to 1.0
"""
def __init__(self, reduction='mean', loss_weight=1.0):
super().__init__()
assert reduction in (None, 'none', 'mean', 'sum')
reduction = 'none' if reduction is None else reduction
self.reduction = reduction
self.loss_weight = loss_weight
def forward(self,
pred,
target,
weight=None,
avg_factor=None,
reduction_override=None):
"""Forward function of loss.
Args:
pred (torch.Tensor): The prediction.
target (torch.Tensor): The learning target of the prediction.
weight (torch.Tensor, optional): Weight of the loss for each
prediction. Defaults to None.
avg_factor (int, optional): Average factor that is used to average
the loss. Defaults to None.
reduction_override (str, optional): The reduction method used to
override the original reduction method of the loss.
Defaults to None.
Returns:
torch.Tensor: The calculated loss
"""
assert reduction_override in (None, 'none', 'mean', 'sum')
reduction = (
reduction_override if reduction_override else self.reduction)
loss = self.loss_weight * mse_loss(
pred, target, weight, reduction=reduction, avg_factor=avg_factor)
return loss
@@ -0,0 +1,109 @@
import functools
import torch
import torch.nn.functional as F
def reduce_loss(loss, reduction):
"""Reduce loss as specified.
Args:
loss (Tensor): Elementwise loss tensor.
reduction (str): Options are "none", "mean" and "sum".
Return:
Tensor: Reduced loss tensor.
"""
reduction_enum = F._Reduction.get_enum(reduction)
# none: 0, elementwise_mean:1, sum: 2
if reduction_enum == 0:
return loss
elif reduction_enum == 1:
return loss.mean()
elif reduction_enum == 2:
return loss.sum()
def weight_reduce_loss(loss, weight=None, reduction='mean', avg_factor=None):
"""Apply element-wise weight and reduce loss.
Args:
loss (Tensor): Element-wise loss.
weight (Tensor): Element-wise weights.
reduction (str): Same as built-in losses of PyTorch.
avg_factor (float): Average factor when computing the mean of losses.
Returns:
Tensor: Processed loss values.
"""
# if weight is specified, apply element-wise weight
if weight is not None:
loss = loss * weight
# if avg_factor is not specified, just reduce the loss
if avg_factor is None:
loss = reduce_loss(loss, reduction)
else:
# if reduction is mean, then average the loss by avg_factor
if reduction == 'mean':
loss = loss.sum() / avg_factor
# if reduction is 'none', then do nothing, otherwise raise an error
elif reduction != 'none':
raise ValueError('avg_factor can not be used with reduction="sum"')
return loss
def weighted_loss(loss_func):
"""Create a weighted version of a given loss function.
To use this decorator, the loss function must have the signature like
`loss_func(pred, target, **kwargs)`. The function only needs to compute
element-wise loss without any reduction. This decorator will add weight
and reduction arguments to the function. The decorated function will have
the signature like `loss_func(pred, target, weight=None, reduction='mean',
avg_factor=None, **kwargs)`.
:Example:
>>> import torch
>>> @weighted_loss
>>> def l1_loss(pred, target):
>>> return (pred - target).abs()
>>> pred = torch.Tensor([0, 2, 3])
>>> target = torch.Tensor([1, 1, 1])
>>> weight = torch.Tensor([1, 0, 1])
>>> l1_loss(pred, target)
tensor(1.3333)
>>> l1_loss(pred, target, weight)
tensor(1.)
>>> l1_loss(pred, target, reduction='none')
tensor([1., 1., 2.])
>>> l1_loss(pred, target, weight, avg_factor=2)
tensor(1.5000)
"""
@functools.wraps(loss_func)
def wrapper(pred,
target,
weight=None,
reduction='mean',
avg_factor=None,
**kwargs):
# get element-wise loss
loss = loss_func(pred, target, **kwargs)
loss = weight_reduce_loss(loss, weight, reduction, avg_factor)
return loss
return wrapper
def convert_to_one_hot(targets: torch.Tensor, classes) -> torch.Tensor:
"""This function converts target class indices to one-hot vectors, given
the number of classes.
Args:
targets (Tensor): The ground truth label of the prediction
with shape (N, 1)
classes (int): the number of classes.
Returns:
Tensor: Processed loss values.
"""
assert (torch.max(targets).item() <
classes), 'Class Index must be less than number of classes'
one_hot_targets = torch.zeros((targets.shape[0], classes),
dtype=torch.long,
device=targets.device)
one_hot_targets.scatter_(1, targets.long(), 1)
return one_hot_targets
@@ -0,0 +1 @@
from .t2m_bigru import T2MMotionEncoder, T2MTextEncoder
@@ -0,0 +1,260 @@
import torch
import torch.nn as nn
import numpy as np
import time
import math
from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence
import torch.nn.functional as F
from ..builder import SUBMODULES
from motiondiff_modules.mogen.models.utils.word_vectorizer import WordVectorizer
def init_weight(m):
if isinstance(m, nn.Conv1d) or isinstance(m, nn.Linear) or isinstance(m, nn.ConvTranspose1d):
nn.init.xavier_normal_(m.weight)
# m.bias.data.fill_(0.01)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def reparameterize(mu, logvar):
s_var = logvar.mul(0.5).exp_()
eps = s_var.data.new(s_var.size()).normal_()
return eps.mul(s_var).add_(mu)
# batch_size, dimension and position
# output: (batch_size, dim)
def positional_encoding(batch_size, dim, pos):
assert batch_size == pos.shape[0]
positions_enc = np.array([
[pos[j] / np.power(10000, (i-i%2)/dim) for i in range(dim)]
for j in range(batch_size)
], dtype=np.float32)
positions_enc[:, 0::2] = np.sin(positions_enc[:, 0::2])
positions_enc[:, 1::2] = np.cos(positions_enc[:, 1::2])
return torch.from_numpy(positions_enc).float()
def get_padding_mask(batch_size, seq_len, cap_lens):
cap_lens = cap_lens.data.tolist()
mask_2d = torch.ones((batch_size, seq_len, seq_len), dtype=torch.float32)
for i, cap_len in enumerate(cap_lens):
mask_2d[i, :, :cap_len] = 0
return mask_2d.bool(), 1 - mask_2d[:, :, 0].clone()
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=300):
super(PositionalEncoding, self).__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
# pe = pe.unsqueeze(0).transpose(0, 1)
self.register_buffer('pe', pe)
def forward(self, pos):
return self.pe[pos]
@SUBMODULES.register_module()
class T2MMotionEncoder(nn.Module):
def __init__(self,
input_size,
movement_hidden_size,
movement_latent_size,
motion_hidden_size,
motion_latent_size):
super().__init__()
self.movement_encoder = MovementConvEncoder(
input_size=input_size-4,
hidden_size=movement_hidden_size,
output_size=movement_latent_size)
self.motion_encoder = MotionEncoderBiGRUCo(
input_size=movement_latent_size,
hidden_size=motion_hidden_size,
output_size=motion_latent_size
)
def load_pretrained(self, ckpt_path):
checkpoint = torch.load(ckpt_path, map_location='cpu')
self.movement_encoder.load_state_dict(checkpoint['movement_encoder'])
self.motion_encoder.load_state_dict(checkpoint['motion_encoder'])
def forward(self, motion, motion_length, motion_mask):
motion = motion.detach().float()
sort_idx = np.argsort(motion_length.data.tolist())[::-1].copy()
rank_idx = np.empty_like(sort_idx)
rank_idx[sort_idx] = np.arange(len(motion_length))
motion = motion[sort_idx]
motion_length = motion_length[sort_idx]
movements = self.movement_encoder(motion[..., :-4]).detach()
m_lens = motion_length // 4
motion_embedding = self.motion_encoder(movements, m_lens)
motion_embedding_ordered = motion_embedding[rank_idx]
return motion_embedding_ordered
@SUBMODULES.register_module()
class T2MTextEncoder(nn.Module):
def __init__(self,
word_size,
pos_size,
hidden_size,
output_size,
max_text_len):
super().__init__()
self.text_encoder = TextEncoderBiGRUCo(
word_size=word_size,
pos_size=pos_size,
hidden_size=hidden_size,
output_size=output_size,
)
self.w_vectorizer = WordVectorizer('./data/glove', 'our_vab')
self.max_text_len = max_text_len
def load_pretrained(self, ckpt_path):
checkpoint = torch.load(ckpt_path, map_location='cpu')
self.text_encoder.load_state_dict(checkpoint['text_encoder'])
def forward(self, text, token, device):
B = len(text)
pos_one_hot = []
word_emb = []
sent_len = []
for i in range(B):
tokens = token[i].split(" ")
if len(tokens) < self.max_text_len:
tokens = ['sos/OTHER'] + tokens + ['eos/OTHER']
batch_sent_len = len(tokens)
tokens = tokens + ['unk/OTHER'] * (self.max_text_len + 2 - batch_sent_len)
else:
tokens = tokens[: self.max_text_len]
tokens = ['sos/OTHER'] + tokens + ['eos/OTHER']
batch_sent_len = len(tokens)
sent_len.append(batch_sent_len)
batch_word_emb = []
batch_pos_one_hot = []
for cur_token in tokens:
cur_word_emb, cur_pos_one_hot = self.w_vectorizer[cur_token]
cur_word_emb = torch.from_numpy(cur_word_emb).float()
cur_pos_one_hot = torch.from_numpy(cur_pos_one_hot).float()
batch_word_emb.append(cur_word_emb)
batch_pos_one_hot.append(cur_pos_one_hot)
batch_word_emb = torch.stack(batch_word_emb, dim=0)
batch_pos_one_hot = torch.stack(batch_pos_one_hot, dim=0)
word_emb.append(batch_word_emb)
pos_one_hot.append(batch_pos_one_hot)
word_emb = torch.stack(word_emb, dim=0).to(device)
pos_one_hot = torch.stack(pos_one_hot, dim=0).to(device)
sent_len = torch.tensor(sent_len, dtype=torch.long).to(device)
text_embedding = self.text_encoder(word_emb, pos_one_hot, sent_len)
return text_embedding
class TextEncoderBiGRUCo(nn.Module):
def __init__(self, word_size, pos_size, hidden_size, output_size):
super(TextEncoderBiGRUCo, self).__init__()
self.pos_emb = nn.Linear(pos_size, word_size)
self.input_emb = nn.Linear(word_size, hidden_size)
self.gru = nn.GRU(hidden_size, hidden_size, batch_first=True, bidirectional=True)
self.output_net = nn.Sequential(
nn.Linear(hidden_size * 2, hidden_size),
nn.LayerNorm(hidden_size),
nn.LeakyReLU(0.2, inplace=True),
nn.Linear(hidden_size, output_size)
)
self.input_emb.apply(init_weight)
self.pos_emb.apply(init_weight)
self.output_net.apply(init_weight)
# self.linear2.apply(init_weight)
# self.batch_size = batch_size
self.hidden_size = hidden_size
self.hidden = nn.Parameter(torch.randn((2, 1, self.hidden_size), requires_grad=True))
# input(batch_size, seq_len, dim)
def forward(self, word_embs, pos_onehot, cap_lens):
num_samples = word_embs.shape[0]
pos_embs = self.pos_emb(pos_onehot)
inputs = word_embs + pos_embs
input_embs = self.input_emb(inputs)
hidden = self.hidden.repeat(1, num_samples, 1)
cap_lens = cap_lens.data.tolist()
emb = pack_padded_sequence(input_embs, cap_lens, batch_first=True, enforce_sorted=False)
gru_seq, gru_last = self.gru(emb, hidden)
gru_last = torch.cat([gru_last[0], gru_last[1]], dim=-1)
return self.output_net(gru_last)
class MovementConvEncoder(nn.Module):
def __init__(self, input_size, hidden_size, output_size):
super(MovementConvEncoder, self).__init__()
self.main = nn.Sequential(
nn.Conv1d(input_size, hidden_size, 4, 2, 1),
nn.Dropout(0.2, inplace=True),
nn.LeakyReLU(0.2, inplace=True),
nn.Conv1d(hidden_size, output_size, 4, 2, 1),
nn.Dropout(0.2, inplace=True),
nn.LeakyReLU(0.2, inplace=True),
)
self.out_net = nn.Linear(output_size, output_size)
self.main.apply(init_weight)
self.out_net.apply(init_weight)
def forward(self, inputs):
inputs = inputs.permute(0, 2, 1)
outputs = self.main(inputs).permute(0, 2, 1)
# print(outputs.shape)
return self.out_net(outputs)
class MotionEncoderBiGRUCo(nn.Module):
def __init__(self, input_size, hidden_size, output_size):
super(MotionEncoderBiGRUCo, self).__init__()
self.input_emb = nn.Linear(input_size, hidden_size)
self.gru = nn.GRU(hidden_size, hidden_size, batch_first=True, bidirectional=True)
self.output_net = nn.Sequential(
nn.Linear(hidden_size*2, hidden_size),
nn.LayerNorm(hidden_size),
nn.LeakyReLU(0.2, inplace=True),
nn.Linear(hidden_size, output_size)
)
self.input_emb.apply(init_weight)
self.output_net.apply(init_weight)
self.hidden_size = hidden_size
self.hidden = nn.Parameter(torch.randn((2, 1, self.hidden_size), requires_grad=True))
# input(batch_size, seq_len, dim)
def forward(self, inputs, m_lens):
num_samples = inputs.shape[0]
input_embs = self.input_emb(inputs)
hidden = self.hidden.repeat(1, num_samples, 1)
cap_lens = m_lens.data.tolist()
emb = pack_padded_sequence(input_embs, cap_lens, batch_first=True)
gru_seq, gru_last = self.gru(emb, hidden)
gru_last = torch.cat([gru_last[0], gru_last[1]], dim=-1)
return self.output_net(gru_last)
@@ -0,0 +1,5 @@
from .actor import ACTOREncoder, ACTORDecoder
from .motiondiffuse import MotionDiffuseTransformer
from .remodiffuse import ReMoDiffuseTransformer
from .mdm import MDMTransformer
from .position_encoding import SinusoidalPositionalEncoding, LearnedPositionalEncoding
@@ -0,0 +1,189 @@
from cv2 import norm
import torch
from torch import layer_norm, nn
from custom_mmpkg.custom_mmcv.runner import BaseModule
import numpy as np
from ..builder import SUBMODULES
from .position_encoding import SinusoidalPositionalEncoding, LearnedPositionalEncoding
import math
@SUBMODULES.register_module()
class ACTOREncoder(BaseModule):
def __init__(self,
max_seq_len=16,
njoints=None,
nfeats=None,
input_feats=None,
latent_dim=256,
output_dim=256,
condition_dim=None,
num_heads=4,
ff_size=1024,
num_layers=8,
activation='gelu',
dropout=0.1,
use_condition=False,
num_class=None,
use_final_proj=False,
output_var=False,
pos_embedding='sinusoidal',
init_cfg=None):
super().__init__(init_cfg=init_cfg)
self.njoints = njoints
self.nfeats = nfeats
if input_feats is None:
assert self.njoints is not None and self.nfeats is not None
self.input_feats = njoints * nfeats
else:
self.input_feats = input_feats
self.max_seq_len = max_seq_len
self.latent_dim = latent_dim
self.condition_dim = condition_dim
self.use_condition = use_condition
self.num_class = num_class
self.use_final_proj = use_final_proj
self.output_var = output_var
self.skelEmbedding = nn.Linear(self.input_feats, self.latent_dim)
if self.use_condition:
if num_class is None:
self.mu_layer = build_MLP(self.condition_dim, self.latent_dim)
if self.output_var:
self.sigma_layer = build_MLP(self.condition_dim, self.latent_dim)
else:
self.mu_layer = nn.Parameter(torch.randn(num_class, self.latent_dim))
if self.output_var:
self.sigma_layer = nn.Parameter(torch.randn(num_class, self.latent_dim))
else:
if self.output_var:
self.query = nn.Parameter(torch.randn(2, self.latent_dim))
else:
self.query = nn.Parameter(torch.randn(1, self.latent_dim))
if pos_embedding == 'sinusoidal':
self.pos_encoder = SinusoidalPositionalEncoding(latent_dim, dropout)
else:
self.pos_encoder = LearnedPositionalEncoding(latent_dim, dropout, max_len=max_seq_len + 2)
seqTransEncoderLayer = nn.TransformerEncoderLayer(
d_model=self.latent_dim,
nhead=num_heads,
dim_feedforward=ff_size,
dropout=dropout,
activation=activation)
self.seqTransEncoder = nn.TransformerEncoder(
seqTransEncoderLayer,
num_layers=num_layers)
def forward(self, motion, motion_mask=None, condition=None):
B, T = motion.shape[:2]
motion = motion.view(B, T, -1)
feature = self.skelEmbedding(motion)
if self.use_condition:
if self.output_var:
if self.num_class is None:
sigma_query = self.sigma_layer(condition).view(B, 1, -1)
else:
sigma_query = self.sigma_layer[condition.long()].view(B, 1, -1)
feature = torch.cat((sigma_query, feature), dim=1)
if self.num_class is None:
mu_query = self.mu_layer(condition).view(B, 1, -1)
else:
mu_query = self.mu_layer[condition.long()].view(B, 1, -1)
feature = torch.cat((mu_query, feature), dim=1)
else:
query = self.query.view(1, -1, self.latent_dim).repeat(B, 1, 1)
feature = torch.cat((query, feature), dim=1)
if self.output_var:
motion_mask = torch.cat((torch.zeros(B, 2).to(motion.device), 1 - motion_mask), dim=1).bool()
else:
motion_mask = torch.cat((torch.zeros(B, 1).to(motion.device), 1 - motion_mask), dim=1).bool()
feature = feature.permute(1, 0, 2).contiguous()
feature = self.pos_encoder(feature)
feature = self.seqTransEncoder(feature, src_key_padding_mask=motion_mask)
if self.use_final_proj:
mu = self.final_mu(feature[0])
if self.output_var:
sigma = self.final_sigma(feature[1])
return mu, sigma
return mu
else:
if self.output_var:
return feature[0], feature[1]
else:
return feature[0]
@SUBMODULES.register_module()
class ACTORDecoder(BaseModule):
def __init__(self,
max_seq_len=16,
njoints=None,
nfeats=None,
input_feats=None,
input_dim=256,
latent_dim=256,
condition_dim=None,
num_heads=4,
ff_size=1024,
num_layers=8,
activation='gelu',
dropout=0.1,
use_condition=False,
num_class=None,
pos_embedding='sinusoidal',
init_cfg=None):
super().__init__(init_cfg=init_cfg)
if input_dim != latent_dim:
self.linear = nn.Linear(input_dim, latent_dim)
else:
self.linear = nn.Identity()
self.njoints = njoints
self.nfeats = nfeats
if input_feats is None:
assert self.njoints is not None and self.nfeats is not None
self.input_feats = njoints * nfeats
else:
self.input_feats = input_feats
self.max_seq_len = max_seq_len
self.input_dim = input_dim
self.latent_dim = latent_dim
self.condition_dim = condition_dim
self.use_condition = use_condition
self.num_class = num_class
if self.use_condition:
if num_class is None:
self.condition_bias = build_MLP(condition_dim, latent_dim)
else:
self.condition_bias = nn.Parameter(torch.randn(num_class, latent_dim))
if pos_embedding == 'sinusoidal':
self.pos_encoder = SinusoidalPositionalEncoding(latent_dim, dropout)
else:
self.pos_encoder = LearnedPositionalEncoding(latent_dim, dropout, max_len=max_seq_len)
seqTransDecoderLayer = nn.TransformerDecoderLayer(
d_model=self.latent_dim,
nhead=num_heads,
dim_feedforward=ff_size,
dropout=dropout,
activation=activation)
self.seqTransDecoder = nn.TransformerDecoder(
seqTransDecoderLayer,
num_layers=num_layers)
self.final = nn.Linear(self.latent_dim, self.input_feats)
def forward(self, input, motion_mask=None, condition=None):
B = input.shape[0]
T = self.max_seq_len
input = self.linear(input)
if self.use_condition:
if self.num_class is None:
condition = self.condition_bias(condition)
else:
condition = self.condition_bias[condition.long()].squeeze(1)
input = input + condition
query = self.pos_encoder.pe[:T, :].view(T, 1, -1).repeat(1, B, 1)
input = input.view(1, B, -1)
feature = self.seqTransDecoder(tgt=query, memory=input, tgt_key_padding_mask=(1 - motion_mask).bool())
pose = self.final(feature).permute(1, 0, 2).contiguous()
return pose
@@ -0,0 +1,251 @@
from abc import ABCMeta, abstractmethod
from cv2 import norm
import torch
from torch import layer_norm, nn
import torch.nn.functional as F
from custom_mmpkg.custom_mmcv.runner import BaseModule
import numpy as np
from ..builder import SUBMODULES, build_attention
from .position_encoding import SinusoidalPositionalEncoding, LearnedPositionalEncoding
from ..utils.stylization_block import StylizationBlock
import math
import clip
def timestep_embedding(timesteps, dim, max_period=10000):
"""
Create sinusoidal timestep embeddings.
:param timesteps: a 1-D Tensor of N indices, one per batch element.
These may be fractional.
:param dim: the dimension of the output.
:param max_period: controls the minimum frequency of the embeddings.
:return: an [N x dim] Tensor of positional embeddings.
"""
half = dim // 2
freqs = torch.exp(
-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
).to(device=timesteps.device)
args = timesteps[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
def set_requires_grad(nets, requires_grad=False):
"""Set requies_grad for all the networks.
Args:
nets (nn.Module | list[nn.Module]): A list of networks or a single
network.
requires_grad (bool): Whether the networks require gradients or not
"""
if not isinstance(nets, list):
nets = [nets]
for net in nets:
if net is not None:
for param in net.parameters():
param.requires_grad = requires_grad
def zero_module(module):
"""
Zero out the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().zero_()
return module
class FFN(nn.Module):
def __init__(self, latent_dim, ffn_dim, dropout, time_embed_dim):
super().__init__()
self.linear1 = nn.Linear(latent_dim, ffn_dim)
self.linear2 = zero_module(nn.Linear(ffn_dim, latent_dim))
self.activation = nn.GELU()
self.dropout = nn.Dropout(dropout)
self.proj_out = StylizationBlock(latent_dim, time_embed_dim, dropout)
def forward(self, x, emb, **kwargs):
y = self.linear2(self.dropout(self.activation(self.linear1(x))))
y = x + self.proj_out(y, emb)
return y
class DecoderLayer(nn.Module):
def __init__(self,
sa_block_cfg=None,
ca_block_cfg=None,
ffn_cfg=None):
super().__init__()
self.sa_block = build_attention(sa_block_cfg)
self.ca_block = build_attention(ca_block_cfg)
self.ffn = FFN(**ffn_cfg)
def forward(self, **kwargs):
if self.sa_block is not None:
x = self.sa_block(**kwargs)
kwargs.update({'x': x})
if self.ca_block is not None:
x = self.ca_block(**kwargs)
kwargs.update({'x': x})
if self.ffn is not None:
x = self.ffn(**kwargs)
return x
class DiffusionTransformer(BaseModule, metaclass=ABCMeta):
def __init__(self,
input_feats,
max_seq_len=240,
latent_dim=512,
time_embed_dim=2048,
num_layers=8,
sa_block_cfg=None,
ca_block_cfg=None,
ffn_cfg=None,
text_encoder=None,
use_cache_for_text=False,
init_cfg=None):
super().__init__(init_cfg=init_cfg)
self.input_feats = input_feats
self.max_seq_len = max_seq_len
self.latent_dim = latent_dim
self.num_layers = num_layers
self.time_embed_dim = time_embed_dim
self.sequence_embedding = nn.Parameter(torch.randn(max_seq_len, latent_dim))
self.use_cache_for_text = use_cache_for_text
if use_cache_for_text:
self.text_cache = {}
self.build_text_encoder(text_encoder)
# Input Embedding
self.joint_embed = nn.Linear(self.input_feats, self.latent_dim)
self.time_embed = nn.Sequential(
nn.Linear(self.latent_dim, self.time_embed_dim),
nn.SiLU(),
nn.Linear(self.time_embed_dim, self.time_embed_dim),
)
self.build_temporal_blocks(sa_block_cfg, ca_block_cfg, ffn_cfg)
# Output Module
self.out = zero_module(nn.Linear(self.latent_dim, self.input_feats))
def build_temporal_blocks(self, sa_block_cfg, ca_block_cfg, ffn_cfg):
self.temporal_decoder_blocks = nn.ModuleList()
for i in range(self.num_layers):
self.temporal_decoder_blocks.append(
DecoderLayer(
sa_block_cfg=sa_block_cfg,
ca_block_cfg=ca_block_cfg,
ffn_cfg=ffn_cfg
)
)
def build_text_encoder(self, text_encoder):
text_latent_dim = text_encoder['latent_dim']
num_text_layers = text_encoder.get('num_layers', 0)
text_ff_size = text_encoder.get('ff_size', 2048)
pretrained_model = text_encoder['pretrained_model']
text_num_heads = text_encoder.get('num_heads', 4)
dropout = text_encoder.get('dropout', 0)
activation = text_encoder.get('activation', 'gelu')
self.use_text_proj = text_encoder.get('use_text_proj', False)
if pretrained_model == 'clip':
self.clip, _ = clip.load('ViT-B/32', "cpu")
set_requires_grad(self.clip, False)
if text_latent_dim != 512:
self.text_pre_proj = nn.Linear(512, text_latent_dim)
else:
self.text_pre_proj = nn.Identity()
else:
raise NotImplementedError()
if num_text_layers > 0:
self.use_text_finetune = True
textTransEncoderLayer = nn.TransformerEncoderLayer(
d_model=text_latent_dim,
nhead=text_num_heads,
dim_feedforward=text_ff_size,
dropout=dropout,
activation=activation)
self.textTransEncoder = nn.TransformerEncoder(
textTransEncoderLayer,
num_layers=num_text_layers)
else:
self.use_text_finetune = False
self.text_ln = nn.LayerNorm(text_latent_dim)
if self.use_text_proj:
self.text_proj = nn.Sequential(
nn.Linear(text_latent_dim, self.time_embed_dim)
)
def encode_text(self, text, clip_feat, device):
B = len(text)
text = clip.tokenize(text, truncate=True).to(device)
if clip_feat is None:
with torch.no_grad():
x = self.clip.token_embedding(text).type(self.clip.dtype) # [batch_size, n_ctx, d_model]
x = x + self.clip.positional_embedding.type(self.clip.dtype)
x = x.permute(1, 0, 2) # NLD -> LND
x = self.clip.transformer(x)
x = self.clip.ln_final(x).type(self.clip.dtype)
else:
x = clip_feat.type(self.clip.dtype).to(device).permute(1, 0, 2)
# T, B, D
x = self.text_pre_proj(x)
xf_out = self.textTransEncoder(x)
xf_out = self.text_ln(xf_out)
if self.use_text_proj:
xf_proj = self.text_proj(xf_out[text.argmax(dim=-1), torch.arange(xf_out.shape[1])])
# B, T, D
xf_out = xf_out.permute(1, 0, 2)
return xf_proj, xf_out
else:
xf_out = xf_out.permute(1, 0, 2)
return xf_out
@abstractmethod
def get_precompute_condition(self, **kwargs):
pass
@abstractmethod
def forward_train(self, h, src_mask, emb, **kwargs):
pass
@abstractmethod
def forward_test(self, h, src_mask, emb, **kwargs):
pass
def forward(self, motion, timesteps, motion_mask=None, **kwargs):
"""
motion: B, T, D
"""
B, T = motion.shape[0], motion.shape[1]
conditions = self.get_precompute_condition(device=motion.device, **kwargs)
if len(motion_mask.shape) == 2:
src_mask = motion_mask.clone().unsqueeze(-1)
else:
src_mask = motion_mask.clone()
if self.use_text_proj:
emb = self.time_embed(timestep_embedding(timesteps, self.latent_dim)) + conditions['xf_proj']
else:
emb = self.time_embed(timestep_embedding(timesteps, self.latent_dim))
# B, T, latent_dim
h = self.joint_embed(motion)
h = h + self.sequence_embedding.unsqueeze(0)[:, :T, :]
if self.training:
return self.forward_train(h=h, src_mask=src_mask, emb=emb, timesteps=timesteps, **conditions)
else:
return self.forward_test(h=h, src_mask=src_mask, emb=emb, timesteps=timesteps, **conditions)
@@ -0,0 +1,212 @@
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import clip
from ..builder import SUBMODULES
def convert_weights(model: nn.Module):
"""Convert applicable model parameters to fp32"""
def _convert_weights_to_fp32(l):
if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Linear)):
l.weight.data = l.weight.data.float()
if l.bias is not None:
l.bias.data = l.bias.data.float()
if isinstance(l, nn.MultiheadAttention):
for attr in [*[f"{s}_proj_weight" for s in ["in", "q", "k", "v"]], "in_proj_bias", "bias_k", "bias_v"]:
tensor = getattr(l, attr)
if tensor is not None:
tensor.data = tensor.data.float()
for name in ["text_projection", "proj"]:
if hasattr(l, name):
attr = getattr(l, name)
if attr is not None:
attr.data = attr.data.float()
model.apply(_convert_weights_to_fp32)
@SUBMODULES.register_module()
class MDMTransformer(nn.Module):
def __init__(self,
input_feats=263,
latent_dim=256,
ff_size=1024,
num_layers=8,
num_heads=4,
dropout=0.1,
activation="gelu",
clip_dim=512,
clip_version=None,
guide_scale=1.0,
cond_mask_prob=0.1,
use_official_ckpt=False,
**kwargs):
super().__init__()
self.latent_dim = latent_dim
self.ff_size = ff_size
self.num_layers = num_layers
self.num_heads = num_heads
self.dropout = dropout
self.activation = activation
self.clip_dim = clip_dim
self.input_feats = input_feats
self.guide_scale = guide_scale
self.use_official_ckpt = use_official_ckpt
self.cond_mask_prob = cond_mask_prob
self.poseEmbedding = nn.Linear(self.input_feats, self.latent_dim)
self.sequence_pos_encoder = PositionalEncoding(self.latent_dim, self.dropout)
seqTransEncoderLayer = nn.TransformerEncoderLayer(d_model=self.latent_dim,
nhead=self.num_heads,
dim_feedforward=self.ff_size,
dropout=self.dropout,
activation=self.activation)
self.seqTransEncoder = nn.TransformerEncoder(seqTransEncoderLayer,
num_layers=self.num_layers)
self.embed_timestep = TimestepEmbedder(self.latent_dim, self.sequence_pos_encoder)
self.embed_text = nn.Linear(self.clip_dim, self.latent_dim)
self.clip_version = clip_version
self.clip_model = self.load_and_freeze_clip(clip_version)
self.poseFinal = nn.Linear(self.latent_dim, self.input_feats)
def load_and_freeze_clip(self, clip_version):
clip_model, clip_preprocess = clip.load(clip_version, device='cpu',
jit=False) # Must set jit=False for training
clip.model.convert_weights(
clip_model) # Actually this line is unnecessary since clip by default already on float16
clip_model.eval()
for p in clip_model.parameters():
p.requires_grad = False
return clip_model
def mask_cond(self, cond, force_mask=False):
bs, d = cond.shape
if force_mask:
return torch.zeros_like(cond)
elif self.training and self.cond_mask_prob > 0.:
mask = torch.bernoulli(torch.ones(bs, device=cond.device) * self.cond_mask_prob).view(bs, 1) # 1-> use null_cond, 0-> use real cond
return cond * (1. - mask)
else:
return cond
def encode_text(self, raw_text):
# raw_text - list (batch_size length) of strings with input text prompts
device = next(self.parameters()).device
max_text_len = 20
if max_text_len is not None:
default_context_length = 77
context_length = max_text_len + 2 # start_token + 20 + end_token
assert context_length < default_context_length
texts = clip.tokenize(raw_text, context_length=context_length, truncate=True).to(device)
zero_pad = torch.zeros([texts.shape[0], default_context_length-context_length], dtype=texts.dtype, device=texts.device)
texts = torch.cat([texts, zero_pad], dim=1)
return self.clip_model.encode_text(texts).float()
def get_precompute_condition(self, text, device=None, **kwargs):
if not self.training and device == torch.device('cpu'):
convert_weights(self.clip_model)
text_feat = self.encode_text(text)
return {'text_feat': text_feat}
def post_process(self, motion):
assert len(motion.shape) == 3
if self.use_official_ckpt:
motion[:, :, :4] = motion[:, :, :4] * 25
return motion
def forward(self, motion, timesteps, text_feat=None, **kwargs):
"""
motion: B, T, D
timesteps: [batch_size] (int)
"""
B, T, D = motion.shape
device = motion.device
if text_feat is None:
enc_text = self.get_precompute_condition(**kwargs)['text_feat']
else:
enc_text = text_feat
if self.training:
# T, B, D
motion = self.poseEmbedding(motion).permute(1, 0, 2)
emb = self.embed_timestep(timesteps) # [1, bs, d]
emb += self.embed_text(self.mask_cond(enc_text, force_mask=False))
xseq = self.sequence_pos_encoder(torch.cat((emb, motion), axis=0))
output = self.seqTransEncoder(xseq)[1:]
# B, T, D
output = self.poseFinal(output).permute(1, 0, 2)
return output
else:
# T, B, D
motion = self.poseEmbedding(motion).permute(1, 0, 2)
emb = self.embed_timestep(timesteps) # [1, bs, d]
emb_uncond = emb + self.embed_text(self.mask_cond(enc_text, force_mask=True))
emb_text = emb + self.embed_text(self.mask_cond(enc_text, force_mask=False))
xseq = self.sequence_pos_encoder(torch.cat((emb_uncond, motion), axis=0))
xseq_text = self.sequence_pos_encoder(torch.cat((emb_text, motion), axis=0))
output = self.seqTransEncoder(xseq)[1:]
output_text = self.seqTransEncoder(xseq_text)[1:]
# B, T, D
output = self.poseFinal(output).permute(1, 0, 2)
output_text = self.poseFinal(output_text).permute(1, 0, 2)
scale = self.guide_scale
output = output + scale * (output_text - output)
return output
class PositionalEncoding(nn.Module):
def __init__(self, d_model, dropout=0.1, max_len=5000):
super(PositionalEncoding, self).__init__()
self.dropout = nn.Dropout(p=dropout)
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-np.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0).transpose(0, 1)
self.register_buffer('pe', pe)
def forward(self, x):
# not used in the final model
x = x + self.pe[:x.shape[0], :]
return self.dropout(x)
class TimestepEmbedder(nn.Module):
def __init__(self, latent_dim, sequence_pos_encoder):
super().__init__()
self.latent_dim = latent_dim
self.sequence_pos_encoder = sequence_pos_encoder
time_embed_dim = self.latent_dim
self.time_embed = nn.Sequential(
nn.Linear(self.latent_dim, time_embed_dim),
nn.SiLU(),
nn.Linear(time_embed_dim, time_embed_dim),
)
def forward(self, timesteps):
return self.time_embed(self.sequence_pos_encoder.pe[timesteps]).permute(1, 0, 2)
@@ -0,0 +1,41 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from ..builder import SUBMODULES
from .diffusion_transformer import DiffusionTransformer
@SUBMODULES.register_module()
class MotionDiffuseTransformer(DiffusionTransformer):
def __init__(self, **kwargs):
super().__init__(**kwargs)
def get_precompute_condition(self,
text=None,
xf_proj=None,
xf_out=None,
device=None,
clip_feat=None,
**kwargs):
if xf_proj is None or xf_out is None:
xf_proj, xf_out = self.encode_text(text, clip_feat, device)
return {'xf_proj': xf_proj, 'xf_out': xf_out}
def post_process(self, motion):
return motion
def forward_train(self, h=None, src_mask=None, emb=None, xf_out=None, **kwargs):
B, T = h.shape[0], h.shape[1]
for module in self.temporal_decoder_blocks:
h = module(x=h, xf=xf_out, emb=emb, src_mask=src_mask)
output = self.out(h).view(B, T, -1).contiguous()
return output
def forward_test(self, h=None, src_mask=None, emb=None, xf_out=None, **kwargs):
B, T = h.shape[0], h.shape[1]
for module in self.temporal_decoder_blocks:
h = module(x=h, xf=xf_out, emb=emb, src_mask=src_mask)
output = self.out(h).view(B, T, -1).contiguous()
return output
@@ -0,0 +1,35 @@
import torch
import torch.nn as nn
import numpy as np
class SinusoidalPositionalEncoding(nn.Module):
def __init__(self, d_model, dropout=0.1, max_len=5000):
super(SinusoidalPositionalEncoding, self).__init__()
self.dropout = nn.Dropout(p=dropout)
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.arange(0, d_model, 2).float()
div_term = div_term * (-np.log(10000.0) / d_model)
div_term = torch.exp(div_term)
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0).transpose(0, 1)
# T, 1, D
self.register_buffer('pe', pe)
def forward(self, x):
x = x + self.pe[:x.shape[0]]
return self.dropout(x)
class LearnedPositionalEncoding(nn.Module):
def __init__(self, d_model, dropout=0.1, max_len=5000):
super(LearnedPositionalEncoding, self).__init__()
self.dropout = nn.Dropout(p=dropout)
self.pe = nn.Parameter(torch.randn(max_len, 1, d_model))
def forward(self, x):
x = x + self.pe[:x.shape[0]]
return self.dropout(x)
@@ -0,0 +1,361 @@
from cv2 import norm
import torch
import torch.nn.functional as F
from torch import layer_norm, nn
import numpy as np
import clip
import random
import math
from ..builder import SUBMODULES, build_attention
from .diffusion_transformer import DiffusionTransformer
def zero_module(module):
"""
Zero out the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().zero_()
return module
def set_requires_grad(nets, requires_grad=False):
"""Set requies_grad for all the networks.
Args:
nets (nn.Module | list[nn.Module]): A list of networks or a single
network.
requires_grad (bool): Whether the networks require gradients or not
"""
if not isinstance(nets, list):
nets = [nets]
for net in nets:
if net is not None:
for param in net.parameters():
param.requires_grad = requires_grad
class FFN(nn.Module):
def __init__(self, latent_dim, ffn_dim, dropout):
super().__init__()
self.linear1 = nn.Linear(latent_dim, ffn_dim)
self.linear2 = zero_module(nn.Linear(ffn_dim, latent_dim))
self.activation = nn.GELU()
self.dropout = nn.Dropout(dropout)
def forward(self, x, **kwargs):
y = self.linear2(self.dropout(self.activation(self.linear1(x))))
y = x + y
return y
class EncoderLayer(nn.Module):
def __init__(self,
sa_block_cfg=None,
ca_block_cfg=None,
ffn_cfg=None):
super().__init__()
self.sa_block = build_attention(sa_block_cfg)
self.ffn = FFN(**ffn_cfg)
def forward(self, **kwargs):
if self.sa_block is not None:
x = self.sa_block(**kwargs)
kwargs.update({'x': x})
if self.ffn is not None:
x = self.ffn(**kwargs)
return x
class RetrievalDatabase(nn.Module):
def __init__(self,
num_retrieval=None,
topk=None,
retrieval_file=None,
latent_dim=512,
output_dim=512,
num_layers=2,
num_motion_layers=4,
kinematic_coef=0.1,
max_seq_len=196,
num_heads=8,
ff_size=1024,
stride=4,
sa_block_cfg=None,
ffn_cfg=None,
dropout=0):
super().__init__()
self.num_retrieval = num_retrieval
self.topk = topk
self.latent_dim = latent_dim
self.stride = stride
self.kinematic_coef = kinematic_coef
self.num_layers = num_layers
self.num_motion_layers = num_motion_layers
self.max_seq_len = max_seq_len
data = np.load(retrieval_file)
self.text_features = torch.Tensor(data['text_features'])
self.captions = data['captions']
self.motions = data['motions']
self.m_lengths = data['m_lengths']
self.clip_seq_features = data['clip_seq_features']
self.train_indexes = data.get('train_indexes', None)
self.test_indexes = data.get('test_indexes', None)
self.latent_dim = latent_dim
self.output_dim = output_dim
self.motion_proj = nn.Linear(self.motions.shape[-1], self.latent_dim)
self.motion_pos_embedding = nn.Parameter(torch.randn(max_seq_len, self.latent_dim))
self.motion_encoder_blocks = nn.ModuleList()
for i in range(num_motion_layers):
self.motion_encoder_blocks.append(
EncoderLayer(
sa_block_cfg=sa_block_cfg,
ffn_cfg=ffn_cfg
)
)
TransEncoderLayer = nn.TransformerEncoderLayer(
d_model=self.latent_dim,
nhead=num_heads,
dim_feedforward=ff_size,
dropout=dropout,
activation="gelu")
self.text_encoder = nn.TransformerEncoder(
TransEncoderLayer,
num_layers=num_layers)
self.results = {}
def extract_text_feature(self, text, clip_model, device):
text = clip.tokenize([text], truncate=True).to(device)
with torch.no_grad():
text_features = clip_model.encode_text(text)
return text_features
def encode_text(self, text, device):
with torch.no_grad():
text = clip.tokenize(text, truncate=True).to(device)
x = self.clip.token_embedding(text).type(self.clip.dtype) # [batch_size, n_ctx, d_model]
x = x + self.clip.positional_embedding.type(self.clip.dtype)
x = x.permute(1, 0, 2) # NLD -> LND
x = self.clip.transformer(x)
x = self.clip.ln_final(x).type(self.clip.dtype)
# B, T, D
xf_out = x.permute(1, 0, 2)
return xf_out
def retrieve(self, caption, length, clip_model, device, idx=None):
if self.training and self.train_indexes is not None and idx is not None:
idx = idx.item()
indexes = self.train_indexes[idx]
data = []
cnt = 0
for retr_idx in indexes:
if retr_idx != idx:
data.append(retr_idx)
cnt += 1
if cnt == self.topk:
break
random.shuffle(data)
return data[:self.num_retrieval]
elif not self.training and self.test_indexes is not None and idx is not None:
idx = idx.item()
indexes = self.test_indexes[idx]
data = []
cnt = 0
for retr_idx in indexes:
data.append(retr_idx)
cnt += 1
if cnt == self.topk:
break
# random.shuffle(data)
return data[:self.num_retrieval]
else:
value = hash(caption)
if value in self.results:
return self.results[value]
text_feature = self.extract_text_feature(caption, clip_model, device)
rel_length = torch.LongTensor(self.m_lengths).to(device)
rel_length = torch.abs(rel_length - length) / torch.clamp(rel_length, min=length)
semantic_score = F.cosine_similarity(self.text_features.to(device), text_feature)
kinematic_score = torch.exp(-rel_length * self.kinematic_coef)
score = semantic_score * kinematic_score
indexes = torch.argsort(score, descending=True)
data = []
cnt = 0
for idx in indexes:
caption, motion, m_length = self.captions[idx], self.motions[idx], self.m_lengths[idx]
if not self.training or m_length != length:
cnt += 1
data.append(idx.item())
if cnt == self.num_retrieval:
self.results[value] = data
return data
assert False
def generate_src_mask(self, T, length):
B = len(length)
src_mask = torch.ones(B, T)
for i in range(B):
for j in range(length[i], T):
src_mask[i, j] = 0
return src_mask
def forward(self, captions, lengths, clip_model, device, idx=None):
B = len(captions)
all_indexes = []
for b_ix in range(B):
length = int(lengths[b_ix])
if idx is None:
batch_indexes = self.retrieve(captions[b_ix], length, clip_model, device)
else:
batch_indexes = self.retrieve(captions[b_ix], length, clip_model, device, idx[b_ix])
all_indexes.extend(batch_indexes)
all_indexes = np.array(all_indexes)
N = all_indexes.shape[0]
all_motions = torch.Tensor(self.motions[all_indexes]).to(device)
all_m_lengths = torch.Tensor(self.m_lengths[all_indexes]).long()
all_captions = self.captions[all_indexes].tolist()
T = all_motions.shape[1]
src_mask = self.generate_src_mask(T, all_m_lengths).to(device)
raw_src_mask = src_mask.clone()
re_motion = self.motion_proj(all_motions) + self.motion_pos_embedding.unsqueeze(0)
for module in self.motion_encoder_blocks:
re_motion = module(x=re_motion, src_mask=src_mask.unsqueeze(-1))
re_motion = re_motion.view(B, self.num_retrieval, T, -1).contiguous()
# stride
re_motion = re_motion[:, :, ::self.stride, :].contiguous()
src_mask = src_mask[:, ::self.stride].contiguous()
src_mask = src_mask.view(B, self.num_retrieval, -1).contiguous()
T = 77
all_text_seq_features = torch.Tensor(self.clip_seq_features[all_indexes]).to(device)
all_text_seq_features = all_text_seq_features.permute(1, 0, 2)
re_text = self.text_encoder(all_text_seq_features)
re_text = re_text.permute(1, 0, 2).view(B, self.num_retrieval, T, -1).contiguous()
re_text = re_text[:, :, -1:, :].contiguous()
# T = re_motion.shape[2]
# re_feat = re_feat.view(B, self.num_retrieval * T, -1).contiguous()
re_dict = dict(
re_text=re_text,
re_motion=re_motion,
re_mask=src_mask,
raw_motion=all_motions,
raw_motion_length=all_m_lengths,
raw_motion_mask=raw_src_mask)
return re_dict
@SUBMODULES.register_module()
class ReMoDiffuseTransformer(DiffusionTransformer):
def __init__(self,
retrieval_cfg=None,
scale_func_cfg=None,
**kwargs):
super().__init__(**kwargs)
self.database = RetrievalDatabase(**retrieval_cfg)
self.scale_func_cfg = scale_func_cfg
def scale_func(self, timestep):
coarse_scale = self.scale_func_cfg['coarse_scale']
w = (1 - (1000 - timestep) / 1000) * coarse_scale + 1
if timestep > 100:
if random.randint(0, 1) == 0:
output = {
'both_coef': w,
'text_coef': 0,
'retr_coef': 1 - w,
'none_coef': 0
}
else:
output = {
'both_coef': 0,
'text_coef': w,
'retr_coef': 0,
'none_coef': 1 - w
}
else:
both_coef = self.scale_func_cfg['both_coef']
text_coef = self.scale_func_cfg['text_coef']
retr_coef = self.scale_func_cfg['retr_coef']
none_coef = 1 - both_coef - text_coef - retr_coef
output = {
'both_coef': both_coef,
'text_coef': text_coef,
'retr_coef': retr_coef,
'none_coef': none_coef
}
return output
def get_precompute_condition(self,
text=None,
motion_length=None,
xf_out=None,
re_dict=None,
device=None,
sample_idx=None,
clip_feat=None,
**kwargs):
if xf_out is None:
xf_out = self.encode_text(text, clip_feat, device)
output = {'xf_out': xf_out}
if re_dict is None:
re_dict = self.database(text, motion_length, self.clip, device, idx=sample_idx)
output['re_dict'] = re_dict
return output
def post_process(self, motion):
return motion
def forward_train(self, h=None, src_mask=None, emb=None, xf_out=None, re_dict=None, **kwargs):
B, T = h.shape[0], h.shape[1]
cond_type = torch.randint(0, 100, size=(B, 1, 1)).to(h.device)
for module in self.temporal_decoder_blocks:
h = module(x=h, xf=xf_out, emb=emb, src_mask=src_mask, cond_type=cond_type, re_dict=re_dict)
output = self.out(h).view(B, T, -1).contiguous()
return output
def forward_test(self, h=None, src_mask=None, emb=None, xf_out=None, re_dict=None, timesteps=None, **kwargs):
B, T = h.shape[0], h.shape[1]
both_cond_type = torch.zeros(B, 1, 1).to(h.device) + 99
text_cond_type = torch.zeros(B, 1, 1).to(h.device) + 1
retr_cond_type = torch.zeros(B, 1, 1).to(h.device) + 10
none_cond_type = torch.zeros(B, 1, 1).to(h.device)
all_cond_type = torch.cat((
both_cond_type, text_cond_type, retr_cond_type, none_cond_type
), dim=0)
h = h.repeat(4, 1, 1)
xf_out = xf_out.repeat(4, 1, 1)
emb = emb.repeat(4, 1)
src_mask = src_mask.repeat(4, 1, 1)
if re_dict['re_motion'].shape[0] != h.shape[0]:
re_dict['re_motion'] = re_dict['re_motion'].repeat(4, 1, 1, 1)
re_dict['re_text'] = re_dict['re_text'].repeat(4, 1, 1, 1)
re_dict['re_mask'] = re_dict['re_mask'].repeat(4, 1, 1)
for module in self.temporal_decoder_blocks:
h = module(x=h, xf=xf_out, emb=emb, src_mask=src_mask, cond_type=all_cond_type, re_dict=re_dict)
out = self.out(h).view(4 * B, T, -1).contiguous()
out_both = out[:B].contiguous()
out_text = out[B: 2 * B].contiguous()
out_retr = out[2 * B: 3 * B].contiguous()
out_none = out[3 * B:].contiguous()
coef_cfg = self.scale_func(int(timesteps[0]))
both_coef = coef_cfg['both_coef']
text_coef = coef_cfg['text_coef']
retr_coef = coef_cfg['retr_coef']
none_coef = coef_cfg['none_coef']
output = out_both * both_coef + out_text * text_coef + out_retr * retr_coef + out_none * none_coef
return output
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,13 @@
import torch.nn as nn
def build_MLP(dim_list, latent_dim):
model_list = []
prev = dim_list[0]
for cur in dim_list[1:]:
model_list.append(nn.Linear(prev, cur))
model_list.append(nn.GELU())
prev = cur
model_list.append(nn.Linear(prev, latent_dim))
model = nn.Sequential(*model_list)
return model
@@ -0,0 +1,40 @@
import torch
import torch.nn as nn
def zero_module(module):
"""
Zero out the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().zero_()
return module
class StylizationBlock(nn.Module):
def __init__(self, latent_dim, time_embed_dim, dropout):
super().__init__()
self.emb_layers = nn.Sequential(
nn.SiLU(),
nn.Linear(time_embed_dim, 2 * latent_dim),
)
self.norm = nn.LayerNorm(latent_dim)
self.out_layers = nn.Sequential(
nn.SiLU(),
nn.Dropout(p=dropout),
zero_module(nn.Linear(latent_dim, latent_dim)),
)
def forward(self, h, emb):
"""
h: B, T, D
emb: B, D
"""
# B, 1, 2D
emb_out = self.emb_layers(emb).unsqueeze(1)
# scale: B, 1, D / shift: B, 1, D
scale, shift = torch.chunk(emb_out, 2, dim=2)
h = self.norm(h) * (1 + scale) + shift
h = self.out_layers(h)
return h
@@ -0,0 +1,80 @@
import numpy as np
import pickle
from os.path import join as pjoin
POS_enumerator = {
'VERB': 0,
'NOUN': 1,
'DET': 2,
'ADP': 3,
'NUM': 4,
'AUX': 5,
'PRON': 6,
'ADJ': 7,
'ADV': 8,
'Loc_VIP': 9,
'Body_VIP': 10,
'Obj_VIP': 11,
'Act_VIP': 12,
'Desc_VIP': 13,
'OTHER': 14,
}
Loc_list = ('left', 'right', 'clockwise', 'counterclockwise', 'anticlockwise', 'forward', 'back', 'backward',
'up', 'down', 'straight', 'curve')
Body_list = ('arm', 'chin', 'foot', 'feet', 'face', 'hand', 'mouth', 'leg', 'waist', 'eye', 'knee', 'shoulder', 'thigh')
Obj_List = ('stair', 'dumbbell', 'chair', 'window', 'floor', 'car', 'ball', 'handrail', 'baseball', 'basketball')
Act_list = ('walk', 'run', 'swing', 'pick', 'bring', 'kick', 'put', 'squat', 'throw', 'hop', 'dance', 'jump', 'turn',
'stumble', 'dance', 'stop', 'sit', 'lift', 'lower', 'raise', 'wash', 'stand', 'kneel', 'stroll',
'rub', 'bend', 'balance', 'flap', 'jog', 'shuffle', 'lean', 'rotate', 'spin', 'spread', 'climb')
Desc_list = ('slowly', 'carefully', 'fast', 'careful', 'slow', 'quickly', 'happy', 'angry', 'sad', 'happily',
'angrily', 'sadly')
VIP_dict = {
'Loc_VIP': Loc_list,
'Body_VIP': Body_list,
'Obj_VIP': Obj_List,
'Act_VIP': Act_list,
'Desc_VIP': Desc_list,
}
class WordVectorizer(object):
def __init__(self, meta_root, prefix):
vectors = np.load(pjoin(meta_root, '%s_data.npy'%prefix))
words = pickle.load(open(pjoin(meta_root, '%s_words.pkl'%prefix), 'rb'))
word2idx = pickle.load(open(pjoin(meta_root, '%s_idx.pkl'%prefix), 'rb'))
self.word2vec = {w: vectors[word2idx[w]] for w in words}
def _get_pos_ohot(self, pos):
pos_vec = np.zeros(len(POS_enumerator))
if pos in POS_enumerator:
pos_vec[POS_enumerator[pos]] = 1
else:
pos_vec[POS_enumerator['OTHER']] = 1
return pos_vec
def __len__(self):
return len(self.word2vec)
def __getitem__(self, item):
word, pos = item.split('/')
if word in self.word2vec:
word_vec = self.word2vec[word]
vip_pos = None
for key, values in VIP_dict.items():
if word in values:
vip_pos = key
break
if vip_pos is not None:
pos_vec = self._get_pos_ohot(vip_pos)
else:
pos_vec = self._get_pos_ohot(pos)
else:
word_vec = self.word2vec['unk']
pos_vec = self._get_pos_ohot('OTHER')
return word_vec, pos_vec
Binary file not shown.

After

Width:  |  Height:  |  Size: 5.1 MiB

File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -0,0 +1,41 @@
import numpy as np
from pathlib import Path
# Map joints Name to SMPL joints idx
JOINT_MAP = {
'MidHip': 0,
'LHip': 1, 'LKnee': 4, 'LAnkle': 7, 'LFoot': 10,
'RHip': 2, 'RKnee': 5, 'RAnkle': 8, 'RFoot': 11,
'LShoulder': 16, 'LElbow': 18, 'LWrist': 20, 'LHand': 22,
'RShoulder': 17, 'RElbow': 19, 'RWrist': 21, 'RHand': 23,
'spine1': 3, 'spine2': 6, 'spine3': 9, 'Neck': 12, 'Head': 15,
'LCollar':13, 'Rcollar' :14,
'Nose':24, 'REye':26, 'LEye':26, 'REar':27, 'LEar':28,
'LHeel': 31, 'RHeel': 34,
'OP RShoulder': 17, 'OP LShoulder': 16,
'OP RHip': 2, 'OP LHip': 1,
'OP Neck': 12,
}
full_smpl_idx = range(24)
key_smpl_idx = [0, 1, 4, 7, 2, 5, 8, 17, 19, 21, 16, 18, 20]
AMASS_JOINT_MAP = {
'MidHip': 0,
'LHip': 1, 'LKnee': 4, 'LAnkle': 7, 'LFoot': 10,
'RHip': 2, 'RKnee': 5, 'RAnkle': 8, 'RFoot': 11,
'LShoulder': 16, 'LElbow': 18, 'LWrist': 20,
'RShoulder': 17, 'RElbow': 19, 'RWrist': 21,
'spine1': 3, 'spine2': 6, 'spine3': 9, 'Neck': 12, 'Head': 15,
'LCollar':13, 'Rcollar' :14,
}
amass_idx = range(22)
amass_smpl_idx = range(22)
SMPL_MODEL_DIR = str(Path(__file__).parent.parent.parent / "body_models")
GMM_MODEL_DIR = str(Path(__file__).parent.parent / "smpl_models")
SMPL_MEAN_FILE = str(Path(__file__).parent.parent / "smpl_models" / "neutral_smpl_mean_params.h5")
# for collsion
Part_Seg_DIR = str(Path(__file__).parent / "smpl_models" / "smplx_parts_segm.pkl")
@@ -0,0 +1,222 @@
import torch
import torch.nn.functional as F
from motiondiff_modules.mogen.smpl.joints2smpl.src import config
# Guassian
def gmof(x, sigma):
"""
Geman-McClure error function
"""
x_squared = x ** 2
sigma_squared = sigma ** 2
return (sigma_squared * x_squared) / (sigma_squared + x_squared)
# angle prior
def angle_prior(pose):
"""
Angle prior that penalizes unnatural bending of the knees and elbows
"""
# We subtract 3 because pose does not include the global rotation of the model
return torch.exp(
pose[:, [55 - 3, 58 - 3, 12 - 3, 15 - 3]] * torch.tensor([1., -1., -1, -1.], device=pose.device)) ** 2
def perspective_projection(points, rotation, translation,
focal_length, camera_center):
"""
This function computes the perspective projection of a set of points.
Input:
points (bs, N, 3): 3D points
rotation (bs, 3, 3): Camera rotation
translation (bs, 3): Camera translation
focal_length (bs,) or scalar: Focal length
camera_center (bs, 2): Camera center
"""
batch_size = points.shape[0]
K = torch.zeros([batch_size, 3, 3], device=points.device)
K[:, 0, 0] = focal_length
K[:, 1, 1] = focal_length
K[:, 2, 2] = 1.
K[:, :-1, -1] = camera_center
# Transform points
points = torch.einsum('bij,bkj->bki', rotation, points)
points = points + translation.unsqueeze(1)
# Apply perspective distortion
projected_points = points / points[:, :, -1].unsqueeze(-1)
# Apply camera intrinsics
projected_points = torch.einsum('bij,bkj->bki', K, projected_points)
return projected_points[:, :, :-1]
def body_fitting_loss(body_pose, betas, model_joints, camera_t, camera_center,
joints_2d, joints_conf, pose_prior,
focal_length=5000, sigma=100, pose_prior_weight=4.78,
shape_prior_weight=5, angle_prior_weight=15.2,
output='sum'):
"""
Loss function for body fitting
"""
batch_size = body_pose.shape[0]
rotation = torch.eye(3, device=body_pose.device).unsqueeze(0).expand(batch_size, -1, -1)
projected_joints = perspective_projection(model_joints, rotation, camera_t,
focal_length, camera_center)
# Weighted robust reprojection error
reprojection_error = gmof(projected_joints - joints_2d, sigma)
reprojection_loss = (joints_conf ** 2) * reprojection_error.sum(dim=-1)
# Pose prior loss
pose_prior_loss = (pose_prior_weight ** 2) * pose_prior(body_pose, betas)
# Angle prior for knees and elbows
angle_prior_loss = (angle_prior_weight ** 2) * angle_prior(body_pose).sum(dim=-1)
# Regularizer to prevent betas from taking large values
shape_prior_loss = (shape_prior_weight ** 2) * (betas ** 2).sum(dim=-1)
total_loss = reprojection_loss.sum(dim=-1) + pose_prior_loss + angle_prior_loss + shape_prior_loss
if output == 'sum':
return total_loss.sum()
elif output == 'reprojection':
return reprojection_loss
# --- get camera fitting loss -----
def camera_fitting_loss(model_joints, camera_t, camera_t_est, camera_center,
joints_2d, joints_conf,
focal_length=5000, depth_loss_weight=100):
"""
Loss function for camera optimization.
"""
# Project model joints
batch_size = model_joints.shape[0]
rotation = torch.eye(3, device=model_joints.device).unsqueeze(0).expand(batch_size, -1, -1)
projected_joints = perspective_projection(model_joints, rotation, camera_t,
focal_length, camera_center)
# get the indexed four
op_joints = ['OP RHip', 'OP LHip', 'OP RShoulder', 'OP LShoulder']
op_joints_ind = [config.JOINT_MAP[joint] for joint in op_joints]
gt_joints = ['RHip', 'LHip', 'RShoulder', 'LShoulder']
gt_joints_ind = [config.JOINT_MAP[joint] for joint in gt_joints]
reprojection_error_op = (joints_2d[:, op_joints_ind] -
projected_joints[:, op_joints_ind]) ** 2
reprojection_error_gt = (joints_2d[:, gt_joints_ind] -
projected_joints[:, gt_joints_ind]) ** 2
# Check if for each example in the batch all 4 OpenPose detections are valid, otherwise use the GT detections
# OpenPose joints are more reliable for this task, so we prefer to use them if possible
is_valid = (joints_conf[:, op_joints_ind].min(dim=-1)[0][:, None, None] > 0).float()
reprojection_loss = (is_valid * reprojection_error_op + (1 - is_valid) * reprojection_error_gt).sum(dim=(1, 2))
# Loss that penalizes deviation from depth estimate
depth_loss = (depth_loss_weight ** 2) * (camera_t[:, 2] - camera_t_est[:, 2]) ** 2
total_loss = reprojection_loss + depth_loss
return total_loss.sum()
# #####--- body fitiing loss -----
def body_fitting_loss_3d(body_pose, preserve_pose,
betas, model_joints, camera_translation,
j3d, pose_prior,
joints3d_conf,
sigma=100, pose_prior_weight=4.78*1.5,
shape_prior_weight=5.0, angle_prior_weight=15.2,
joint_loss_weight=500.0,
pose_preserve_weight=0.0,
use_collision=False,
model_vertices=None, model_faces=None,
search_tree=None, pen_distance=None, filter_faces=None,
collision_loss_weight=1000
):
"""
Loss function for body fitting
"""
batch_size = body_pose.shape[0]
#joint3d_loss = (joint_loss_weight ** 2) * gmof((model_joints + camera_translation) - j3d, sigma).sum(dim=-1)
joint3d_error = gmof((model_joints + camera_translation) - j3d, sigma)
joint3d_loss_part = (joints3d_conf ** 2) * joint3d_error.sum(dim=-1)
joint3d_loss = ((joint_loss_weight ** 2) * joint3d_loss_part).sum(dim=-1)
# Pose prior loss
pose_prior_loss = (pose_prior_weight ** 2) * pose_prior(body_pose, betas)
# Angle prior for knees and elbows
angle_prior_loss = (angle_prior_weight ** 2) * angle_prior(body_pose).sum(dim=-1)
# Regularizer to prevent betas from taking large values
shape_prior_loss = (shape_prior_weight ** 2) * (betas ** 2).sum(dim=-1)
collision_loss = 0.0
# Calculate the loss due to interpenetration
if use_collision:
triangles = torch.index_select(
model_vertices, 1,
model_faces).view(batch_size, -1, 3, 3)
with torch.no_grad():
collision_idxs = search_tree(triangles)
# Remove unwanted collisions
if filter_faces is not None:
collision_idxs = filter_faces(collision_idxs)
if collision_idxs.ge(0).sum().item() > 0:
collision_loss = torch.sum(collision_loss_weight * pen_distance(triangles, collision_idxs))
pose_preserve_loss = (pose_preserve_weight ** 2) * ((body_pose - preserve_pose) ** 2).sum(dim=-1)
# print('joint3d_loss', joint3d_loss.shape)
# print('pose_prior_loss', pose_prior_loss.shape)
# print('angle_prior_loss', angle_prior_loss.shape)
# print('shape_prior_loss', shape_prior_loss.shape)
# print('collision_loss', collision_loss)
# print('pose_preserve_loss', pose_preserve_loss.shape)
total_loss = joint3d_loss + pose_prior_loss + angle_prior_loss + shape_prior_loss + collision_loss + pose_preserve_loss
return total_loss.sum()
# #####--- get camera fitting loss -----
def camera_fitting_loss_3d(model_joints, camera_t, camera_t_est,
j3d, joints_category="orig", depth_loss_weight=100.0):
"""
Loss function for camera optimization.
"""
model_joints = model_joints + camera_t
# # get the indexed four
# op_joints = ['OP RHip', 'OP LHip', 'OP RShoulder', 'OP LShoulder']
# op_joints_ind = [config.JOINT_MAP[joint] for joint in op_joints]
#
# j3d_error_loss = (j3d[:, op_joints_ind] -
# model_joints[:, op_joints_ind]) ** 2
gt_joints = ['RHip', 'LHip', 'RShoulder', 'LShoulder']
gt_joints_ind = [config.JOINT_MAP[joint] for joint in gt_joints]
if joints_category=="orig":
select_joints_ind = [config.JOINT_MAP[joint] for joint in gt_joints]
elif joints_category=="AMASS":
select_joints_ind = [config.AMASS_JOINT_MAP[joint] for joint in gt_joints]
else:
print("NO SUCH JOINTS CATEGORY!")
j3d_error_loss = (j3d[:, select_joints_ind] -
model_joints[:, gt_joints_ind]) ** 2
# Loss that penalizes deviation from depth estimate
depth_loss = (depth_loss_weight**2) * (camera_t - camera_t_est)**2
total_loss = j3d_error_loss + depth_loss
return total_loss.sum()
@@ -0,0 +1,240 @@
# -*- coding: utf-8 -*-
# Max-Planck-Gesellschaft zur Förderung der Wissenschaften e.V. (MPG) is
# holder of all proprietary rights on this computer program.
# You can only use this computer program if you have closed
# a license agreement with MPG or you get the right to use the computer
# program from someone who is authorized to grant you that right.
# Any use of the computer program without a valid license is prohibited and
# liable to prosecution.
#
# Copyright©2019 Max-Planck-Gesellschaft zur Förderung
# der Wissenschaften e.V. (MPG). acting on behalf of its Max Planck Institute
# for Intelligent Systems. All rights reserved.
#
# Contact: ps-license@tuebingen.mpg.de
from __future__ import absolute_import
from __future__ import print_function
from __future__ import division
import sys
import os
import time
import pickle
import numpy as np
import torch
import torch.nn as nn
DEFAULT_DTYPE = torch.float32
def create_prior(prior_type, **kwargs):
if prior_type == 'gmm':
prior = MaxMixturePrior(**kwargs)
elif prior_type == 'l2':
return L2Prior(**kwargs)
elif prior_type == 'angle':
return SMPLifyAnglePrior(**kwargs)
elif prior_type == 'none' or prior_type is None:
# Don't use any pose prior
def no_prior(*args, **kwargs):
return 0.0
prior = no_prior
else:
raise ValueError('Prior {}'.format(prior_type) + ' is not implemented')
return prior
class SMPLifyAnglePrior(nn.Module):
def __init__(self, dtype=torch.float32, **kwargs):
super(SMPLifyAnglePrior, self).__init__()
# Indices for the roration angle of
# 55: left elbow, 90deg bend at -np.pi/2
# 58: right elbow, 90deg bend at np.pi/2
# 12: left knee, 90deg bend at np.pi/2
# 15: right knee, 90deg bend at np.pi/2
angle_prior_idxs = np.array([55, 58, 12, 15], dtype=np.int64)
angle_prior_idxs = torch.tensor(angle_prior_idxs, dtype=torch.long)
self.register_buffer('angle_prior_idxs', angle_prior_idxs)
angle_prior_signs = np.array([1, -1, -1, -1],
dtype=np.float32 if dtype == torch.float32
else np.float64)
angle_prior_signs = torch.tensor(angle_prior_signs,
dtype=dtype)
self.register_buffer('angle_prior_signs', angle_prior_signs)
def forward(self, pose, with_global_pose=False):
''' Returns the angle prior loss for the given pose
Args:
pose: (Bx[23 + 1] * 3) torch tensor with the axis-angle
representation of the rotations of the joints of the SMPL model.
Kwargs:
with_global_pose: Whether the pose vector also contains the global
orientation of the SMPL model. If not then the indices must be
corrected.
Returns:
A sze (B) tensor containing the angle prior loss for each element
in the batch.
'''
angle_prior_idxs = self.angle_prior_idxs - (not with_global_pose) * 3
return torch.exp(pose[:, angle_prior_idxs] *
self.angle_prior_signs).pow(2)
class L2Prior(nn.Module):
def __init__(self, dtype=DEFAULT_DTYPE, reduction='sum', **kwargs):
super(L2Prior, self).__init__()
def forward(self, module_input, *args):
return torch.sum(module_input.pow(2))
def fix_pickle(original, destination):
content = ''
outsize = 0
with open(original, 'rb') as infile:
content = infile.read()
with open(destination, 'wb') as output:
for line in content.splitlines():
outsize += len(line) + 1
output.write(line + str.encode('\n'))
class MaxMixturePrior(nn.Module):
def __init__(self, prior_folder='prior',
num_gaussians=6, dtype=DEFAULT_DTYPE, epsilon=1e-16,
use_merged=True,
**kwargs):
super(MaxMixturePrior, self).__init__()
if dtype == DEFAULT_DTYPE:
np_dtype = np.float32
elif dtype == torch.float64:
np_dtype = np.float64
else:
print('Unknown float type {}, exiting!'.format(dtype))
sys.exit(-1)
self.num_gaussians = num_gaussians
self.epsilon = epsilon
self.use_merged = use_merged
gmm_fn = 'gmm_{:02d}.pkl'.format(num_gaussians)
full_gmm_fn = os.path.join(prior_folder, gmm_fn)
if not os.path.exists(full_gmm_fn):
print('The path to the mixture prior "{}"'.format(full_gmm_fn) +
' does not exist, exiting!')
sys.exit(-1)
fix_pickle(full_gmm_fn, full_gmm_fn) #https://stackoverflow.com/a/46587771
with open(full_gmm_fn, 'rb') as f:
gmm = pickle.load(f, encoding='latin1')
if type(gmm) == dict:
means = gmm['means'].astype(np_dtype)
covs = gmm['covars'].astype(np_dtype)
weights = gmm['weights'].astype(np_dtype)
elif 'sklearn.mixture.gmm.GMM' in str(type(gmm)):
means = gmm.means_.astype(np_dtype)
covs = gmm.covars_.astype(np_dtype)
weights = gmm.weights_.astype(np_dtype)
else:
print('Unknown type for the prior: {}, exiting!'.format(type(gmm)))
sys.exit(-1)
self.register_buffer('means', torch.tensor(means, dtype=dtype))
self.register_buffer('covs', torch.tensor(covs, dtype=dtype))
precisions = [np.linalg.inv(cov) for cov in covs]
precisions = np.stack(precisions).astype(np_dtype)
self.register_buffer('precisions',
torch.tensor(precisions, dtype=dtype))
# The constant term:
sqrdets = np.array([(np.sqrt(np.linalg.det(c)))
for c in gmm['covars']])
const = (2 * np.pi)**(69 / 2.)
nll_weights = np.asarray(gmm['weights'] / (const *
(sqrdets / sqrdets.min())))
nll_weights = torch.tensor(nll_weights, dtype=dtype).unsqueeze(dim=0)
self.register_buffer('nll_weights', nll_weights)
weights = torch.tensor(gmm['weights'], dtype=dtype).unsqueeze(dim=0)
self.register_buffer('weights', weights)
self.register_buffer('pi_term',
torch.log(torch.tensor(2 * np.pi, dtype=dtype)))
cov_dets = [np.log(np.linalg.det(cov.astype(np_dtype)) + epsilon)
for cov in covs]
self.register_buffer('cov_dets',
torch.tensor(cov_dets, dtype=dtype))
# The dimensionality of the random variable
self.random_var_dim = self.means.shape[1]
def get_mean(self):
''' Returns the mean of the mixture '''
mean_pose = torch.matmul(self.weights, self.means)
return mean_pose
def merged_log_likelihood(self, pose, betas):
diff_from_mean = pose.unsqueeze(dim=1) - self.means
prec_diff_prod = torch.einsum('mij,bmj->bmi',
[self.precisions, diff_from_mean])
diff_prec_quadratic = (prec_diff_prod * diff_from_mean).sum(dim=-1)
curr_loglikelihood = 0.5 * diff_prec_quadratic - \
torch.log(self.nll_weights)
# curr_loglikelihood = 0.5 * (self.cov_dets.unsqueeze(dim=0) +
# self.random_var_dim * self.pi_term +
# diff_prec_quadratic
# ) - torch.log(self.weights)
min_likelihood, _ = torch.min(curr_loglikelihood, dim=1)
return min_likelihood
def log_likelihood(self, pose, betas, *args, **kwargs):
''' Create graph operation for negative log-likelihood calculation
'''
likelihoods = []
for idx in range(self.num_gaussians):
mean = self.means[idx]
prec = self.precisions[idx]
cov = self.covs[idx]
diff_from_mean = pose - mean
curr_loglikelihood = torch.einsum('bj,ji->bi',
[diff_from_mean, prec])
curr_loglikelihood = torch.einsum('bi,bi->b',
[curr_loglikelihood,
diff_from_mean])
cov_term = torch.log(torch.det(cov) + self.epsilon)
curr_loglikelihood += 0.5 * (cov_term +
self.random_var_dim *
self.pi_term)
likelihoods.append(curr_loglikelihood)
log_likelihoods = torch.stack(likelihoods, dim=1)
min_idx = torch.argmin(log_likelihoods, dim=1)
weight_component = self.nll_weights[:, min_idx]
weight_component = -torch.log(weight_component)
return weight_component + log_likelihoods[:, min_idx]
def forward(self, pose, betas):
if self.use_merged:
return self.merged_log_likelihood(pose, betas)
else:
return self.log_likelihood(pose, betas)
@@ -0,0 +1,294 @@
import torch
import os, sys
import pickle
import smplx
import numpy as np
sys.path.append(os.path.dirname(__file__))
from customloss import (camera_fitting_loss,
body_fitting_loss,
camera_fitting_loss_3d,
body_fitting_loss_3d,
)
from prior import MaxMixturePrior
from motiondiff_modules.mogen.smpl.joints2smpl.src import config
from tqdm import tqdm
import comfy.utils
@torch.no_grad()
def guess_init_3d(model_joints,
j3d,
joints_category="orig"):
"""Initialize the camera translation via triangle similarity, by using the torso joints .
:param model_joints: SMPL model with pre joints
:param j3d: 25x3 array of Kinect Joints
:returns: 3D vector corresponding to the estimated camera translation
"""
# get the indexed four
gt_joints = ['RHip', 'LHip', 'RShoulder', 'LShoulder']
gt_joints_ind = [config.JOINT_MAP[joint] for joint in gt_joints]
if joints_category=="orig":
joints_ind_category = [config.JOINT_MAP[joint] for joint in gt_joints]
elif joints_category=="AMASS":
joints_ind_category = [config.AMASS_JOINT_MAP[joint] for joint in gt_joints]
else:
print("NO SUCH JOINTS CATEGORY!")
sum_init_t = (j3d[:, joints_ind_category] - model_joints[:, gt_joints_ind]).sum(dim=1)
init_t = sum_init_t / 4.0
return init_t
# SMPLIfy 3D
class SMPLify3D():
"""Implementation of SMPLify, use 3D joints."""
def __init__(self,
smplxmodel,
step_size=1e-2,
batch_size=1,
num_iters=100,
use_collision=False,
use_lbfgs=True,
joints_category="orig",
device=torch.device('cuda:0')
):
# Store options
self.batch_size = batch_size
self.device = device
self.step_size = step_size
self.num_iters = num_iters
# --- choose optimizer
self.use_lbfgs = use_lbfgs
# GMM pose prior
self.pose_prior = MaxMixturePrior(prior_folder=config.GMM_MODEL_DIR,
num_gaussians=8,
dtype=torch.float32).to(device)
# collision part
self.use_collision = use_collision
if self.use_collision:
self.part_segm_fn = config.Part_Seg_DIR
# reLoad SMPL-X model
self.smpl = smplxmodel
self.model_faces = smplxmodel.faces_tensor.view(-1)
# select joint joint_category
self.joints_category = joints_category
if joints_category=="orig":
self.smpl_index = config.full_smpl_idx
self.corr_index = config.full_smpl_idx
elif joints_category=="AMASS":
self.smpl_index = config.amass_smpl_idx
self.corr_index = config.amass_idx
else:
self.smpl_index = None
self.corr_index = None
print("NO SUCH JOINTS CATEGORY!")
# ---- get the man function here ------
def __call__(self, init_pose, init_betas, init_cam_t, j3d, conf_3d=1.0, seq_ind=0):
"""Perform body fitting.
Input:
init_pose: SMPL pose estimate
init_betas: SMPL betas estimate
init_cam_t: Camera translation estimate
j3d: joints 3d aka keypoints
conf_3d: confidence for 3d joints
seq_ind: index of the sequence
Returns:
vertices: Vertices of optimized shape
joints: 3D joints of optimized shape
pose: SMPL pose parameters of optimized shape
betas: SMPL beta parameters of optimized shape
camera_translation: Camera translation
"""
# # # add the mesh inter-section to avoid
search_tree = None
pen_distance = None
filter_faces = None
if self.use_collision:
from mesh_intersection.bvh_search_tree import BVH
import mesh_intersection.loss as collisions_loss
from mesh_intersection.filter_faces import FilterFaces
search_tree = BVH(max_collisions=8)
pen_distance = collisions_loss.DistanceFieldPenetrationLoss(
sigma=0.5, point2plane=False, vectorized=True, penalize_outside=True)
if self.part_segm_fn:
# Read the part segmentation
part_segm_fn = os.path.expandvars(self.part_segm_fn)
with open(part_segm_fn, 'rb') as faces_parents_file:
face_segm_data = pickle.load(faces_parents_file, encoding='latin1')
faces_segm = face_segm_data['segm']
faces_parents = face_segm_data['parents']
# Create the module used to filter invalid collision pairs
filter_faces = FilterFaces(
faces_segm=faces_segm, faces_parents=faces_parents,
ign_part_pairs=None).to(device=self.device)
# Split SMPL pose to body pose and global orientation
body_pose = init_pose[:, 3:].detach().clone()
global_orient = init_pose[:, :3].detach().clone()
betas = init_betas.detach().clone()
# use guess 3d to get the initial
smpl_output = self.smpl(global_orient=global_orient,
body_pose=body_pose,
betas=betas)
model_joints = smpl_output.joints
init_cam_t = guess_init_3d(model_joints, j3d, self.joints_category).unsqueeze(1).detach()
camera_translation = init_cam_t.clone()
preserve_pose = init_pose[:, 3:].detach().clone()
# -------------Step 1: Optimize camera translation and body orientation--------
# Optimize only camera translation and body orientation
print("Optimizing step 1/2: camera translation and body orientation...")
body_pose.requires_grad = False
betas.requires_grad = False
global_orient.requires_grad = True
camera_translation.requires_grad = True
camera_opt_params = [global_orient, camera_translation]
if self.use_lbfgs:
camera_optimizer = torch.optim.LBFGS(camera_opt_params, max_iter=self.num_iters,
lr=self.step_size, line_search_fn='strong_wolfe')
pbar_comfy = comfy.utils.ProgressBar(10)
pbar = tqdm(range(10))
for i in pbar:
def closure():
camera_optimizer.zero_grad()
smpl_output = self.smpl(global_orient=global_orient,
body_pose=body_pose,
betas=betas)
model_joints = smpl_output.joints
# print('model_joints', model_joints.shape)
# print('camera_translation', camera_translation.shape)
# print('init_cam_t', init_cam_t.shape)
# print('j3d', j3d.shape)
loss = camera_fitting_loss_3d(model_joints, camera_translation,
init_cam_t, j3d, self.joints_category)
loss.backward()
pbar.set_postfix(loss=loss.item())
return loss
pbar_comfy.update(1)
camera_optimizer.step(closure)
else:
camera_optimizer = torch.optim.Adam(camera_opt_params, lr=self.step_size, betas=(0.9, 0.999))
pbar_comfy = comfy.utils.ProgressBar(20)
pbar = tqdm(range(20))
for i in pbar:
smpl_output = self.smpl(global_orient=global_orient,
body_pose=body_pose,
betas=betas)
model_joints = smpl_output.joints
loss = camera_fitting_loss_3d(model_joints[:, self.smpl_index], camera_translation,
init_cam_t, j3d[:, self.corr_index], self.joints_category)
camera_optimizer.zero_grad()
loss.backward()
camera_optimizer.step()
pbar.set_postfix(loss=loss.item())
pbar_comfy.update(1)
# Fix camera translation after optimizing camera
# --------Step 2: Optimize body joints --------------------------
# Optimize only the body pose and global orientation of the body
print("Optimizing step 2/2: body joints...")
body_pose.requires_grad = True
global_orient.requires_grad = True
camera_translation.requires_grad = True
# --- if we use the sequence, fix the shape
if seq_ind == 0:
betas.requires_grad = True
body_opt_params = [body_pose, betas, global_orient, camera_translation]
else:
betas.requires_grad = False
body_opt_params = [body_pose, global_orient, camera_translation]
if self.use_lbfgs:
body_optimizer = torch.optim.LBFGS(body_opt_params, max_iter=self.num_iters,
lr=self.step_size, line_search_fn='strong_wolfe')
pbar_comfy = comfy.utils.ProgressBar(self.num_iters)
pbar = tqdm(range(self.num_iters))
for i in pbar:
def closure():
body_optimizer.zero_grad()
smpl_output = self.smpl(global_orient=global_orient,
body_pose=body_pose,
betas=betas)
model_joints = smpl_output.joints
model_vertices = smpl_output.vertices
loss = body_fitting_loss_3d(body_pose, preserve_pose, betas, model_joints[:, self.smpl_index], camera_translation,
j3d[:, self.corr_index], self.pose_prior,
joints3d_conf=conf_3d,
joint_loss_weight=600.0,
pose_preserve_weight=5.0,
use_collision=self.use_collision,
model_vertices=model_vertices, model_faces=self.model_faces,
search_tree=search_tree, pen_distance=pen_distance, filter_faces=filter_faces)
loss.backward()
pbar.set_postfix(loss=loss.item())
return loss
pbar_comfy.update(1)
body_optimizer.step(closure)
else:
body_optimizer = torch.optim.Adam(body_opt_params, lr=self.step_size, betas=(0.9, 0.999))
pbar = tqdm(range(self.num_iters))
pbar_comfy = comfy.utils.ProgressBar(self.num_iters)
for i in pbar:
smpl_output = self.smpl(global_orient=global_orient,
body_pose=body_pose,
betas=betas)
model_joints = smpl_output.joints
model_vertices = smpl_output.vertices
loss = body_fitting_loss_3d(body_pose, preserve_pose, betas, model_joints[:, self.smpl_index], camera_translation,
j3d[:, self.corr_index], self.pose_prior,
joints3d_conf=conf_3d,
joint_loss_weight=600.0,
use_collision=self.use_collision,
model_vertices=model_vertices, model_faces=self.model_faces,
search_tree=search_tree, pen_distance=pen_distance, filter_faces=filter_faces)
body_optimizer.zero_grad()
loss.backward()
body_optimizer.step()
pbar_comfy.update(1)
# Get final loss value
with torch.no_grad():
smpl_output = self.smpl(global_orient=global_orient,
body_pose=body_pose,
betas=betas, return_full_pose=True)
model_joints = smpl_output.joints
model_vertices = smpl_output.vertices
final_loss = body_fitting_loss_3d(body_pose, preserve_pose, betas, model_joints[:, self.smpl_index], camera_translation,
j3d[:, self.corr_index], self.pose_prior,
joints3d_conf=conf_3d,
joint_loss_weight=600.0,
use_collision=self.use_collision, model_vertices=model_vertices, model_faces=self.model_faces,
search_tree=search_tree, pen_distance=pen_distance, filter_faces=filter_faces)
vertices = smpl_output.vertices.detach()
joints = smpl_output.joints.detach()
pose = torch.cat([global_orient, body_pose], dim=-1).detach()
betas = betas.detach()
return vertices, joints, pose, betas, camera_translation, final_loss
@@ -0,0 +1,442 @@
#Based on https://github.com/Mael-zys/T2M-GPT/blob/main/render_final.py
from motiondiff_modules.mogen.smpl.rotation2xyz import Rotation2xyz
import numpy as np
from trimesh import Trimesh
import os
#https://stackoverflow.com/a/45756291
if os.name == 'posix' and "DISPLAY" not in os.environ:
os.environ['PYOPENGL_PLATFORM'] = "egl"
import torch
import comfy.utils
from motiondiff_modules.mogen.smpl.simplify_loc2rot import joints2smpl
import pyrender
from pyrender.shader_program import ShaderProgramCache
from shapely import geometry
import trimesh
from pyrender.constants import RenderFlags
from comfy.model_management import get_torch_device
from tqdm import tqdm
from PIL import Image
shader_dir = os.path.join(os.path.dirname(__file__), 'shaders')
class WeakPerspectiveCamera(pyrender.Camera):
def __init__(self,
scale,
translation,
znear=pyrender.camera.DEFAULT_Z_NEAR,
zfar=None,
name=None):
super(WeakPerspectiveCamera, self).__init__(
znear=znear,
zfar=zfar,
name=name,
)
self.scale = scale
self.translation = translation
def get_projection_matrix(self, width=None, height=None):
P = np.eye(4)
P[0, 0] = self.scale[0]
P[1, 1] = self.scale[1]
P[0, 3] = self.translation[0] * self.scale[0]
P[1, 3] = -self.translation[1] * self.scale[1]
P[2, 2] = -1
return P
def render(motions):
frames, njoints, nfeats = motions.shape
MINS = motions.min(axis=0).min(axis=0)
MAXS = motions.max(axis=0).max(axis=0)
height_offset = MINS[1]
motions[:, :, 1] -= height_offset
trajec = motions[:, 0, [0, 2]]
j2s = joints2smpl(num_frames=frames, device=get_torch_device())
rot2xyz = Rotation2xyz(device=get_torch_device())
faces = rot2xyz.smpl_model.faces
print(f'Running SMPLify, it may take a few minutes.')
motion_tensor, opt_dict = j2s.forward(motions) # [nframes, njoints, 3]
vertices = rot2xyz(torch.tensor(motion_tensor).clone(), mask=None,
pose_rep='rot6d', translation=True, glob=True,
jointstype='vertices',
vertstrans=True)
frames = vertices.shape[3] # shape: 1, nb_frames, 3, nb_joints
MINS = torch.min(torch.min(vertices[0], axis=0)[0], axis=1)[0]
MAXS = torch.max(torch.max(vertices[0], axis=0)[0], axis=1)[0]
# vertices[:,:,1,:] -= MINS[1] + 1e-5
out_list = []
minx = MINS[0] - 0.5
maxx = MAXS[0] + 0.5
minz = MINS[2] - 0.5
maxz = MAXS[2] + 0.5
polygon = geometry.Polygon([[minx, minz], [minx, maxz], [maxx, maxz], [maxx, minz]])
polygon_mesh = trimesh.creation.extrude_polygon(polygon, 1e-5)
vid = []
for i in range(frames):
if i % 10 == 0:
print(i)
mesh = Trimesh(vertices=vertices[0, :, :, i].squeeze().tolist(), faces=faces)
base_color = (0.11, 0.53, 0.8, 0.5)
## OPAQUE rendering without alpha
## BLEND rendering consider alpha
material = pyrender.MetallicRoughnessMaterial(
metallicFactor=0.7,
alphaMode='OPAQUE',
baseColorFactor=base_color
)
mesh = pyrender.Mesh.from_trimesh(mesh, material=material)
polygon_mesh.visual.face_colors = [0, 0, 0, 0.21]
polygon_render = pyrender.Mesh.from_trimesh(polygon_mesh, smooth=False)
bg_color = [1, 1, 1, 0.8]
scene = pyrender.Scene(bg_color=bg_color, ambient_light=(0.4, 0.4, 0.4))
sx, sy, tx, ty = [0.75, 0.75, 0, 0.10]
camera = pyrender.PerspectiveCamera(yfov=(np.pi / 3.0))
light = pyrender.DirectionalLight(color=[1,1,1], intensity=300)
scene.add(mesh)
c = np.pi / 2
scene.add(polygon_render, pose=np.array([[ 1, 0, 0, 0],
[ 0, np.cos(c), -np.sin(c), MINS[1].cpu().numpy()],
[ 0, np.sin(c), np.cos(c), 0],
[ 0, 0, 0, 1]]))
light_pose = np.eye(4)
light_pose[:3, 3] = [0, -1, 1]
scene.add(light, pose=light_pose.copy())
light_pose[:3, 3] = [0, 1, 1]
scene.add(light, pose=light_pose.copy())
light_pose[:3, 3] = [1, 1, 2]
scene.add(light, pose=light_pose.copy())
c = -np.pi / 6
scene.add(camera, pose=[[ 1, 0, 0, (minx+maxx).cpu().numpy()/2],
[ 0, np.cos(c), -np.sin(c), 1.5],
[ 0, np.sin(c), np.cos(c), max(4, minz.cpu().numpy()+(1.5-MINS[1].cpu().numpy())*2, (maxx-minx).cpu().numpy())],
[ 0, 0, 0, 1]
])
# render scene
r = pyrender.OffscreenRenderer(960, 960)
color, _ = r.render(scene, flags=RenderFlags.RGBA)
# Image.fromarray(color).save(outdir+name+'_'+str(i)+'.png')
vid.append(color)
r.delete()
out = np.stack(vid, axis=0)
return out
def render_from_smpl(thetas, yfov, move_x, move_y, move_z, x_rot, y_rot, z_rot, frame_width, frame_height, draw_platform=True, depth_only=False, normals=False, smpl_model_path=None, shape_parameters=None, normalized_to_vertices=False):
if shape_parameters is not None:
betas_tensor = torch.tensor([shape_parameters], dtype=torch.float32)
batch_size = thetas.shape[3]
betas_batch = betas_tensor.repeat(batch_size, 1) # Replicates the single sample across the batch
betas_batch = betas_batch.to(device=get_torch_device())
else:
betas_batch = None
rot2xyz = Rotation2xyz(device=get_torch_device(), smpl_model_path=smpl_model_path, betas=betas_batch)
faces = rot2xyz.smpl_model.faces
vertices = rot2xyz(thetas.clone().to(get_torch_device()).detach(), mask=None,
pose_rep='xyz' if normalized_to_vertices else 'rot6d', translation=True, glob=True,
jointstype='vertices',
vertstrans=True)
frames = vertices.shape[3] # shape: 1, nb_frames, 3, nb_joints
MINS = torch.min(torch.min(vertices[0], axis=0)[0], axis=1)[0]
MAXS = torch.max(torch.max(vertices[0], axis=0)[0], axis=1)[0]
minx = MINS[0] - 0.5
maxx = MAXS[0] + 0.5
minz = MINS[2] - 0.5
maxz = MAXS[2] + 0.5
if draw_platform:
polygon = geometry.Polygon([[minx, minz], [minx, maxz], [maxx, maxz], [maxx, minz]])
polygon_mesh = trimesh.creation.extrude_polygon(polygon, 1e-5)
polygon_mesh.visual.face_colors = [0, 0, 0, 0.21]
polygon_render = pyrender.Mesh.from_trimesh(polygon_mesh, smooth=False)
c = np.pi / 2
platform_pose=np.array([[ 1, 0, 0, 0],
[ 0, np.cos(c), -np.sin(c), MINS[1].cpu().numpy()],
[ 0, np.sin(c), np.cos(c), 0],
[ 0, 0, 0, 1]])
base_color = (0.11, 0.53, 0.8, 0.5)
material = pyrender.MetallicRoughnessMaterial(
metallicFactor=0.7,
alphaMode='OPAQUE',
baseColorFactor=base_color
)
x_translation = move_x #X-axis translation value
y_translation = move_y # Y-axis translation value
z_translation = move_z # Z-axis translation value
initial_pos = [(minx+maxx).cpu().numpy()/2 + x_translation,
y_translation,
max(4, minz.cpu().numpy()+(1.5-MINS[1].cpu().numpy())*2, (maxx-minx).cpu().numpy()) + z_translation]
alpha = np.radians(x_rot)
beta = np.radians(y_rot)
gamma = np.radians(z_rot)
# Rotation matrix around X-axis
R_x = [[1, 0, 0, 0],
[0, np.cos(alpha), -np.sin(alpha), 0],
[0, np.sin(alpha), np.cos(alpha), 0],
[0, 0, 0, 1]]
# Rotation matrix around Y-axis
R_y = [[np.cos(beta), 0, np.sin(beta), 0],
[0, 1, 0, 0],
[-np.sin(beta), 0, np.cos(beta), 0],
[0, 0, 0, 1]]
# Rotation matrix around Z-axis
R_z = [[np.cos(gamma), -np.sin(gamma), 0, 0],
[np.sin(gamma), np.cos(gamma), 0, 0],
[0, 0, 1, 0],
[0, 0, 0, 1]]
# Combine rotations, order of multiplication depends on the desired rotation order
R = np.dot(R_z, np.dot(R_y, R_x))
# Now, R is a 4x4 matrix that represents the rotation around X, Y, and Z
# Translation vector
T = [initial_pos[0], initial_pos[1], initial_pos[2], 1]
# Combine the rotation and translation into the final transformation matrix
camera_pose = np.dot(R, np.array([[1, 0, 0, T[0]],
[0, 1, 0, T[1]],
[0, 0, 1, T[2]],
[0, 0, 0, 1]]))
if normals and not depth_only:
r = pyrender.OffscreenRenderer(frame_width, frame_height)
r._renderer._program_cache = ShaderProgramCache(shader_dir=shader_dir)
else:
r = pyrender.OffscreenRenderer(frame_width, frame_height)
light = pyrender.DirectionalLight(color=[1,1,1], intensity=300)
light_positions = [
[0, -1, 1],
[0, 1, 1],
[1, 1, 2]
]
# Create transformation matrices for each light
light_poses = [np.eye(4) for _ in light_positions]
for i, position in enumerate(light_positions):
light_poses[i][:3, 3] = position
#Build the scene
camera = pyrender.PerspectiveCamera(yfov)
bg_color = [1, 1, 1, 0.8]
scene = pyrender.Scene(bg_color=bg_color, ambient_light=(0.4, 0.4, 0.4))
scene.add(camera, pose=camera_pose)
if draw_platform:
scene.add(polygon_render, pose=platform_pose)
if not normals:
for pose in light_poses:
scene.add(light, pose=pose)
# Render loop
vid = []
vid_depth = []
print("Rendering SMPL human mesh...")
pbar = comfy.utils.ProgressBar(frames)
for i in tqdm(range(frames)):
mesh = Trimesh(vertices=vertices[0, :, :, i].squeeze().tolist(), faces=faces)
mesh = pyrender.Mesh.from_trimesh(mesh, material=material)
mesh_node = pyrender.Node(mesh=mesh)
scene.add_node(mesh_node)
if depth_only:
depth = r.render(scene, flags=RenderFlags.DEPTH_ONLY)
color = np.zeros([frame_width, frame_height, 3])
else:
color, depth = r.render(scene, flags=RenderFlags.RGBA)
vid.append(color)
vid_depth.append(depth)
scene.remove_node(mesh_node)
pbar.update(1)
r = None
return np.stack(vid, axis=0), np.stack(vid_depth, axis=0)
# verts_frames: list of [num_subjects, num_verts, 3]
# cam_t_frames: list of [num_subjects, 3]
def render_from_smpl_multiple_subjects(verts_frames, cam_t_frames, focal_length, fx_offset, fy_offset, move_x, move_y, move_z, x_rot, y_rot, z_rot, frame_width, frame_height, draw_platform=True, depth_only=False, normals=False, smpl_model_path=None):
def vertices_to_trimesh(vertices, camera_translation, faces, rot_axis=[1,0,0], rot_angle=0,):
mesh = trimesh.Trimesh(vertices + camera_translation, faces.copy())
rot = trimesh.transformations.rotation_matrix(
np.radians(rot_angle), rot_axis)
mesh.apply_transform(rot)
rot = trimesh.transformations.rotation_matrix(
np.radians(180), [1, 0, 0])
mesh.apply_transform(rot)
return mesh
rot2xyz = Rotation2xyz(device="cpu", smpl_model_path=smpl_model_path)
faces = rot2xyz.smpl_model.faces
MINS = torch.stack([verts_frame.min(0).values.min(0).values for verts_frame in verts_frames]).min(0).values
MAXS = torch.stack([verts_frame.max(0).values.max(0).values for verts_frame in verts_frames]).max(0).values
minx = MINS[0] - 0.5
maxx = MAXS[0] + 0.5
minz = MINS[2] - 0.5
maxz = MAXS[2] + 0.5
if draw_platform:
polygon = geometry.Polygon([[minx, minz], [minx, maxz], [maxx, maxz], [maxx, minz]])
polygon_mesh = trimesh.creation.extrude_polygon(polygon, 1e-5)
polygon_mesh.visual.face_colors = [0, 0, 0, 0.21]
polygon_render = pyrender.Mesh.from_trimesh(polygon_mesh, smooth=False)
c = np.pi / 2
platform_pose=np.array([[ 1, 0, 0, 0],
[ 0, np.cos(c), -np.sin(c), MINS[1].cpu().numpy()],
[ 0, np.sin(c), np.cos(c), 0],
[ 0, 0, 0, 1]])
base_color = (0.11, 0.53, 0.8, 0.5)
material = pyrender.MetallicRoughnessMaterial(
metallicFactor=0.7,
alphaMode='OPAQUE',
baseColorFactor=base_color
)
x_translation = move_x #X-axis translation value
y_translation = move_y # Y-axis translation value
z_translation = move_z # Z-axis translation value
initial_pos = [(minx+maxx).cpu().numpy()/2 + x_translation,
y_translation,
max(4, minz.cpu().numpy()+(1.5-MINS[1].cpu().numpy())*2, (maxx-minx).cpu().numpy()) + z_translation]
alpha = np.radians(x_rot)
beta = np.radians(y_rot)
gamma = np.radians(z_rot)
# Rotation matrix around X-axis
R_x = [[1, 0, 0, 0],
[0, np.cos(alpha), -np.sin(alpha), 0],
[0, np.sin(alpha), np.cos(alpha), 0],
[0, 0, 0, 1]]
# Rotation matrix around Y-axis
R_y = [[np.cos(beta), 0, np.sin(beta), 0],
[0, 1, 0, 0],
[-np.sin(beta), 0, np.cos(beta), 0],
[0, 0, 0, 1]]
# Rotation matrix around Z-axis
R_z = [[np.cos(gamma), -np.sin(gamma), 0, 0],
[np.sin(gamma), np.cos(gamma), 0, 0],
[0, 0, 1, 0],
[0, 0, 0, 1]]
# Combine rotations, order of multiplication depends on the desired rotation order
R = np.dot(R_z, np.dot(R_y, R_x))
# Now, R is a 4x4 matrix that represents the rotation around X, Y, and Z
# Translation vector
T = [initial_pos[0], initial_pos[1], initial_pos[2], 1]
# Combine the rotation and translation into the final transformation matrix
camera_pose = np.dot(R, np.array([[1, 0, 0, T[0]],
[0, 1, 0, T[1]],
[0, 0, 1, T[2]],
[0, 0, 0, 1]]))
if normals and not depth_only:
r = pyrender.OffscreenRenderer(frame_width, frame_height)
r._renderer._program_cache = ShaderProgramCache(shader_dir=shader_dir)
else:
r = pyrender.OffscreenRenderer(frame_width, frame_height)
light = pyrender.DirectionalLight(color=[1,1,1], intensity=300)
light_positions = [
[0, -1, 1],
[0, 1, 1],
[1, 1, 2]
]
# Create transformation matrices for each light
light_poses = [np.eye(4) for _ in light_positions]
for i, position in enumerate(light_positions):
light_poses[i][:3, 3] = position
#Build the scene
camera = pyrender.IntrinsicsCamera(fx=focal_length + fx_offset, fy=focal_length + fy_offset,
cx=frame_width / 2, cy=frame_height / 2, zfar=1e12)
bg_color = [1, 1, 1, 0.8]
scene = pyrender.Scene(bg_color=bg_color, ambient_light=(0.4, 0.4, 0.4))
scene.add(camera, pose=camera_pose)
if draw_platform:
scene.add(polygon_render, pose=platform_pose)
if not normals:
for pose in light_poses:
scene.add(light, pose=pose)
# Render loop
vid = []
vid_depth = []
print("Rendering SMPL human mesh...")
pbar = comfy.utils.ProgressBar(len(verts_frames))
for i in tqdm(range(len(verts_frames))):
subjects = verts_frames[i]
cam_t_subjects = cam_t_frames[i]
mesh_nodes = []
for subject_vertices, cam_t in zip(subjects, cam_t_subjects):
mesh = vertices_to_trimesh(subject_vertices, cam_t, faces)
mesh = pyrender.Mesh.from_trimesh(mesh, material=material)
mesh_node = pyrender.Node(mesh=mesh)
scene.add_node(mesh_node)
mesh_nodes.append(mesh_node)
if depth_only:
depth = r.render(scene, flags=RenderFlags.DEPTH_ONLY)
color = np.zeros([frame_width, frame_height, 3])
else:
color, depth = r.render(scene, flags=RenderFlags.RGBA)
vid.append(color)
vid_depth.append(depth)
for mesh_node in mesh_nodes: scene.remove_node(mesh_node)
pbar.update(1)
r = None
return np.stack(vid, axis=0), np.stack(vid_depth, axis=0)
@@ -0,0 +1,92 @@
# This code is based on https://github.com/Mathux/ACTOR.git
import torch
from motiondiff_modules.mogen.smpl import rotation_conversions as geometry
from .smpl import SMPL, JOINTSTYPE_ROOT
# from .get_model import JOINTSTYPES
JOINTSTYPES = ["a2m", "a2mpl", "smpl", "vibe", "vertices"]
class Rotation2xyz:
def __init__(self, device, dataset='amass', smpl_model_path=None, betas=None):
self.device = device
self.dataset = dataset
self.smpl_model = SMPL(smpl_model_path).eval().to(device)
self.betas = betas
def __call__(self, x, mask, pose_rep, translation, glob,
jointstype, vertstrans, beta=0,
glob_rot=None, get_rotations_back=False, **kwargs):
if pose_rep == "xyz":
return x
if mask is None:
mask = torch.ones((x.shape[0], x.shape[-1]), dtype=bool, device=x.device)
if not glob and glob_rot is None:
raise TypeError("You must specify global rotation if glob is False")
if jointstype not in JOINTSTYPES:
raise NotImplementedError("This jointstype is not implemented.")
if translation:
x_translations = x[:, -1, :3]
x_rotations = x[:, :-1]
else:
x_rotations = x
x_rotations = x_rotations.permute(0, 3, 1, 2)
nsamples, time, njoints, feats = x_rotations.shape
# Compute rotations (convert only masked sequences output)
if pose_rep == "rotvec":
rotations = geometry.axis_angle_to_matrix(x_rotations[mask])
elif pose_rep == "rotmat":
rotations = x_rotations[mask].view(-1, njoints, 3, 3)
elif pose_rep == "rotquat":
rotations = geometry.quaternion_to_matrix(x_rotations[mask])
elif pose_rep == "rot6d":
rotations = geometry.rotation_6d_to_matrix(x_rotations[mask])
else:
raise NotImplementedError("No geometry for this one.")
if not glob:
global_orient = torch.tensor(glob_rot, device=x.device)
global_orient = geometry.axis_angle_to_matrix(global_orient).view(1, 1, 3, 3)
global_orient = global_orient.repeat(len(rotations), 1, 1, 1)
else:
global_orient = rotations[:, 0]
rotations = rotations[:, 1:]
if self.betas is None:
self.betas = torch.zeros([rotations.shape[0], self.smpl_model.num_betas],
dtype=rotations.dtype, device=rotations.device)
self.betas[:, 1] = beta
# import ipdb; ipdb.set_trace()
out = self.smpl_model(body_pose=rotations, global_orient=global_orient, betas=self.betas)
# get the desirable joints
joints = out[jointstype]
x_xyz = torch.empty(nsamples, time, joints.shape[1], 3, device=x.device, dtype=x.dtype)
x_xyz[~mask] = 0
x_xyz[mask] = joints
x_xyz = x_xyz.permute(0, 2, 3, 1).contiguous()
# the first translation root at the origin on the prediction
if jointstype != "vertices":
rootindex = JOINTSTYPE_ROOT[jointstype]
x_xyz = x_xyz - x_xyz[:, [rootindex], :, :]
if translation and vertstrans:
# the first translation root at the origin
x_translations = x_translations - x_translations[:, :, [0]]
# add the translation to all the joints
x_xyz = x_xyz + x_translations[:, None, :, :]
if get_rotations_back:
return x_xyz, rotations, global_orient
else:
return x_xyz
@@ -0,0 +1,532 @@
# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved.
# Check PYTORCH3D_LICENCE before use
import functools
from typing import Optional
import torch
import torch.nn.functional as F
"""
The transformation matrices returned from the functions in this file assume
the points on which the transformation will be applied are column vectors.
i.e. the R matrix is structured as
R = [
[Rxx, Rxy, Rxz],
[Ryx, Ryy, Ryz],
[Rzx, Rzy, Rzz],
] # (3, 3)
This matrix can be applied to column vectors by post multiplication
by the points e.g.
points = [[0], [1], [2]] # (3 x 1) xyz coordinates of a point
transformed_points = R * points
To apply the same matrix to points which are row vectors, the R matrix
can be transposed and pre multiplied by the points:
e.g.
points = [[0, 1, 2]] # (1 x 3) xyz coordinates of a point
transformed_points = points * R.transpose(1, 0)
"""
def quaternion_to_matrix(quaternions):
"""
Convert rotations given as quaternions to rotation matrices.
Args:
quaternions: quaternions with real part first,
as tensor of shape (..., 4).
Returns:
Rotation matrices as tensor of shape (..., 3, 3).
"""
r, i, j, k = torch.unbind(quaternions, -1)
two_s = 2.0 / (quaternions * quaternions).sum(-1)
o = torch.stack(
(
1 - two_s * (j * j + k * k),
two_s * (i * j - k * r),
two_s * (i * k + j * r),
two_s * (i * j + k * r),
1 - two_s * (i * i + k * k),
two_s * (j * k - i * r),
two_s * (i * k - j * r),
two_s * (j * k + i * r),
1 - two_s * (i * i + j * j),
),
-1,
)
return o.reshape(quaternions.shape[:-1] + (3, 3))
def _copysign(a, b):
"""
Return a tensor where each element has the absolute value taken from the,
corresponding element of a, with sign taken from the corresponding
element of b. This is like the standard copysign floating-point operation,
but is not careful about negative 0 and NaN.
Args:
a: source tensor.
b: tensor whose signs will be used, of the same shape as a.
Returns:
Tensor of the same shape as a with the signs of b.
"""
signs_differ = (a < 0) != (b < 0)
return torch.where(signs_differ, -a, a)
def _sqrt_positive_part(x):
"""
Returns torch.sqrt(torch.max(0, x))
but with a zero subgradient where x is 0.
"""
ret = torch.zeros_like(x)
positive_mask = x > 0
ret[positive_mask] = torch.sqrt(x[positive_mask])
return ret
def matrix_to_quaternion(matrix):
"""
Convert rotations given as rotation matrices to quaternions.
Args:
matrix: Rotation matrices as tensor of shape (..., 3, 3).
Returns:
quaternions with real part first, as tensor of shape (..., 4).
"""
if matrix.size(-1) != 3 or matrix.size(-2) != 3:
raise ValueError(f"Invalid rotation matrix shape f{matrix.shape}.")
m00 = matrix[..., 0, 0]
m11 = matrix[..., 1, 1]
m22 = matrix[..., 2, 2]
o0 = 0.5 * _sqrt_positive_part(1 + m00 + m11 + m22)
x = 0.5 * _sqrt_positive_part(1 + m00 - m11 - m22)
y = 0.5 * _sqrt_positive_part(1 - m00 + m11 - m22)
z = 0.5 * _sqrt_positive_part(1 - m00 - m11 + m22)
o1 = _copysign(x, matrix[..., 2, 1] - matrix[..., 1, 2])
o2 = _copysign(y, matrix[..., 0, 2] - matrix[..., 2, 0])
o3 = _copysign(z, matrix[..., 1, 0] - matrix[..., 0, 1])
return torch.stack((o0, o1, o2, o3), -1)
def _axis_angle_rotation(axis: str, angle):
"""
Return the rotation matrices for one of the rotations about an axis
of which Euler angles describe, for each value of the angle given.
Args:
axis: Axis label "X" or "Y or "Z".
angle: any shape tensor of Euler angles in radians
Returns:
Rotation matrices as tensor of shape (..., 3, 3).
"""
cos = torch.cos(angle)
sin = torch.sin(angle)
one = torch.ones_like(angle)
zero = torch.zeros_like(angle)
if axis == "X":
R_flat = (one, zero, zero, zero, cos, -sin, zero, sin, cos)
if axis == "Y":
R_flat = (cos, zero, sin, zero, one, zero, -sin, zero, cos)
if axis == "Z":
R_flat = (cos, -sin, zero, sin, cos, zero, zero, zero, one)
return torch.stack(R_flat, -1).reshape(angle.shape + (3, 3))
def euler_angles_to_matrix(euler_angles, convention: str):
"""
Convert rotations given as Euler angles in radians to rotation matrices.
Args:
euler_angles: Euler angles in radians as tensor of shape (..., 3).
convention: Convention string of three uppercase letters from
{"X", "Y", and "Z"}.
Returns:
Rotation matrices as tensor of shape (..., 3, 3).
"""
if euler_angles.dim() == 0 or euler_angles.shape[-1] != 3:
raise ValueError("Invalid input euler angles.")
if len(convention) != 3:
raise ValueError("Convention must have 3 letters.")
if convention[1] in (convention[0], convention[2]):
raise ValueError(f"Invalid convention {convention}.")
for letter in convention:
if letter not in ("X", "Y", "Z"):
raise ValueError(f"Invalid letter {letter} in convention string.")
matrices = map(_axis_angle_rotation, convention, torch.unbind(euler_angles, -1))
return functools.reduce(torch.matmul, matrices)
def _angle_from_tan(
axis: str, other_axis: str, data, horizontal: bool, tait_bryan: bool
):
"""
Extract the first or third Euler angle from the two members of
the matrix which are positive constant times its sine and cosine.
Args:
axis: Axis label "X" or "Y or "Z" for the angle we are finding.
other_axis: Axis label "X" or "Y or "Z" for the middle axis in the
convention.
data: Rotation matrices as tensor of shape (..., 3, 3).
horizontal: Whether we are looking for the angle for the third axis,
which means the relevant entries are in the same row of the
rotation matrix. If not, they are in the same column.
tait_bryan: Whether the first and third axes in the convention differ.
Returns:
Euler Angles in radians for each matrix in data as a tensor
of shape (...).
"""
i1, i2 = {"X": (2, 1), "Y": (0, 2), "Z": (1, 0)}[axis]
if horizontal:
i2, i1 = i1, i2
even = (axis + other_axis) in ["XY", "YZ", "ZX"]
if horizontal == even:
return torch.atan2(data[..., i1], data[..., i2])
if tait_bryan:
return torch.atan2(-data[..., i2], data[..., i1])
return torch.atan2(data[..., i2], -data[..., i1])
def _index_from_letter(letter: str):
if letter == "X":
return 0
if letter == "Y":
return 1
if letter == "Z":
return 2
def matrix_to_euler_angles(matrix, convention: str):
"""
Convert rotations given as rotation matrices to Euler angles in radians.
Args:
matrix: Rotation matrices as tensor of shape (..., 3, 3).
convention: Convention string of three uppercase letters.
Returns:
Euler angles in radians as tensor of shape (..., 3).
"""
if len(convention) != 3:
raise ValueError("Convention must have 3 letters.")
if convention[1] in (convention[0], convention[2]):
raise ValueError(f"Invalid convention {convention}.")
for letter in convention:
if letter not in ("X", "Y", "Z"):
raise ValueError(f"Invalid letter {letter} in convention string.")
if matrix.size(-1) != 3 or matrix.size(-2) != 3:
raise ValueError(f"Invalid rotation matrix shape f{matrix.shape}.")
i0 = _index_from_letter(convention[0])
i2 = _index_from_letter(convention[2])
tait_bryan = i0 != i2
if tait_bryan:
central_angle = torch.asin(
matrix[..., i0, i2] * (-1.0 if i0 - i2 in [-1, 2] else 1.0)
)
else:
central_angle = torch.acos(matrix[..., i0, i0])
o = (
_angle_from_tan(
convention[0], convention[1], matrix[..., i2], False, tait_bryan
),
central_angle,
_angle_from_tan(
convention[2], convention[1], matrix[..., i0, :], True, tait_bryan
),
)
return torch.stack(o, -1)
def random_quaternions(
n: int, dtype: Optional[torch.dtype] = None, device=None, requires_grad=False
):
"""
Generate random quaternions representing rotations,
i.e. versors with nonnegative real part.
Args:
n: Number of quaternions in a batch to return.
dtype: Type to return.
device: Desired device of returned tensor. Default:
uses the current device for the default tensor type.
requires_grad: Whether the resulting tensor should have the gradient
flag set.
Returns:
Quaternions as tensor of shape (N, 4).
"""
o = torch.randn((n, 4), dtype=dtype, device=device, requires_grad=requires_grad)
s = (o * o).sum(1)
o = o / _copysign(torch.sqrt(s), o[:, 0])[:, None]
return o
def random_rotations(
n: int, dtype: Optional[torch.dtype] = None, device=None, requires_grad=False
):
"""
Generate random rotations as 3x3 rotation matrices.
Args:
n: Number of rotation matrices in a batch to return.
dtype: Type to return.
device: Device of returned tensor. Default: if None,
uses the current device for the default tensor type.
requires_grad: Whether the resulting tensor should have the gradient
flag set.
Returns:
Rotation matrices as tensor of shape (n, 3, 3).
"""
quaternions = random_quaternions(
n, dtype=dtype, device=device, requires_grad=requires_grad
)
return quaternion_to_matrix(quaternions)
def random_rotation(
dtype: Optional[torch.dtype] = None, device=None, requires_grad=False
):
"""
Generate a single random 3x3 rotation matrix.
Args:
dtype: Type to return
device: Device of returned tensor. Default: if None,
uses the current device for the default tensor type
requires_grad: Whether the resulting tensor should have the gradient
flag set
Returns:
Rotation matrix as tensor of shape (3, 3).
"""
return random_rotations(1, dtype, device, requires_grad)[0]
def standardize_quaternion(quaternions):
"""
Convert a unit quaternion to a standard form: one in which the real
part is non negative.
Args:
quaternions: Quaternions with real part first,
as tensor of shape (..., 4).
Returns:
Standardized quaternions as tensor of shape (..., 4).
"""
return torch.where(quaternions[..., 0:1] < 0, -quaternions, quaternions)
def quaternion_raw_multiply(a, b):
"""
Multiply two quaternions.
Usual torch rules for broadcasting apply.
Args:
a: Quaternions as tensor of shape (..., 4), real part first.
b: Quaternions as tensor of shape (..., 4), real part first.
Returns:
The product of a and b, a tensor of quaternions shape (..., 4).
"""
aw, ax, ay, az = torch.unbind(a, -1)
bw, bx, by, bz = torch.unbind(b, -1)
ow = aw * bw - ax * bx - ay * by - az * bz
ox = aw * bx + ax * bw + ay * bz - az * by
oy = aw * by - ax * bz + ay * bw + az * bx
oz = aw * bz + ax * by - ay * bx + az * bw
return torch.stack((ow, ox, oy, oz), -1)
def quaternion_multiply(a, b):
"""
Multiply two quaternions representing rotations, returning the quaternion
representing their composition, i.e. the versor with nonnegative real part.
Usual torch rules for broadcasting apply.
Args:
a: Quaternions as tensor of shape (..., 4), real part first.
b: Quaternions as tensor of shape (..., 4), real part first.
Returns:
The product of a and b, a tensor of quaternions of shape (..., 4).
"""
ab = quaternion_raw_multiply(a, b)
return standardize_quaternion(ab)
def quaternion_invert(quaternion):
"""
Given a quaternion representing rotation, get the quaternion representing
its inverse.
Args:
quaternion: Quaternions as tensor of shape (..., 4), with real part
first, which must be versors (unit quaternions).
Returns:
The inverse, a tensor of quaternions of shape (..., 4).
"""
return quaternion * quaternion.new_tensor([1, -1, -1, -1])
def quaternion_apply(quaternion, point):
"""
Apply the rotation given by a quaternion to a 3D point.
Usual torch rules for broadcasting apply.
Args:
quaternion: Tensor of quaternions, real part first, of shape (..., 4).
point: Tensor of 3D points of shape (..., 3).
Returns:
Tensor of rotated points of shape (..., 3).
"""
if point.size(-1) != 3:
raise ValueError(f"Points are not in 3D, f{point.shape}.")
real_parts = point.new_zeros(point.shape[:-1] + (1,))
point_as_quaternion = torch.cat((real_parts, point), -1)
out = quaternion_raw_multiply(
quaternion_raw_multiply(quaternion, point_as_quaternion),
quaternion_invert(quaternion),
)
return out[..., 1:]
def axis_angle_to_matrix(axis_angle):
"""
Convert rotations given as axis/angle to rotation matrices.
Args:
axis_angle: Rotations given as a vector in axis angle form,
as a tensor of shape (..., 3), where the magnitude is
the angle turned anticlockwise in radians around the
vector's direction.
Returns:
Rotation matrices as tensor of shape (..., 3, 3).
"""
return quaternion_to_matrix(axis_angle_to_quaternion(axis_angle))
def matrix_to_axis_angle(matrix):
"""
Convert rotations given as rotation matrices to axis/angle.
Args:
matrix: Rotation matrices as tensor of shape (..., 3, 3).
Returns:
Rotations given as a vector in axis angle form, as a tensor
of shape (..., 3), where the magnitude is the angle
turned anticlockwise in radians around the vector's
direction.
"""
return quaternion_to_axis_angle(matrix_to_quaternion(matrix))
def axis_angle_to_quaternion(axis_angle):
"""
Convert rotations given as axis/angle to quaternions.
Args:
axis_angle: Rotations given as a vector in axis angle form,
as a tensor of shape (..., 3), where the magnitude is
the angle turned anticlockwise in radians around the
vector's direction.
Returns:
quaternions with real part first, as tensor of shape (..., 4).
"""
angles = torch.norm(axis_angle, p=2, dim=-1, keepdim=True)
half_angles = 0.5 * angles
eps = 1e-6
small_angles = angles.abs() < eps
sin_half_angles_over_angles = torch.empty_like(angles)
sin_half_angles_over_angles[~small_angles] = (
torch.sin(half_angles[~small_angles]) / angles[~small_angles]
)
# for x small, sin(x/2) is about x/2 - (x/2)^3/6
# so sin(x/2)/x is about 1/2 - (x*x)/48
sin_half_angles_over_angles[small_angles] = (
0.5 - (angles[small_angles] * angles[small_angles]) / 48
)
quaternions = torch.cat(
[torch.cos(half_angles), axis_angle * sin_half_angles_over_angles], dim=-1
)
return quaternions
def quaternion_to_axis_angle(quaternions):
"""
Convert rotations given as quaternions to axis/angle.
Args:
quaternions: quaternions with real part first,
as tensor of shape (..., 4).
Returns:
Rotations given as a vector in axis angle form, as a tensor
of shape (..., 3), where the magnitude is the angle
turned anticlockwise in radians around the vector's
direction.
"""
norms = torch.norm(quaternions[..., 1:], p=2, dim=-1, keepdim=True)
half_angles = torch.atan2(norms, quaternions[..., :1])
angles = 2 * half_angles
eps = 1e-6
small_angles = angles.abs() < eps
sin_half_angles_over_angles = torch.empty_like(angles)
sin_half_angles_over_angles[~small_angles] = (
torch.sin(half_angles[~small_angles]) / angles[~small_angles]
)
# for x small, sin(x/2) is about x/2 - (x/2)^3/6
# so sin(x/2)/x is about 1/2 - (x*x)/48
sin_half_angles_over_angles[small_angles] = (
0.5 - (angles[small_angles] * angles[small_angles]) / 48
)
return quaternions[..., 1:] / sin_half_angles_over_angles
def rotation_6d_to_matrix(d6: torch.Tensor) -> torch.Tensor:
"""
Converts 6D rotation representation by Zhou et al. [1] to rotation matrix
using Gram--Schmidt orthogonalisation per Section B of [1].
Args:
d6: 6D rotation representation, of size (*, 6)
Returns:
batch of rotation matrices of size (*, 3, 3)
[1] Zhou, Y., Barnes, C., Lu, J., Yang, J., & Li, H.
On the Continuity of Rotation Representations in Neural Networks.
IEEE Conference on Computer Vision and Pattern Recognition, 2019.
Retrieved from http://arxiv.org/abs/1812.07035
"""
a1, a2 = d6[..., :3], d6[..., 3:]
b1 = F.normalize(a1, dim=-1)
b2 = a2 - (b1 * a2).sum(-1, keepdim=True) * b1
b2 = F.normalize(b2, dim=-1)
b3 = torch.cross(b1, b2, dim=-1)
return torch.stack((b1, b2, b3), dim=-2)
def matrix_to_rotation_6d(matrix: torch.Tensor) -> torch.Tensor:
"""
Converts rotation matrices to 6D rotation representation by Zhou et al. [1]
by dropping the last row. Note that 6D representation is not unique.
Args:
matrix: batch of rotation matrices of size (*, 3, 3)
Returns:
6D rotation representation, of size (*, 6)
[1] Zhou, Y., Barnes, C., Lu, J., Yang, J., & Li, H.
On the Continuity of Rotation Representations in Neural Networks.
IEEE Conference on Computer Vision and Pattern Recognition, 2019.
Retrieved from http://arxiv.org/abs/1812.07035
"""
return matrix[..., :2, :].clone().reshape(*matrix.size()[:-2], 6)
def canonicalize_smplh(poses, trans = None):
bs, nframes, njoints = poses.shape[:3]
global_orient = poses[:, :, 0]
# first global rotations
rot2d = matrix_to_axis_angle(global_orient[:, 0])
#rot2d[:, :2] = 0 # Remove the rotation along the vertical axis
rot2d = axis_angle_to_matrix(rot2d)
# Rotate the global rotation to eliminate Z rotations
global_orient = torch.einsum("ikj,imkl->imjl", rot2d, global_orient)
# Construct canonicalized version of x
xc = torch.cat((global_orient[:, :, None], poses[:, :, 1:]), dim=2)
if trans is not None:
vel = trans[:, 1:] - trans[:, :-1]
# Turn the translation as well
vel = torch.einsum("ikj,ilk->ilj", rot2d, vel)
trans = torch.cat((torch.zeros(bs, 1, 3, device=vel.device),
torch.cumsum(vel, 1)), 1)
return xc, trans
else:
return xc
@@ -0,0 +1,13 @@
#version 330 core
in vec3 frag_position;
in vec3 frag_normal;
out vec4 frag_color;
void main()
{
vec3 normal = normalize(frag_normal);
frag_color = vec4(normal * 0.5 + 0.5, 1.0);
}
@@ -0,0 +1,25 @@
#version 330 core
// Vertex Attributes
layout(location = 0) in vec3 position;
layout(location = NORMAL_LOC) in vec3 normal;
layout(location = INST_M_LOC) in mat4 inst_m;
// Uniforms
uniform mat4 M; // Model matrix
uniform mat4 V; // View matrix
uniform mat4 P; // Projection matrix
// Outputs
out vec3 frag_position;
out vec3 frag_normal;
void main()
{
mat4 modelView = V * M * inst_m; // Compute model-view matrix
gl_Position = P * modelView * vec4(position, 1);
frag_position = vec3(modelView * vec4(position, 1.0)); // Position in camera space
mat3 normalMatrix = transpose(inverse(mat3(modelView))); // Normal matrix in camera space
frag_normal = normalize(normalMatrix * normal); // Transform normal to camera space
}
@@ -0,0 +1,135 @@
import numpy as np
import os
import torch
from .joints2smpl.src import config
import smplx
import h5py
from .joints2smpl.src.smplify import SMPLify3D
from tqdm import tqdm
from motiondiff_modules.mogen.smpl import rotation_conversions as geometry
import argparse
import comfy.utils
class joints2smpl:
def __init__(self, num_frames, device, num_smplify_iters=150, smplify_step_size=1e-2, fix_foot=False, smpl_model_path=config.SMPL_MODEL_DIR):
self.device = device
# self.device = torch.device("cpu")
self.batch_size = num_frames
self.num_joints = 22 # for HumanML3D
self.joint_category = "AMASS"
self.num_smplify_iters = num_smplify_iters
self.smplify_step_size = smplify_step_size
self.fix_foot = fix_foot
smplmodel = smplx.create(smpl_model_path,
model_type="smpl", gender="neutral", ext="pkl",
batch_size=self.batch_size).to(self.device)
# ## --- load the mean pose as original ----
smpl_mean_file = config.SMPL_MEAN_FILE
file = h5py.File(smpl_mean_file, 'r')
self.init_mean_pose = torch.from_numpy(file['pose'][:]).unsqueeze(0).repeat(self.batch_size, 1).float().to(self.device)
self.init_mean_shape = torch.from_numpy(file['shape'][:]).unsqueeze(0).repeat(self.batch_size, 1).float().to(self.device)
self.cam_trans_zero = torch.Tensor([0.0, 0.0, 0.0]).unsqueeze(0).to(self.device)
#
# # #-------------initialize SMPLify
self.smplify = SMPLify3D(smplxmodel=smplmodel,
batch_size=self.batch_size,
joints_category=self.joint_category,
num_iters=self.num_smplify_iters,
device=self.device,
step_size=self.smplify_step_size)
def npy2smpl(self, npy_path):
out_path = npy_path.replace('.npy', '_rot.npy')
motions = np.load(npy_path, allow_pickle=True)[None][0]
# print_batch('', motions)
n_samples = motions['motion'].shape[0]
all_thetas = []
pbar = comfy.utils.ProgressBar(n_samples)
for sample_i in tqdm(range(n_samples)):
thetas, _ = self.joint2smpl(motions['motion'][sample_i].transpose(2, 0, 1)) # [nframes, njoints, 3]
all_thetas.append(thetas.cpu().numpy())
pbar.update(1)
motions['motion'] = np.concatenate(all_thetas, axis=0)
print('motions', motions['motion'].shape)
print(f'Saving [{out_path}]')
np.save(out_path, motions)
exit()
def joint2smpl(self, input_joints, init_params=None):
_smplify = self.smplify # if init_params is None else self.smplify_fast
pred_pose = torch.zeros(self.batch_size, 72).to(self.device)
pred_betas = torch.zeros(self.batch_size, 10).to(self.device)
pred_cam_t = torch.zeros(self.batch_size, 3).to(self.device)
keypoints_3d = torch.zeros(self.batch_size, self.num_joints, 3).to(self.device)
# run the whole seqs
num_seqs = input_joints.shape[0]
# joints3d = input_joints[idx] # *1.2 #scale problem [check first]
keypoints_3d = torch.from_numpy(input_joints).to(self.device).float()
# if idx == 0:
if init_params is None:
pred_betas = self.init_mean_shape
pred_pose = self.init_mean_pose
pred_cam_t = self.cam_trans_zero
else:
pred_betas = init_params['betas']
pred_pose = init_params['pose']
pred_cam_t = init_params['cam']
if self.joint_category == "AMASS":
confidence_input = torch.ones(self.num_joints)
# make sure the foot and ankle
if self.fix_foot == True:
confidence_input[7] = 1.5
confidence_input[8] = 1.5
confidence_input[10] = 1.5
confidence_input[11] = 1.5
else:
print("Such category not settle down!")
self.confidence_input = confidence_input
new_opt_vertices, new_opt_joints, new_opt_pose, new_opt_betas, \
new_opt_cam_t, new_opt_joint_loss = _smplify(
pred_pose.detach(),
pred_betas.detach(),
pred_cam_t.detach(),
keypoints_3d,
conf_3d=confidence_input.to(self.device),
# seq_ind=idx
)
thetas = new_opt_pose.reshape(self.batch_size, 24, 3)
thetas = geometry.matrix_to_rotation_6d(geometry.axis_angle_to_matrix(thetas)) # [bs, 24, 6]
root_loc = keypoints_3d[:, 0].clone().detach() # [bs, 3]
root_loc = torch.cat([root_loc, torch.zeros_like(root_loc)], dim=-1).unsqueeze(1) # [bs, 1, 6]
thetas = torch.cat([thetas, root_loc], dim=1).unsqueeze(0).permute(0, 2, 3, 1) # [1, 25, 6, 196]
return thetas.clone().detach(), {'joints': new_opt_joints.clone().detach(), 'pose': new_opt_joints[0, :24].flatten().clone().detach(), 'betas': new_opt_betas.clone().detach(), 'cam': new_opt_cam_t.clone().detach()}
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument("--input_path", type=str, required=True, help='Blender file or dir with blender files')
parser.add_argument("--cuda", type=bool, default=True, help='')
parser.add_argument("--device", type=int, default=0, help='')
params = parser.parse_args()
simplify = joints2smpl(device_id=params.device, cuda=params.cuda)
if os.path.isfile(params.input_path) and params.input_path.endswith('.npy'):
simplify.npy2smpl(params.input_path)
elif os.path.isdir(params.input_path):
files = [os.path.join(params.input_path, f) for f in os.listdir(params.input_path) if f.endswith('.npy')]
for f in files:
simplify.npy2smpl(f)
+124
View File
@@ -0,0 +1,124 @@
# This code is based on https://github.com/Mathux/ACTOR.git
import numpy as np
import torch
import contextlib
from smplx import SMPLLayer as _SMPLLayer
from smplx.lbs import vertices2joints
# action2motion_joints = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 21, 24, 38]
# change 0 and 8
action2motion_joints = [8, 1, 2, 3, 4, 5, 6, 7, 0, 9, 10, 11, 12, 13, 14, 21, 24, 38]
#https://github.com/Mael-zys/T2M-GPT/blob/main/utils/config.py
import os
from pathlib import Path
SMPL_DATA_PATH = Path(__file__).parent / "body_models" / "smpl"
SMPL_KINTREE_PATH = os.path.join(SMPL_DATA_PATH, "kintree_table.pkl")
SMPL_MODEL_PATH = os.path.join(SMPL_DATA_PATH, "SMPL_NEUTRAL.pkl")
JOINT_REGRESSOR_TRAIN_EXTRA = os.path.join(SMPL_DATA_PATH, 'J_regressor_extra.npy')
ROT_CONVENTION_TO_ROT_NUMBER = {
'legacy': 23,
'no_hands': 21,
'full_hands': 51,
'mitten_hands': 33,
}
GENDERS = ['neutral', 'male', 'female']
NUM_BETAS = 10
JOINTSTYPE_ROOT = {"a2m": 0, # action2motion
"smpl": 0,
"a2mpl": 0, # set(smpl, a2m)
"vibe": 8} # 0 is the 8 position: OP MidHip below
JOINT_MAP = {
'OP Nose': 24, 'OP Neck': 12, 'OP RShoulder': 17,
'OP RElbow': 19, 'OP RWrist': 21, 'OP LShoulder': 16,
'OP LElbow': 18, 'OP LWrist': 20, 'OP MidHip': 0,
'OP RHip': 2, 'OP RKnee': 5, 'OP RAnkle': 8,
'OP LHip': 1, 'OP LKnee': 4, 'OP LAnkle': 7,
'OP REye': 25, 'OP LEye': 26, 'OP REar': 27,
'OP LEar': 28, 'OP LBigToe': 29, 'OP LSmallToe': 30,
'OP LHeel': 31, 'OP RBigToe': 32, 'OP RSmallToe': 33, 'OP RHeel': 34,
'Right Ankle': 8, 'Right Knee': 5, 'Right Hip': 45,
'Left Hip': 46, 'Left Knee': 4, 'Left Ankle': 7,
'Right Wrist': 21, 'Right Elbow': 19, 'Right Shoulder': 17,
'Left Shoulder': 16, 'Left Elbow': 18, 'Left Wrist': 20,
'Neck (LSP)': 47, 'Top of Head (LSP)': 48,
'Pelvis (MPII)': 49, 'Thorax (MPII)': 50,
'Spine (H36M)': 51, 'Jaw (H36M)': 52,
'Head (H36M)': 53, 'Nose': 24, 'Left Eye': 26,
'Right Eye': 25, 'Left Ear': 28, 'Right Ear': 27
}
JOINT_NAMES = [
'OP Nose', 'OP Neck', 'OP RShoulder',
'OP RElbow', 'OP RWrist', 'OP LShoulder',
'OP LElbow', 'OP LWrist', 'OP MidHip',
'OP RHip', 'OP RKnee', 'OP RAnkle',
'OP LHip', 'OP LKnee', 'OP LAnkle',
'OP REye', 'OP LEye', 'OP REar',
'OP LEar', 'OP LBigToe', 'OP LSmallToe',
'OP LHeel', 'OP RBigToe', 'OP RSmallToe', 'OP RHeel',
'Right Ankle', 'Right Knee', 'Right Hip',
'Left Hip', 'Left Knee', 'Left Ankle',
'Right Wrist', 'Right Elbow', 'Right Shoulder',
'Left Shoulder', 'Left Elbow', 'Left Wrist',
'Neck (LSP)', 'Top of Head (LSP)',
'Pelvis (MPII)', 'Thorax (MPII)',
'Spine (H36M)', 'Jaw (H36M)',
'Head (H36M)', 'Nose', 'Left Eye',
'Right Eye', 'Left Ear', 'Right Ear'
]
# adapted from VIBE/SPIN to output smpl_joints, vibe joints and action2motion joints
class SMPL(_SMPLLayer):
""" Extension of the official SMPL implementation to support more joints """
def __init__(self, model_path=SMPL_MODEL_PATH, **kwargs):
kwargs["model_path"] = model_path
# remove the verbosity for the 10-shapes beta parameters
with contextlib.redirect_stdout(None):
super(SMPL, self).__init__(**kwargs)
J_regressor_extra = np.load(JOINT_REGRESSOR_TRAIN_EXTRA)
self.register_buffer('J_regressor_extra', torch.tensor(J_regressor_extra, dtype=torch.float32))
vibe_indexes = np.array([JOINT_MAP[i] for i in JOINT_NAMES])
a2m_indexes = vibe_indexes[action2motion_joints]
smpl_indexes = np.arange(24)
a2mpl_indexes = np.unique(np.r_[smpl_indexes, a2m_indexes])
self.maps = {"vibe": vibe_indexes,
"a2m": a2m_indexes,
"smpl": smpl_indexes,
"a2mpl": a2mpl_indexes}
def forward(self, *args, **kwargs):
smpl_output = super(SMPL, self).forward(*args, **kwargs)
extra_joints = vertices2joints(self.J_regressor_extra, smpl_output.vertices)
all_joints = torch.cat([smpl_output.joints, extra_joints], dim=1)
output = {"vertices": smpl_output.vertices}
for joinstype, indexes in self.maps.items():
output[joinstype] = all_joints[:, indexes]
return output
@@ -0,0 +1,18 @@
from motiondiff_modules.mogen.utils.collect_env import collect_env
from motiondiff_modules.mogen.utils.dist_utils import DistOptimizerHook, allreduce_grads
from motiondiff_modules.mogen.utils.logger import get_root_logger
from motiondiff_modules.mogen.utils.misc import multi_apply, torch_to_numpy
from motiondiff_modules.mogen.utils.path_utils import (
Existence,
check_input_path,
check_path_existence,
check_path_suffix,
prepare_output_path,
)
__all__ = [
'collect_env', 'DistOptimizerHook', 'allreduce_grads', 'get_root_logger',
'multi_apply', 'torch_to_numpy', 'Existence', 'check_input_path',
'check_path_existence', 'check_path_suffix', 'prepare_output_path'
]
@@ -0,0 +1,16 @@
from custom_mmpkg.custom_mmcv.utils import collect_env as collect_base_env
from custom_mmpkg.custom_mmcv.utils import get_git_hash
import motiondiff_modules.mogen as mogen
def collect_env():
"""Collect the information of the running environments."""
env_info = collect_base_env()
env_info['mogen'] = mogen.__version__ + '+' + get_git_hash()[:7]
return env_info
if __name__ == '__main__':
for name, val in collect_env().items():
print(f'{name}: {val}')
@@ -0,0 +1,59 @@
from collections import OrderedDict
import torch.distributed as dist
from custom_mmpkg.custom_mmcv.runner import OptimizerHook
from torch._utils import (
_flatten_dense_tensors,
_take_tensors,
_unflatten_dense_tensors,
)
def _allreduce_coalesced(tensors, world_size, bucket_size_mb=-1):
if bucket_size_mb > 0:
bucket_size_bytes = bucket_size_mb * 1024 * 1024
buckets = _take_tensors(tensors, bucket_size_bytes)
else:
buckets = OrderedDict()
for tensor in tensors:
tp = tensor.type()
if tp not in buckets:
buckets[tp] = []
buckets[tp].append(tensor)
buckets = buckets.values()
for bucket in buckets:
flat_tensors = _flatten_dense_tensors(bucket)
dist.all_reduce(flat_tensors)
flat_tensors.div_(world_size)
for tensor, synced in zip(
bucket, _unflatten_dense_tensors(flat_tensors, bucket)):
tensor.copy_(synced)
def allreduce_grads(params, coalesce=True, bucket_size_mb=-1):
grads = [
param.grad.data for param in params
if param.requires_grad and param.grad is not None
]
world_size = dist.get_world_size()
if coalesce:
_allreduce_coalesced(grads, world_size, bucket_size_mb)
else:
for tensor in grads:
dist.all_reduce(tensor.div_(world_size))
class DistOptimizerHook(OptimizerHook):
def __init__(self, grad_clip=None, coalesce=True, bucket_size_mb=-1):
self.grad_clip = grad_clip
self.coalesce = coalesce
self.bucket_size_mb = bucket_size_mb
def after_train_iter(self, runner):
runner.optimizer.zero_grad()
runner.outputs['loss'].backward()
if self.grad_clip is not None:
self.clip_grads(runner.model.parameters())
runner.optimizer.step()
+7
View File
@@ -0,0 +1,7 @@
import logging
from custom_mmpkg.custom_mmcv.utils import get_logger
def get_root_logger(log_file=None, log_level=logging.INFO):
return get_logger('mogen', log_file, log_level)
+14
View File
@@ -0,0 +1,14 @@
from functools import partial
import torch
def multi_apply(func, *args, **kwargs):
pfunc = partial(func, **kwargs) if kwargs else func
map_results = map(pfunc, *args)
return tuple(map(list, zip(*map_results)))
def torch_to_numpy(x):
assert isinstance(x, torch.Tensor)
return x.detach().cpu().numpy()
@@ -0,0 +1,232 @@
import os
import warnings
from enum import Enum
from pathlib import Path
from typing import List, Union
try:
from typing import Literal
except ImportError:
from typing_extensions import Literal
def check_path_suffix(path_str: str,
allowed_suffix: Union[str, List[str]] = '') -> bool:
"""Check whether the suffix of the path is allowed.
Args:
path_str (str):
Path to check.
allowed_suffix (List[str], optional):
What extension names are allowed.
Offer a list like ['.jpg', ',jpeg'].
When it's [], all will be received.
Use [''] then directory is allowed.
Defaults to [].
Returns:
bool:
True: suffix test passed
False: suffix test failed
"""
if isinstance(allowed_suffix, str):
allowed_suffix = [allowed_suffix]
pathinfo = Path(path_str)
suffix = pathinfo.suffix.lower()
if len(allowed_suffix) == 0:
return True
if pathinfo.is_dir():
if '' in allowed_suffix:
return True
else:
return False
else:
for index, tmp_suffix in enumerate(allowed_suffix):
if not tmp_suffix.startswith('.'):
tmp_suffix = '.' + tmp_suffix
allowed_suffix[index] = tmp_suffix.lower()
if suffix in allowed_suffix:
return True
else:
return False
class Existence(Enum):
"""State of file existence."""
FileExist = 0
DirectoryExistEmpty = 1
DirectoryExistNotEmpty = 2
MissingParent = 3
DirectoryNotExist = 4
FileNotExist = 5
def check_path_existence(
path_str: str,
path_type: Literal['file', 'dir', 'auto'] = 'auto',
) -> Existence:
"""Check whether a file or a directory exists at the expected path.
Args:
path_str (str):
Path to check.
path_type (Literal[, optional):
What kind of file do we expect at the path.
Choose among `file`, `dir`, `auto`.
Defaults to 'auto'. path_type = path_type.lower()
Raises:
KeyError: if `path_type` conflicts with `path_str`
Returns:
Existence:
0. FileExist: file at path_str exists.
1. DirectoryExistEmpty: folder at path exists and.
2. DirectoryExistNotEmpty: folder at path_str exists and not empty.
3. MissingParent: its parent doesn't exist.
4. DirectoryNotExist: expect a folder at path_str, but not found.
5. FileNotExist: expect a file at path_str, but not found.
"""
path_type = path_type.lower()
assert path_type in {'file', 'dir', 'auto'}
pathinfo = Path(path_str)
if not pathinfo.parent.is_dir():
return Existence.MissingParent
suffix = pathinfo.suffix.lower()
if path_type == 'dir' or\
path_type == 'auto' and suffix == '':
if pathinfo.is_dir():
if len(os.listdir(path_str)) == 0:
return Existence.DirectoryExistEmpty
else:
return Existence.DirectoryExistNotEmpty
else:
return Existence.DirectoryNotExist
elif path_type == 'file' or\
path_type == 'auto' and suffix != '':
if pathinfo.is_file():
return Existence.FileExist
elif pathinfo.is_dir():
if len(os.listdir(path_str)) == 0:
return Existence.DirectoryExistEmpty
else:
return Existence.DirectoryExistNotEmpty
if path_str.endswith('/'):
return Existence.DirectoryNotExist
else:
return Existence.FileNotExist
def prepare_output_path(output_path: str,
allowed_suffix: List[str] = [],
tag: str = 'output file',
path_type: Literal['file', 'dir', 'auto'] = 'auto',
overwrite: bool = True) -> None:
"""Check output folder or file.
Args:
output_path (str): could be folder or file.
allowed_suffix (List[str], optional):
Check the suffix of `output_path`. If folder, should be [] or [''].
If could both be folder or file, should be [suffixs..., ''].
Defaults to [].
tag (str, optional): The `string` tag to specify the output type.
Defaults to 'output file'.
path_type (Literal[, optional):
Choose `file` for file and `dir` for folder.
Choose `auto` if allowed to be both.
Defaults to 'auto'.
overwrite (bool, optional):
Whether overwrite the existing file or folder.
Defaults to True.
Raises:
FileNotFoundError: suffix does not match.
FileExistsError: file or folder already exists and `overwrite` is
False.
Returns:
None
"""
if path_type.lower() == 'dir':
allowed_suffix = []
exist_result = check_path_existence(output_path, path_type=path_type)
if exist_result == Existence.MissingParent:
warnings.warn(
f'The parent folder of {tag} does not exist: {output_path},' +
f' will make dir {Path(output_path).parent.absolute().__str__()}')
os.makedirs(
Path(output_path).parent.absolute().__str__(), exist_ok=True)
elif exist_result == Existence.DirectoryNotExist:
os.mkdir(output_path)
print(f'Making directory {output_path} for saving results.')
elif exist_result == Existence.FileNotExist:
suffix_matched = \
check_path_suffix(output_path, allowed_suffix=allowed_suffix)
if not suffix_matched:
raise FileNotFoundError(
f'The {tag} should be {", ".join(allowed_suffix)}: '
f'{output_path}.')
elif exist_result == Existence.FileExist:
if not overwrite:
raise FileExistsError(
f'{output_path} exists (set overwrite = True to overwrite).')
else:
print(f'Overwriting {output_path}.')
elif exist_result == Existence.DirectoryExistEmpty:
pass
elif exist_result == Existence.DirectoryExistNotEmpty:
if not overwrite:
raise FileExistsError(
f'{output_path} is not empty (set overwrite = '
'True to overwrite the files).')
else:
print(f'Overwriting {output_path} and its files.')
else:
raise FileNotFoundError(f'No Existence type for {output_path}.')
def check_input_path(
input_path: str,
allowed_suffix: List[str] = [],
tag: str = 'input file',
path_type: Literal['file', 'dir', 'auto'] = 'auto',
):
"""Check input folder or file.
Args:
input_path (str): input folder or file path.
allowed_suffix (List[str], optional):
Check the suffix of `input_path`. If folder, should be [] or [''].
If could both be folder or file, should be [suffixs..., ''].
Defaults to [].
tag (str, optional): The `string` tag to specify the output type.
Defaults to 'output file'.
path_type (Literal[, optional):
Choose `file` for file and `directory` for folder.
Choose `auto` if allowed to be both.
Defaults to 'auto'.
Raises:
FileNotFoundError: file does not exists or suffix does not match.
Returns:
None
"""
if path_type.lower() == 'dir':
allowed_suffix = []
exist_result = check_path_existence(input_path, path_type=path_type)
if exist_result in [
Existence.FileExist, Existence.DirectoryExistEmpty,
Existence.DirectoryExistNotEmpty
]:
suffix_matched = \
check_path_suffix(input_path, allowed_suffix=allowed_suffix)
if not suffix_matched:
raise FileNotFoundError(
f'The {tag} should be {", ".join(allowed_suffix)}:' +
f'{input_path}.')
else:
raise FileNotFoundError(f'The {tag} does not exist: {input_path}.')
@@ -0,0 +1,251 @@
"""
This code is borrowed from https://github.com/EricGuo5513/text-to-motion
Modified to be compatible with matplotlib 3.8.0+
"""
import torch
import numpy as np
import math
import matplotlib
import matplotlib.pyplot as plt
from mpl_toolkits.mplot3d import Axes3D
from matplotlib.animation import FuncAnimation, PillowWriter
from mpl_toolkits.mplot3d.art3d import Poly3DCollection
import mpl_toolkits.mplot3d.axes3d as p3
from tempfile import TemporaryDirectory
from pathlib import Path
from uuid import uuid4
# Define a kinematic tree for the skeletal struture
kit_kinematic_chain = [[0, 11, 12, 13, 14, 15], [0, 16, 17, 18, 19, 20], [0, 1, 2, 3, 4], [3, 5, 6, 7], [3, 8, 9, 10]]
kit_raw_offsets = np.array(
[
[0, 0, 0],
[0, 1, 0],
[0, 1, 0],
[0, 1, 0],
[0, 1, 0],
[1, 0, 0],
[0, -1, 0],
[0, -1, 0],
[-1, 0, 0],
[0, -1, 0],
[0, -1, 0],
[1, 0, 0],
[0, -1, 0],
[0, -1, 0],
[0, 0, 1],
[0, 0, 1],
[-1, 0, 0],
[0, -1, 0],
[0, -1, 0],
[0, 0, 1],
[0, 0, 1]
]
)
t2m_raw_offsets = np.array([[0,0,0],
[1,0,0],
[-1,0,0],
[0,1,0],
[0,-1,0],
[0,-1,0],
[0,1,0],
[0,-1,0],
[0,-1,0],
[0,1,0],
[0,0,1],
[0,0,1],
[0,1,0],
[1,0,0],
[-1,0,0],
[0,0,1],
[0,-1,0],
[0,-1,0],
[0,-1,0],
[0,-1,0],
[0,-1,0],
[0,-1,0]])
t2m_kinematic_chain = [[0, 2, 5, 8, 11], [0, 1, 4, 7, 10], [0, 3, 6, 9, 12, 15], [9, 14, 17, 19, 21], [9, 13, 16, 18, 20]]
t2m_left_hand_chain = [[20, 22, 23, 24], [20, 34, 35, 36], [20, 25, 26, 27], [20, 31, 32, 33], [20, 28, 29, 30]]
t2m_right_hand_chain = [[21, 43, 44, 45], [21, 46, 47, 48], [21, 40, 41, 42], [21, 37, 38, 39], [21, 49, 50, 51]]
def qinv(q):
assert q.shape[-1] == 4, 'q must be a tensor of shape (*, 4)'
mask = torch.ones_like(q)
mask[..., 1:] = -mask[..., 1:]
return q * mask
def qrot(q, v):
"""
Rotate vector(s) v about the rotation described by quaternion(s) q.
Expects a tensor of shape (*, 4) for q and a tensor of shape (*, 3) for v,
where * denotes any number of dimensions.
Returns a tensor of shape (*, 3).
"""
assert q.shape[-1] == 4
assert v.shape[-1] == 3
assert q.shape[:-1] == v.shape[:-1]
original_shape = list(v.shape)
# print(q.shape)
q = q.contiguous().view(-1, 4)
v = v.contiguous().view(-1, 3)
qvec = q[:, 1:]
uv = torch.cross(qvec, v, dim=1)
uuv = torch.cross(qvec, uv, dim=1)
return (v + 2 * (q[:, :1] * uv + uuv)).view(original_shape)
def recover_root_rot_pos(data):
rot_vel = data[..., 0]
r_rot_ang = torch.zeros_like(rot_vel).to(data.device)
'''Get Y-axis rotation from rotation velocity'''
r_rot_ang[..., 1:] = rot_vel[..., :-1]
r_rot_ang = torch.cumsum(r_rot_ang, dim=-1)
r_rot_quat = torch.zeros(data.shape[:-1] + (4,)).to(data.device)
r_rot_quat[..., 0] = torch.cos(r_rot_ang)
r_rot_quat[..., 2] = torch.sin(r_rot_ang)
r_pos = torch.zeros(data.shape[:-1] + (3,)).to(data.device)
r_pos[..., 1:, [0, 2]] = data[..., :-1, 1:3]
'''Add Y-axis rotation to root position'''
r_pos = qrot(qinv(r_rot_quat), r_pos)
r_pos = torch.cumsum(r_pos, dim=-2)
r_pos[..., 1] = data[..., 3]
return r_rot_quat, r_pos
def recover_from_ric(data, joints_num):
r_rot_quat, r_pos = recover_root_rot_pos(data)
positions = data[..., 4:(joints_num - 1) * 3 + 4]
positions = positions.view(positions.shape[:-1] + (-1, 3))
'''Add Y-axis rotation to local joints'''
positions = qrot(qinv(r_rot_quat[..., None, :]).expand(positions.shape[:-1] + (4,)), positions)
'''Add root XZ to joints'''
positions[..., 0] += r_pos[..., 0:1]
positions[..., 2] += r_pos[..., 2:3]
'''Concate root and joints'''
positions = torch.cat([r_pos.unsqueeze(-2), positions], dim=-2)
return positions
def plot_3d_motion(save_path, kinematic_tree, joints, distance, elevation, rotation, poselinewidth, title, figsize=(10, 10), fps=120, radius=4, visualization="original", save_as_pil_lists=False):
matplotlib.use('Agg')
title_sp = title.split(' ')
if len(title_sp) > 20:
title = '\n'.join([' '.join(title_sp[:10]), ' '.join(title_sp[10:20]), ' '.join(title_sp[20:])])
elif len(title_sp) > 10:
title = '\n'.join([' '.join(title_sp[:10]), ' '.join(title_sp[10:])])
def init():
ax.set_xlim3d([-radius / 4, radius / 4])
ax.set_ylim3d([0, radius / 2])
ax.set_zlim3d([0, radius / 2])
if visualization == "original":
fig.suptitle(title, fontsize=20)
ax.grid(b=False)
fig.add_axes(ax) #As a user, you do not instantiate Axes directly, but use Axes creation methods instead; e.g. from pyplot or Figure: subplots, subplot_mosaic or Figure.add_axes.
if visualization == "pseudo-openpose":
ax.set_facecolor("black")
def plot_xzPlane(minx, maxx, miny, minz, maxz):
verts = [
[minx, miny, minz],
[minx, miny, maxz],
[maxx, miny, maxz],
[maxx, miny, minz]
]
xz_plane = Poly3DCollection([verts])
xz_plane.set_facecolor((0.5, 0.5, 0.5, 0.5))
ax.add_collection3d(xz_plane)
# (seq_len, joints_num, 3)
data = joints.copy().reshape(len(joints), -1, 3)
fig = plt.figure(figsize=figsize)
ax = p3.Axes3D(fig)
init()
MINS = data.min(axis=0).min(axis=0)
MAXS = data.max(axis=0).max(axis=0)
if visualization == "pseudo-openpose":
colors = [('#009900', '#009900', '#009933', '#009966'), #left leg
('#009999', '#009999', '#006699', '#003399'), #right leg
('#000099', ), #body
('#990000', '#990000', '#996600', '#999900'), #left arm
('#993300', '#993300', '#669900', '#339900'), #right arm
'darkblue', 'darkblue', 'darkblue', 'darkblue', 'darkblue',
'darkred', 'darkred', 'darkred', 'darkred', 'darkred']
else:
colors = ['red', 'blue', 'black', 'red', 'blue',
'darkblue', 'darkblue', 'darkblue', 'darkblue', 'darkblue',
'darkred', 'darkred', 'darkred', 'darkred', 'darkred']
frame_number = data.shape[0]
height_offset = MINS[1]
data[:, :, 1] -= height_offset
trajec = data[:, 0, [0, 2]]
data[..., 0] -= data[:, 0:1, 0]
data[..., 2] -= data[:, 0:1, 2]
def update(index):
for line in ax.lines:
line.remove()
for collection in ax.collections:
collection.remove()
ax.view_init(elev=elevation, azim=rotation)
ax._dist = distance #ax.dist = 7.5, https://matplotlib.org/stable/api/prev_api_changes/api_changes_3.8.0.html#axes3d
if visualization == "original":
plot_xzPlane(MINS[0] - trajec[index, 0], MAXS[0] - trajec[index, 0], 0, MINS[2] - trajec[index, 1],
MAXS[2] - trajec[index, 1])
if index > 1:
ax.plot3D(trajec[:index, 0] - trajec[index, 0], np.zeros_like(trajec[:index, 0]),
trajec[:index, 1] - trajec[index, 1], linewidth=1.0,
color='blue')
for i, (chain, color) in enumerate(zip(kinematic_tree, colors)):
if i < 5:
linewidth = 4.0
else:
linewidth = 2.0
if visualization == "pseudo-openpose":
for j in range(len(data[index, chain, 0]) - 1):
ax.plot3D(data[index, chain, 0][j:j+2], data[index, chain, 1][j:j+2], data[index, chain, 2][j:j+2], linewidth=poselinewidth,
color=color[j % len(color)])
else:
ax.plot3D(data[index, chain, 0], data[index, chain, 1], data[index, chain, 2], linewidth=linewidth,
color=color)
plt.axis('off')
ax.set_xticklabels([])
ax.set_yticklabels([])
ax.set_zticklabels([])
ani = FuncAnimation(fig, update, frames=frame_number, interval=1000 / fps, repeat=False)
if save_as_pil_lists:
pil_writer = PillowWriter()
with TemporaryDirectory() as tmpdir:
ani.save(Path(tmpdir, f"{uuid4()}.gif"), writer=pil_writer)
plt.close()
return pil_writer._frames
else:
ani.save(save_path, fps=fps)
plt.close()
+25
View File
@@ -0,0 +1,25 @@
__version__ = '0.0.1'
def parse_version_info(version_str):
"""Parse a version string into a tuple.
Args:
version_str (str): The version string.
Returns:
tuple[int | str]: The version info, e.g., "1.3.0" is parsed into
(1, 3, 0), and "2.0.0rc1" is parsed into (2, 0, 0, 'rc1').
"""
version_info = []
for x in version_str.split('.'):
if x.isdigit():
version_info.append(int(x))
elif x.find('rc') != -1:
patch_version = x.split('rc')
version_info.append(int(patch_version[0]))
version_info.append(f'rc{patch_version[1]}')
return tuple(version_info)
version_info = parse_version_info(__version__)
__all__ = ['__version__', 'version_info', 'parse_version_info']