Add 4DHuman
This commit is contained in:
@@ -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']
|
||||
@@ -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'
|
||||
]
|
||||
@@ -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
|
||||
@@ -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.
Binary file not shown.
Binary file not shown.
Binary file not shown.
|
After Width: | Height: | Size: 5.1 MiB |
Binary file not shown.
File diff suppressed because one or more lines are too long
Binary file not shown.
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)
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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']
|
||||
Reference in New Issue
Block a user