feat: init project

Signed-off-by: storyicon <storyicon@foxmail.com>
This commit is contained in:
storyicon
2024-05-21 08:01:08 +00:00
commit c792760074
1612 changed files with 723752 additions and 0 deletions
+134
View File
@@ -0,0 +1,134 @@
# Byte-compiled / optimized / DLL files
__pycache__/
*.py[cod]
*$py.class
# C extensions
*.so
# Distribution / packaging
.Python
build/
develop-eggs/
dist/
downloads/
eggs/
.eggs/
lib/
lib64/
parts/
sdist/
var/
wheels/
pip-wheel-metadata/
share/python-wheels/
*.egg-info/
.installed.cfg
*.egg
MANIFEST
# PyInstaller
# Usually these files are written by a python script from a template
# before PyInstaller builds the exe, so as to inject date/other infos into it.
*.manifest
*.spec
# Installer logs
pip-log.txt
pip-delete-this-directory.txt
# Unit test / coverage reports
htmlcov/
.tox/
.nox/
.coverage
.coverage.*
.cache
nosetests.xml
coverage.xml
*.cover
.hypothesis/
.pytest_cache/
# Translations
*.mo
*.pot
# Django stuff:
*.log
local_settings.py
db.sqlite3
# Flask stuff:
instance/
.webassets-cache
# Scrapy stuff:
.scrapy
# Sphinx documentation
docs/_build/
# PyBuilder
target/
# Jupyter Notebook
.ipynb_checkpoints
# IPython
profile_default/
ipython_config.py
# pyenv
.python-version
# pipenv
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
# However, in case of collaboration, if having platform-specific dependencies or dependencies
# having no cross-platform support, pipenv may install dependencies that don’t work, or not
# install all needed dependencies.
#Pipfile.lock
# celery beat schedule file
celerybeat-schedule
# SageMath parsed files
*.sage.py
# Environments
.env
.venv
env/
venv/
ENV/
env.bak/
venv.bak/
# Spyder project settings
.spyderproject
.spyproject
# Rope project settings
.ropeproject
# mkdocs documentation
/site
# mypy
.mypy_cache/
.dmypy.json
dmypy.json
# Pyre type checker
.pyre/
*.swp
.*.swp
.DS_Store
# project
outputs/
results/
scripts/codetest/
# configs/train/video_creation_anchorxia_*
+19
View File
@@ -0,0 +1,19 @@
FROM anchorxia/musev:1.0.0
#MAINTAINER 维护者信息
LABEL MAINTAINER="anchorxia"
LABEL Email="anchorxia@tencent.com"
LABEL Description="musev gpu runtime image, base docker is pytorch/pytorch:2.0.1-cuda11.7-cudnn8-devel"
ARG DEBIAN_FRONTEND=noninteractive
USER root
SHELL ["/bin/bash", "--login", "-c"]
RUN . /opt/conda/etc/profile.d/conda.sh \
&& echo "source activate musev" >> ~/.bashrc \
&& conda activate musev \
&& conda env list \
&& pip install cuid
USER root
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2024 TMElyralab
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+139
View File
@@ -0,0 +1,139 @@
# Byte-compiled / optimized / DLL files
__pycache__/
*.py[cod]
*$py.class
# C extensions
*.so
# Distribution / packaging
.Python
build/
develop-eggs/
dist/
downloads/
eggs/
.eggs/
lib/
lib64/
parts/
sdist/
var/
wheels/
pip-wheel-metadata/
share/python-wheels/
*.egg-info/
.installed.cfg
*.egg
MANIFEST
# PyInstaller
# Usually these files are written by a python script from a template
# before PyInstaller builds the exe, so as to inject date/other infos into it.
*.manifest
*.spec
# Installer logs
pip-log.txt
pip-delete-this-directory.txt
# Unit test / coverage reports
htmlcov/
.tox/
.nox/
.coverage
.coverage.*
.cache
nosetests.xml
coverage.xml
*.cover
.hypothesis/
.pytest_cache/
# Translations
*.mo
*.pot
# Django stuff:
*.log
local_settings.py
db.sqlite3
# Flask stuff:
instance/
.webassets-cache
# Scrapy stuff:
.scrapy
# Sphinx documentation
docs/_build/
# PyBuilder
target/
# Jupyter Notebook
.ipynb_checkpoints
# IPython
profile_default/
ipython_config.py
# pyenv
.python-version
# pipenv
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
# However, in case of collaboration, if having platform-specific dependencies or dependencies
# having no cross-platform support, pipenv may install dependencies that don’t work, or not
# install all needed dependencies.
#Pipfile.lock
# celery beat schedule file
celerybeat-schedule
# SageMath parsed files
*.sage.py
# Environments
.env
.venv
env/
venv/
ENV/
env.bak/
venv.bak/
# Spyder project settings
.spyderproject
.spyproject
# Rope project settings
.ropeproject
# mkdocs documentation
/site
# mypy
.mypy_cache/
.dmypy.json
dmypy.json
# Pyre type checker
.pyre/
*.swp
.*.swp
dataset/files
experiments
log
csvs
.idea
.vscode
__pycache__/
*.code-workspace
.DS_Store
third_party/
.polaris_cache/
*.lock
+83
View File
@@ -0,0 +1,83 @@
# FROM mirrors.tencent.com/todacc/venus-std-base-cuda11.8:0.1.0
FROM mirrors.tencent.com/todacc/venus-std-ext-cuda11.8-pytorch2.0-tf2.12-py3.10:0.7.0
#MAINTAINER 维护者信息
LABEL MAINTAINER="anchorxia"
LABEL Email="xzqjack@hotmail.com"
LABEL Description="gpu development image, from mirrors.tencent.com/todacc/venus-std-ext-cuda11.8-pytorch2.0-tf2.12-py3.10:0.7.0"
USER root
# 安装必须软件
# RUN GENERIC_REPO_URL="http://mirrors.tencent.com/repository/generic/venus_repo/image_res" \
# && cd /data/ \
# && wget -q $GENERIC_REPO_URL/gcc/gcc-11.2.0.zip \
# && unzip -q gcc-11.2.0.zip \
# && cd gcc-releases-gcc-11.2.0 \
# && ./contrib/download_prerequisites \
# && ./configure --enable-bootstrap --enable-languages=c,c++ --enable-threads=posix --enable-checking=release --enable-multilib --with-system-zlib \
# && make --silent -j10 \
# && make --silent install \
# && gcc -v \
# && rm -rf /data/gcc-releases-gcc-11.2.0 /data/gcc-11.2.0.zip
# RUN yum update -y \
# && yum install -y epel-release \
# && yum install -y ffmpeg \
# && yum install -y Xvfb \
# && yum install -y centos-release-scl devtoolset-11
RUN yum install -y wget zsh git curl tmux cmake htop iotop git-lfs zip \
&& yum install -y autojump autojump-zsh portaudio portaudio-devel \
&& yum clean all
USER mqq
RUN source ~/.bashrc \
&& GENERIC_REPO_URL="http://mirrors.tencent.com/repository/generic/venus_repo/image_res" \
&& conda deactivate \
# && conda remove -y -n env-2.7.18 --all \
# && conda remove -y -n env-3.6.8 --all \
# && conda remove -y -n env-3.7.7 --all \
# && conda remove -y -n env-3.8.8 --all \
# && conda remove -y -n env-3.9.2 --all \
# && conda remove -y -n env-novelai --all \
&& conda create -n projectv python=3.10.6 -y \
&& conda activate projectv \
&& pip install venus-sdk -q -i https://mirrors.tencent.com/repository/pypi/tencent_pypi/simple \
--extra-index-url https://mirrors.tencent.com/pypi/simple/ \
&& pip install tensorflow==2.12.0 tensorboard==2.12.0 \
&& pip install torch==2.0.1+cu118 torchvision==0.15.2+cu118 -f https://mirror.sjtu.edu.cn/pytorch-wheels/torch_stable.html -i https://mirrors.bfsu.edu.cn/pypi/web/simple -U \
# 安装xformers,支持不同型号gpu
&& pip install ninja==1.11.1 \
# && git clone https://github.com/facebookresearch/xformers.git \
# && cd xformers \
# && git checkout v0.0.17rc482 \
# && git submodule update --init --recursive \
# && pip install numpy==1.23.4 pyre-extensions==0.0.23 \
# && FORCE_CUDA="1" MAX_JOBS=1 TORCH_CUDA_ARCH_LIST="6.1;7.0;7.5;8.0;8.6" pip install -e . \
# && cd .. \
# 安装一堆包
&& pip install --no-cache-dir transformers bitsandbytes decord accelerate xformers omegaconf einops imageio==2.31.1 \
&& pip install --no-cache-dir pandas h5py matplotlib modelcards pynvml black pytest moviepy torch-tb-profiler scikit-learn librosa ffmpeg easydict webp controlnet_aux mediapipe \
&& pip install --no-cache-dir Cython easydict gdown infomap insightface ipython librosa onnx onnxruntime onnxsim opencv_python Pillow protobuf pytube PyYAML \
&& pip install --no-cache-dir requests scipy six tqdm gradio albumentations opencv-contrib-python imageio-ffmpeg pytorch-lightning test-tube \
&& pip install --no-cache-dir timm addict yapf prettytable safetensors basicsr fvcore pycocotools wandb gunicorn \
&& pip install --no-cache-dir streamlit webdataset kornia open_clip_torch streamlit-drawable-canvas torchmetrics \
# 安装暗水印
&& pip install --no-cache-dir invisible-watermark==0.1.5 gdown==4.5.3 ftfy==6.1.1 modelcards==0.1.6 \
# 安装openmm相关包
&& pip install--no-cache-dir -U openmim \
&& mim install mmengine \
&& mim install "mmcv>=2.0.1" \
&& mim install "mmdet>=3.1.0" \
&& mim install "mmpose>=1.1.0" \
# jupyters
&& pip install ipywidgets==8.0.3 \
&& python -m ipykernel install --user --name projectv --display-name "python(projectv)" \
&& pip install --no-cache-dir matplotlib==3.6.2 redis==4.5.1 pydantic[dotenv]==1.10.2 loguru==0.6.0 IProgress==0.4 \
&& pip install --no-cache-dir cos-python-sdk-v5==1.9.22 coscmd==1.8.6.30 \
# 必须放在最后pip,避免和jupyter的不兼容
&& pip install --no-cache-dir markupsafe==2.0.1 \
&& wget -P /tmp $GENERIC_REPO_URL/cpu/clean-layer.sh \
&& sh /tmp/clean-layer.sh
ENV LD_LIBRARY_PATH=/usr/local/lib64:$LD_LIBRARY_PATH
USER root
+2
View File
@@ -0,0 +1,2 @@
# MMCM
Process package for multi media, cross multi modal.
+6
View File
@@ -0,0 +1,6 @@
from .audio import *
from .data import *
from .music import *
from .text import *
from .vision import *
from .t2p import *
View File
+9
View File
@@ -0,0 +1,9 @@
from .general.items import Items, Item
from .emb.emb import MediaMapEmb
from .emb.h5py_emb import H5pyMediaMapEmb, H5pyMediaMapEmbProxy
from .media_map.media_map import MediaMap, MetaInfo, MetaInfoList, MediaMapSeq
from .media_map.media_map_process import get_sub_mediamap_by_clip_idx, get_sub_mediamap_by_stage, get_subseq_by_time
from .clip.clip import Clip, ClipSeq
from .clip.clipid import ClipIds, ClipIdsSeq, MatchedClipIds, MatchedClipIdsSeq
+324
View File
@@ -0,0 +1,324 @@
from copy import deepcopy
from typing import Iterable
import logging
import numpy as np
from ..utils.util import convert_class_attr_to_dict
logger = logging.getLogger(__name__) # pylint: disable=invalid-name
class Clip(object, Item):
"""媒体片段, 指转场点与转场点之间的部分"""
def __init__(
self,
time_start,
duration,
clipid=None,
media_type=None,
mediaid=None,
timepoint_type=None,
text=None,
stage=None,
path=None,
duration_num=None,
group_time_start=0,
group_clipid=None,
original_clipid=None,
emb=None,
multi_factor=None,
similar_clipseq=None,
rythm: float = None,
**kwargs
):
"""
Args:
time_start (float): 开始时间,秒为单位,对应该媒体文件的, 和media_map.json上的序号一一对应
duration (_type_): 片段持续时间
clipid (int, or [int]): 由media_map提供的片段序号, 和media_map.json上的序号一一对应
media_type (str, optional): music, video,text, Defaults to None.
mediaid (int): 多媒体id, 当clipid是列表时,表示该片段是个融合片段
timepoint_type(int, ): 开始点的转场类型. Defaults to None.
text(str, optional): 该片段的文本描述,音乐可以是歌词,视频可以是台词,甚至可以是弹幕. Defaults to None.
stage(str, optional): 该片段在整个媒体文件中的结构位置,如音乐的intro、chrous、vesa,视频的片头、片尾、开始、高潮、转场等. Defaults to None.
path (_type_, optional): 该媒体文件的路径,用于后续媒体读取、处理. Defaults to None.
duration_num (_type_, optional): 片段持续帧数, Defaults to None.
group_time_start (int, optional): 当多歌曲、多视频剪辑时,group_time_start 表示该片段所对应的子媒体前所有子媒体的片段时长总和。
默认0, 表示只有1个媒体文件. Defaults to 0.
group_clipid (int, optional): # MediaInfo.sub_meta_info 中的实际序号.
original_clipid (None or [int], optional): 有些片段由其他片段合并,该字段用于片段来源,id是 media_map.json 中的实际序号. Defaults to None.
emb (np.array, optional): 片段 综合emb,. Defaults to None.
multi_factor (MultiFactorFeature), optional): 多维度特征. Defaults to None.
similar_clipseq ([Clip]], optional): 与该片段相似的片段,具体结构待定义. Defaults to None.
"""
self.media_type = media_type
self.mediaid = mediaid
self.time_start = time_start
self.duration = duration
self.clipid = clipid
self.path = path
self.timepoint_type = timepoint_type
self.text = text
self.stage = stage
self.group_time_start = group_time_start
self.group_clipid = group_clipid
self.duration_num = duration_num
self.original_clipid = original_clipid if original_clipid is not None else []
self.emb = emb
self.multi_factor = multi_factor
self.similar_clipseq = similar_clipseq
self.rythm = rythm
# TODO: 目前谱面中会有一些不必要的中间结果,比较占内存,现在代码里删掉,待后续数据协议确定
kwargs = {k: v for k, v in kwargs.items()}
self.__dict__.update(kwargs)
self.preprocess()
def preprocess(self):
pass
def spread_parameters(self):
pass
@property
def time_end(
self,
):
return self.time_start + self.duration
@property
def mvp_clip(self):
"""读取实际的片段数据为moviepy格式
Raises:
NotImplementedError: _description_
"""
raise NotImplementedError
class ClipSeq(object):
"""媒体片段序列"""
ClipClass = Clip
def __init__(self, clips) -> None:
"""_summary_
Args:
clips ([Clip]]): 媒体片段序列
"""
if not isinstance(clips, list):
clips = [clips]
if len(clips) == 0:
self.clips = []
elif isinstance(clips[0], dict):
self.clips = [self.ClipClass(**d) for d in clips]
else:
self.clips = clips
def set_clip_value(self, k, v):
"""给序列中的每一个clip 赋值"""
for i in range(len(self.clips)):
self.clips[i].__setattr__(k, v)
def __len__(
self,
):
return len(self.clips)
def merge(self, other, group_time_start_delta=None, groupid_delta=None):
"""融合其他ClipSeq。media_info 融合时需要记录 clip 所在的 groupid 和 group_time_start,delta用于表示变化
Args:
other (ClipSeq): 待融合的ClipSeq
group_time_start_delta (float, optional): . Defaults to None.
groupid_delta (int, optional): _description_. Defaults to None.
"""
if group_time_start_delta is not None or groupid_delta is not None:
for i, clip in enumerate(other):
if group_time_start_delta is not None:
clip.group_time_start += group_time_start_delta
if groupid_delta is not None:
clip.groupid += groupid_delta
self.clips.extend(other.clips)
for i in range(len(self.clips)):
self.clips[i].group_clipid = i
@property
def duration(
self,
):
"""Clip.duration的和
Returns:
float: 序列总时长
"""
if len(self.clips) == 0:
return 0
else:
return sum([c.duration for c in self.clips])
def __getitem__(self, i) -> Clip:
"""支持索引和切片操作,如果输入是整数则返回Clip,如果是切片,则返回ClipSeq
Args:
i (int or slice): 索引
Raises:
ValueError: 需要按照给的输入类型索引
Returns:
Clip or ClipSeq:
"""
if "int" in str(type(i)):
i = int(i)
if isinstance(i, int):
clip = self.clips[i]
return clip
elif isinstance(i, Iterable):
clips = [self.__getitem__(x) for x in i]
clipseq = ClipSeq(clips)
return clipseq
elif isinstance(i, slice):
if i.step is None:
step = 1
else:
step = i.step
clips = [self.__getitem__(x) for x in range(i.start, i.stop, step)]
clipseq = ClipSeq(clips)
return clipseq
else:
raise ValueError(
"unsupported input, should be int or slice, but given {}, type={}".format(
i, type(i)
)
)
def insert(self, idx, obj):
self.clips.insert(idx, obj)
def append(self, obj):
self.clips.append(obj)
def extend(self, objs):
self.clips.extend(objs)
@property
def duration_seq_emb(
self,
):
emb = np.array([c.duration for c in self.clips])
return emb
@property
def timestamp_seq_emb(self):
emb = np.array([c.time_start for c in self.clips])
return emb
@property
def rela_timestamp_seq_emb(self):
emb = self.timestamp_seq_emb / self.duration
return emb
def get_factor_seq_emb(self, factor, dim):
emb = []
for c in self.clips:
if factor not in c.multi_factor or c.multi_factor[factor] is None:
v = np.full(dim, np.inf)
else:
v = c.multi_factor[factor]
emb.append(v)
emb = np.stack(emb, axis=0)
return emb
def semantic_seq_emb(self, dim):
return self.get_factor_seq_emb(factor="semantics", dim=dim)
def emotion_seq_emb(self, dim):
return self.get_factor_seq_emb(factor="emotion", dim=dim)
def theme_seq_emb(self, dim):
return self.get_factor_seq_emb(factor="theme", dim=dim)
def to_dct(
self,
target_keys=None,
ignored_keys=None,
):
if ignored_keys is None:
ignored_keys = ["kwargs", "audio_path", "lyric_path", "start", "end"]
clips = [
clip.to_dct(target_keys=target_keys, ignored_keys=ignored_keys)
for clip in self.clips
]
return clips
@property
def mvp_clip(self):
"""读取实际的片段数据为moviepy格式
Raises:
NotImplementedError: _description_
"""
raise NotImplementedError
class ClipIds(object):
def __init__(
self,
clipids: list or int,
) -> None:
"""ClipSeq 中的 Clip序号,主要用于多个 Clip 融合后的 Clip, 使用场景如
1. 一个 MusicClip 可以匹配到多个 VideoClip,VideoClip 的索引便可以使用 ClipIds 定义。
Args:
clipids (list or int): ClipSeq 中的序号
"""
self.clipids = clipids if isinstance(clipids, list) else [clipids]
class ClipIdsSeq(object):
def __init__(self, clipids_seq: list) -> None:
"""多个 ClipIds,使用场景可以是
1. 将MediaClipSeq 进行重组,拆分重组成更粗粒度的ClipSeq;
Args:
clipids_seq (list): 组合后的 ClipIds 列表
"""
self.clipids_seq = (
clipids_seq if isinstance(clipids_seq, ClipIds) else [clipids_seq]
)
# TODO: metric后续可能是字典
class MatchedClipIds(object):
def __init__(
self, id1: ClipIds, id2: ClipIds, metric: float = None, **kwargs
) -> None:
"""两种模态数据的片段匹配对,使用场景 可以是
1. 音乐片段和视频片段 之间的匹配关系,
Args:
id1 (ClipIds): 第一种模态的片段
id2 (ClipIds): 第二种模态的片段
metric (float): 匹配度量距离
"""
self.id1 = id1 if isinstance(id1, ClipIds) else ClipIds(id1)
self.id2 = id2 if isinstance(id2, ClipIds) else ClipIds(id2)
self.metric = metric
self.__dict__.update(**kwargs)
class MatchedClipIdsSeq(object):
def __init__(self, seq: list, metric: float = None, **kwargs) -> None:
"""两种模态数据的序列匹配对,使用场景可以是
1. 音乐片段序列和视频片段序列 之间的匹配,每一个元素都是MatchedClipIds:
Args:
seq (list): 两种模态数据的序列匹配对列表
metric (float): 匹配度量距离
"""
self.seq = seq
self.metric = metric
self.__dict__.update(**kwargs)
+5
View File
@@ -0,0 +1,5 @@
from .clip import Clip, ClipSeq
from .clipid import ClipIds, MatchedClipIds, ClipIdsSeq, MatchedClipIdsSeq
from .clip_process import find_idx_by_time, find_idx_by_clip, get_subseq_by_time, get_subseq_by_idx, clip_is_top, clip_is_middle, clip_is_end, abadon_old_return_new, reset_clipseq_id, insert_endclip, insert_startclip, drop_start_end_by_time, complete_clipseq, complete_gap
from .clip_stat import stat_clipseq_duration
from .clip_filter import ClipFilter, ClipSeqFilter
+197
View File
@@ -0,0 +1,197 @@
from __future__ import annotations
from copy import deepcopy
from typing import Iterable, List, Tuple, Dict, Hashable, Any, Union
import numpy as np
from ...utils.util import convert_class_attr_to_dict
from ..general.items import Items, Item
from .clipid import MatchedClipIds
import logging
logger = logging.getLogger(__name__) # pylint: disable=invalid-name
__all__ = ["Clip", "ClipSeq"]
class Clip(Item):
"""媒体片段, 指转场点与转场点之间的部分"""
def __init__(
self,
time_start: float,
duration: float,
clipid: int = None,
media_type: str = None,
mediaid: str = None,
timepoint_type: str = None,
text: str = None,
stage: str = None,
path: str = None,
duration_num: int = None,
similar_clipseq: MatchedClipIds = None,
dynamic: float = None,
**kwargs,
):
"""
Args:
time_start (float): 开始时间,秒为单位,对应该媒体文件的, 和media_map.json上的序号一一对应
duration (_type_): 片段持续时间
clipid (int, or [int]): 由media_map提供的片段序号, 和media_map.json上的序号一一对应
media_type (str, optional): music, video,text, Defaults to None.
mediaid (int): 多媒体id, 当clipid是列表时,表示该片段是个融合片段
timepoint_type(int, ): 开始点的转场类型. Defaults to None.
text(str, optional): 该片段的文本描述,音乐可以是歌词,视频可以是台词,甚至可以是弹幕. Defaults to None.
stage(str, optional): 该片段在整个媒体文件中的结构位置,如音乐的intro、chrous、vesa,视频的片头、片尾、开始、高潮、转场等. Defaults to None.
path (str, optional): 该媒体文件的路径,用于后续媒体读取、处理. Defaults to None.
duration_num (_type_, optional): 片段持续帧数, Defaults to None.
similar_clipseq ([Clip]], optional): 与该片段相似的片段,具体结构待定义. Defaults to None.
"""
self.media_type = media_type
self.mediaid = mediaid
self.time_start = time_start
self.duration = duration
self.clipid = clipid
self.path = path
self.timepoint_type = timepoint_type
self.text = text
self.stage = stage
self.duration_num = duration_num
self.similar_clipseq = similar_clipseq
self.dynamic = dynamic
self.__dict__.update(**kwargs)
def preprocess(self):
pass
def spread_parameters(self):
pass
@property
def time_end(
self,
) -> float:
return self.time_start + self.duration
def get_emb(self, key: str, idx: int) -> np.float:
return self.emb.get_value(key, idx)
class ClipSeq(Items):
"""媒体片段序列"""
def __init__(self, items: List[Clip] = None):
super().__init__(items)
self.clipseq = self.data
def preprocess(self):
pass
def set_clip_value(self, k: Hashable, v: Any) -> None:
"""给序列中的每一个clip 赋值"""
for i in range(len(self.clipseq)):
self.clipseq[i].__setattr__(k, v)
def __len__(
self,
) -> int:
return len(self.clipseq)
@property
def duration(
self,
) -> float:
"""Clip.duration的和
Returns:
float: 序列总时长
"""
if len(self.clipseq) == 0:
return 0
else:
return sum([c.duration for c in self.clipseq])
def __getitem__(self, i: Union[int, Iterable]) -> Union[Clip, ClipSeq]:
"""支持索引和切片操作,如果输入是整数则返回Clip,如果是切片,则返回ClipSeq
Args:
i (int or slice): 索引
Raises:
ValueError: 需要按照给的输入类型索引
Returns:
Clip or ClipSeq:
"""
if "int" in str(type(i)):
i = int(i)
if isinstance(i, int):
clip = self.clipseq[i]
return clip
elif isinstance(i, Iterable):
clipseq = [self.__getitem__(x) for x in i]
clipseq = ClipSeq(clipseq)
return clipseq
elif isinstance(i, slice):
if i.step is None:
step = 1
else:
step = i.step
clipseq = [self.__getitem__(x) for x in range(i.start, i.stop, step)]
clipseq = ClipSeq(clipseq)
return clipseq
else:
raise ValueError(
"unsupported input, should be int or slice, but given {}, type={}".format(
i, type(i)
)
)
@property
def mvp_clip(self):
"""读取实际的片段数据为moviepy格式
Raises:
NotImplementedError: _description_
"""
raise NotImplementedError
@property
def duration_seq_emb(
self,
) -> np.array:
emb = np.array([c.duration for c in self.clipseq])
return emb
@property
def timestamp_seq_emb(self) -> np.array:
emb = np.array([c.time_start for c in self.clipseq])
return emb
@property
def rela_timestamp_seq_emb(self) -> np.array:
duration_seq = [c.duration for c in self.clipseq]
emb = np.cumsum(duration_seq) / self.duration
return emb
def get_emb(self, key: str, idx: int) -> np.float:
clip_start_idx = self.clipseq[0].clipid
clip_end_idx = self.clipseq[-1].clipid
# TODO: 待修改为更通用的形式
if idx is None:
idx = range(clip_start_idx, clip_end_idx + 1)
elif isinstance(idx, int):
idx += clip_start_idx
elif isinstance(idx, Iterable):
idx = [x + clip_start_idx for x in idx]
else:
raise ValueError(
f"idx only support None, int, Iterable, but given {idx},type is {type(idx)}"
)
return self.emb.get_value(key, idx=idx)
+46
View File
@@ -0,0 +1,46 @@
from typing import Callable, List, Union
from .clip import ClipSeq
from .clip_process import reset_clipseq_id
class ClipFilter(object):
"""clip滤波器,判断 Clip 是否符合标准
Args:
object (bool): 是否符合输入函数
"""
def __init__(self, funcs: Union[Callable, List[Callable]], logic_func: Callable=all) -> None:
"""多个 clip 判断函数,通过 逻辑与、或当综合结果。
Args:
funcs (list of func): 列表判断函数
logic_func (func, optional): all or any. Defaults to all.
"""
self.funcs = funcs if isinstance(funcs, list) else [funcs]
self.logic_func = logic_func
def __call__(self, clip) -> bool:
flag = [func(clip) for func in self.funcs]
flag = self.logic_func(flag)
return flag
# TODO
class ClipSeqFilter(object):
def __init__(self, filter: Callable) -> None:
self.filter = filter
def __call__(self, clipseq: ClipSeq) -> ClipSeq:
new_clipseq = []
n_clipseq = len(clipseq)
for i in range(n_clipseq):
clip = clipseq[i]
if self.filter(clip):
new_clipseq.append(clip)
new_clipseq = reset_clipseq_id(new_clipseq)
# logger.debug("ClipSeqFilter: clipseq length before={}, after={}".format(n_clipseq, len(new_clipseq)))
return new_clipseq
+64
View File
@@ -0,0 +1,64 @@
from typing import List, Union, Callable
from copy import deepcopy
from .clip import ClipSeq
from .clip_process import reset_clipseq_id
import logging
logger = logging.getLogger(__name__) # pylint: disable=invalid-name
# TODO: 不同类型的clip需要不同的融合方式
def fuse_clips(s1: ClipSeq, s2: ClipSeq) -> ClipSeq:
"""合并2个clip
Args:
s1 (Clip):
s2 (Clip):
Returns:
Clip: 合并后Clip
"""
if not isinstance(s2, list):
s2 = [s2]
s1 = deepcopy(s1)
for other_clip in s2:
s1.duration += other_clip.duration
if s1.stage is not None and other_clip.stage is not None:
# TODO:如何保留融合的clip信息
s1.stage = "{}_{}".format(s1.stage, other_clip.stage)
s1.origin_clipid.extend(other_clip.origin_clipid)
if s1.timepoint_type is not None and other_clip.timepoint_type is not None:
s1.timepoint_type = "{}_{}".format(
s1.timepoint_type, other_clip.timepoint_type
)
return s1
# TODO: 不同的filter和fusion函数不适用同一种流程,待优化
class ClipSeqFusion(object):
"""_summary_
Args:
object (_type_): _description_
"""
def __init__(self, filter: Callable, fuse_func: Callable = None) -> None:
self.filter = filter
self.fuse_func = fuse_func
def __call__(self, clipseq: ClipSeq) -> ClipSeq:
new_clipseq = []
n_clipseq = len(clipseq)
for i in range(n_clipseq):
clip = clipseq[i]
if self.filter(clip):
new_clipseq.append(clip)
new_clipseq = reset_clipseq_id(new_clipseq)
logger.debug(
"ClipSeqFilter: clipseq length before={}, after={}".format(
n_clipseq, len(new_clipseq)
)
)
return new_clipseq
+366
View File
@@ -0,0 +1,366 @@
from functools import partial
from copy import deepcopy
from typing import Iterable, List, Tuple, Union
import bisect
import logging
import numpy as np
from .clip import Clip, ClipSeq
from .clipid import ClipIds, ClipIdsSeq, MatchedClipIds, MatchedClipIdsSeq
logger = logging.getLogger(__name__) # pylint: disable=invalid-name
__all__ = [
"find_idx_by_rela_time",
"find_idx_by_time",
"find_idx_by_clip",
"get_subseq_by_time",
"get_subseq_by_idx",
"clip_is_top",
"clip_is_middle",
"clip_is_end",
"abadon_old_return_new",
"reset_clipseq_id",
"insert_endclip",
"insert_startclip",
"drop_start_end_by_time",
"complete_clipseq",
"complete_gap",
"get_subseq_by_stages",
"find_time_by_stage",
]
def find_idx_by_rela_time(clipseq: ClipSeq, timepoint: float) -> int:
clipseq_duration = clipseq.duration
timepoint = clipseq_duration * timepoint
clipseq_times = [c.duration for c in clipseq]
clipseq_times.insert(0, 0)
clipseq_times = np.cumsum(clipseq_times)
idx = bisect.bisect_right(clipseq_times, timepoint)
idx = min(max(0, idx - 1), len(clipseq) - 1)
return idx
def find_idx_by_time(clipseq: ClipSeq, timepoint: float) -> int:
"""寻找指定时间timepoint 在 clipseq 中的片段位置
Args:
clipseq (ClipSeq): 待寻找的片段序列
timepoint (float): 指定时间位置
Returns:
_type_: _description_
"""
clipseq_times = [c.time_start for c in clipseq]
idx = bisect.bisect_right(clipseq_times, timepoint)
idx = min(max(0, idx - 1), len(clipseq) - 1)
return idx
def find_idx_by_clip(clipseq: ClipSeq, clip: Clip, eps: float = 1e-4) -> int:
"""通过计算目标clip和clipseq中所有候选clip的交集占比来找最近clip
Args:
clipseq (ClipSeq): 候选clip序列
clip (Clip): 目标clip
eps (float, optional): 最小交集占比. Defaults to 1e-4.
Returns:
int: 目标clip在候选clip序列的位置,若无则为None
"""
timepoints = np.array([[c.time_start, c.time_start + c.duration] for c in clipseq])
clip_time_start = clip.time_start
clip_duraiton = clip.duration
clip_time_end = clip_time_start + clip_duraiton
max_time_start = np.maximum(timepoints[:, 0], clip_time_start)
min_time_end = np.minimum(timepoints[:, 1], clip_time_end)
intersection = min_time_end - max_time_start
intersection_ratio = intersection / clip_duraiton
max_intersection_ratio = np.max(intersection_ratio)
idx = np.argmax(intersection_ratio) if max_intersection_ratio > eps else None
return idx
def get_subseq_by_time(
clipseq: ClipSeq,
start: float = 0,
duration: float = None,
end: float = 1,
eps: float = 1e-2,
) -> ClipSeq:
"""根据时间对媒体整体做掐头去尾,保留中间部分。,也可以是大于1的数。
start和end如果是0-1的小数,则认为是是相对时间位置,实际位置会乘以duration;
start和end如果是大于1的数,则是绝对时间位置。
Args:
clipseq (ClipSeq): 待处理的序列
start (float,): 保留部分的开始,. Defaults to 0.
duration (float, optional): 媒体文件当前总时长
end (float, optional): 保留部分的结尾. Defaults to 1.
Returns:
ClipSeq: 处理后的序列
"""
if (start == 0 or start is None) and (end is None or end == 1):
logger.warning("you should set start or end")
return clipseq
if duration is None:
duration = clipseq.duration
if start is None or start == 0:
clip_start_idx = 0
else:
if start < 1:
start = start * duration
clip_start_idx = find_idx_by_time(clipseq, start)
if end is None or end == 1 or np.abs(duration - end) < eps:
clip_end_idx = -1
else:
if end < 1:
end = end * duration
clip_end_idx = find_idx_by_time(clipseq, end)
if clip_end_idx != -1 and clip_start_idx >= clip_end_idx:
logger.error(
f"clip_end_idx({clip_end_idx}) should be > clip_start_idx({clip_start_idx})"
)
subseq = get_subseq_by_idx(clipseq, clip_start_idx, clip_end_idx)
return subseq
def get_subseq_by_idx(clipseq: ClipSeq, start: int = None, end: int = None) -> ClipSeq:
"""通过指定索引范围,切片子序列
Args:
clipseq (ClipSeq):
start (int, optional): 开始索引. Defaults to None.
end (int, optional): 结尾索引. Defaults to None.
Returns:
_type_: _description_
"""
if start is None and end is None:
return clipseq
if start is None:
start = 0
if end is None:
end = len(clipseq)
return clipseq[start:end]
def clip_is_top(clip: Clip, total: float, th: float = 0.1) -> bool:
"""判断Clip是否属于开始部分
Args:
clip (Clip):
total (float): 所在ClipSeq总时长
th (float, optional): 开始范围的截止位置. Defaults to 0.05.
Returns:
Bool: 是不是头部Clip
"""
clip_time = clip.time_start
if clip_time / total <= th:
return True
else:
return False
def clip_is_end(clip: Clip, total: float, th: float = 0.9) -> bool:
"""判断Clip是否属于结尾部分
Args:
clip (Clip):
total (float): 所在ClipSeq总时长
th (float, optional): 结尾范围的开始位置. Defaults to 0.9.
Returns:
Bool: 是不是尾部Clip
"""
clip_time = clip.time_start + clip.duration
if clip_time / total >= th:
return True
else:
return False
def clip_is_middle(
clip: Clip, total: float, start: float = 0.05, end: float = 0.9
) -> bool:
"""判断Clip是否属于中间部分
Args:
clip (Clip):
total (float): 所在ClipSeq总时长
start (float, optional): 中间范围的开始位置. Defaults to 0.05.
start (float, optional): 中间范围的截止位置. Defaults to 0.9.
Returns:
Bool: 是不是中间Clip
"""
if start >= 0 and start < 1:
start = total * start
if end > 0 and end <= 1:
end = total * end
clip_time_start = clip.time_start
clip_time_end = clip.time_start + clip.duration
if (clip_time_start >= start) and (clip_time_end <= end):
return True
else:
return False
def abadon_old_return_new(s1: Clip, s2: Clip) -> Clip:
"""特殊的融合方式
Args:
s1 (Clip): 靠前的clip
s2 (Clip): 靠后的clip
Returns:
Clip: 融合后的Clip
"""
return s2
# TODO:待确认是否要更新clipid,不方便对比着json进行debug
def reset_clipseq_id(clipseq: ClipSeq) -> ClipSeq:
for i in range(len(clipseq)):
if isinstance(clipseq[i], dict):
clipseq[i]["clipid"] = i
else:
clipseq[i].clipid = i
return clipseq
def insert_startclip(clipseq: ClipSeq) -> ClipSeq:
"""给ClipSeq插入一个开始片段。
Args:
clipseq (ClipSeq):
clip_class (Clip, optional): 插入的Clip类型. Defaults to Clip.
Returns:
ClipSeq: 插入头部Clip的新ClipSeq
"""
if clipseq[0].time_start > 0:
start = clipseq.ClipClass(
time_start=0, duration=round(clipseq[0].time_start, 3), timepoint_type=0
)
clipseq.insert(0, start)
clipseq = reset_clipseq_id(clipseq)
return clipseq
def insert_endclip(clipseq: ClipSeq, duration: float) -> ClipSeq:
"""给ClipSeq插入一个尾部片段。
Args:
clipseq (ClipSeq):
duration(float, ): 序列的总时长
clip_class (Clip, optional): 插入的Clip类型. Defaults to Clip.
Returns:
ClipSeq: 插入尾部Clip的新ClipSeq
"""
clipseq_endtime = clipseq[-1].time_start + clipseq[-1].duration
if duration - clipseq_endtime > 1:
end = clipseq.ClipClass(
time_start=round(clipseq_endtime, 3),
duration=round(duration - clipseq_endtime, 3),
timepoint_type=0,
)
clipseq.append(end)
clipseq = reset_clipseq_id(clipseq)
return clipseq
def drop_start_end_by_time(
clipseq: ClipSeq, start: float, end: float, duration: float = None
):
return get_subseq_by_time(clipseq=clipseq, start=start, end=end, duration=duration)
def complete_clipseq(
clipseq: ClipSeq, duration: float = None, gap_th: float = 2
) -> ClipSeq:
"""绝大多数需要clipseq中的时间信息是连续、完备的,有时候是空的,需要补足的部分。
如歌词时间戳生成的music_map缺头少尾、中间有空的部分。
Args:
clipseq (ClipSeq): 待补集的序列
duration (float, optional): 整个序列持续时间. Defaults to None.
gap_th (float, optional): 有时候中间空隙过短就会被融合到上一个片段中. Defaults to 2.
Returns:
ClipSeq: 补集后的序列,时间连续、完备。
"""
if isinstance(clipseq, list):
clipseq = ClipSeq(clipseq)
return complete_clipseq(clipseq=clipseq, duration=duration, gap_th=gap_th)
clipseq = complete_gap(clipseq, th=gap_th)
clipseq = insert_startclip(clipseq)
if duration is not None:
clipseq = insert_endclip(clipseq, duration)
return clipseq
def complete_gap(clipseq: ClipSeq, th: float = 2) -> ClipSeq:
"""generate blank clip timepoint = 0,如果空白时间过短,则空白附到上一个歌词片段中。
Args:
clipseq (ClipSeq): 原始的歌词生成的MusicClipSeq
th (float, optional): 有时候中间空隙过短就会被融合到上一个片段中. Defaults to 2.
Returns:
ClipSeq: 补全后的
"""
gap_clipseq = []
clipid = 0
for i in range(len(clipseq) - 1):
time_start = clipseq[i].time_start
duration = clipseq[i].duration
time_end = time_start + duration
next_time_start = clipseq[i + 1].time_start
time_diff = next_time_start - time_end
if time_diff >= th:
blank_clip = clipseq.ClipClass(
time_start=time_end,
duration=time_diff,
timepoint_type=0,
clipid=clipid,
)
gap_clipseq.append(blank_clip)
clipid += 1
else:
clipseq[i].duration = next_time_start - time_start
clipseq.extend(gap_clipseq)
clipseq.clips = sorted(clipseq.clips, key=lambda clip: clip.time_start)
reset_clipseq_id(clipseq)
return clipseq
def find_time_by_stage(
clipseq: ClipSeq, stages: Union[str, List[str]] = None
) -> Tuple[float, float]:
if isinstance(stages, list):
stages = [stages]
for clip in clipseq:
if clip.stage in stages:
return clip.time_start, clip.time_end
return None, None
def get_subseq_by_stages(clipseq: ClipSeq, stages: Union[str, List[str]]) -> ClipSeq:
if isinstance(stages, List):
stages = [stages]
start, _ = find_time_by_stage(clipseq, stages[0])
_, end = find_time_by_stage(clipseq, stages[-1])
if start1 is None:
start1 = 0
if end2 is None:
end2 = clipseq.duration
subseq = get_subseq_by_time(clipseq=clipseq, start=start, end=end)
return subseq
+13
View File
@@ -0,0 +1,13 @@
from typing import Tuple
import numpy as np
from .clip import ClipSeq
def stat_clipseq_duration(
clipseq: ClipSeq,
) -> Tuple[np.array, np.array]:
clip_duration = [clip.duration for clip in clipseq]
(hist, bin_edges) = np.histogram(clip_duration)
return hist, bin_edges
+70
View File
@@ -0,0 +1,70 @@
from __future__ import annotations
from typing import Union, List
__all__ = [
"ClipIds",
"ClipIdsSeq",
"MatchedClipIds",
"MatchedClipIdsSeq",
]
class ClipIds(object):
def __init__(
self,
clipids: Union[int, List[int]],
) -> None:
"""ClipSeq 中的 Clip序号,主要用于多个 Clip 融合后的 Clip, 使用场景如
1. 一个 MusicClip 可以匹配到多个 VideoClip,VideoClip 的索引便可以使用 ClipIds 定义。
Args:
clipids (list or int): ClipSeq 中的序号
"""
self.clipids = clipids if isinstance(clipids, list) else [clipids]
class ClipIdsSeq(object):
def __init__(self, clipids_seq: List[ClipIds]) -> None:
"""多个 ClipIds,使用场景可以是
1. 将MediaClipSeq 进行重组,拆分重组成更粗粒度的ClipSeq;
Args:
clipids_seq (list): 组合后的 ClipIds 列表
"""
self.clipids_seq = (
clipids_seq if isinstance(clipids_seq, ClipIds) else [clipids_seq]
)
# TODO: metric后续可能是字典
class MatchedClipIds(object):
def __init__(
self, id1: ClipIds, id2: ClipIds, metric: float = None, **kwargs
) -> None:
"""两种模态数据的片段匹配对,使用场景 可以是
1. 音乐片段和视频片段 之间的匹配关系,
Args:
id1 (ClipIds): 第一种模态的片段
id2 (ClipIds): 第二种模态的片段
metric (float): 匹配度量距离
"""
self.id1 = id1 if isinstance(id1, ClipIds) else ClipIds(id1)
self.id2 = id2 if isinstance(id2, ClipIds) else ClipIds(id2)
self.metric = metric
self.__dict__.update(**kwargs)
class MatchedClipIdsSeq(object):
def __init__(self, seq: List[MatchedClipIds], metric: float = None, **kwargs) -> None:
"""两种模态数据的序列匹配对,使用场景可以是
1. 音乐片段序列和视频片段序列 之间的匹配,每一个元素都是MatchedClipIds:
Args:
seq (list): 两种模态数据的序列匹配对列表
metric (float): 匹配度量距离
"""
self.seq = seq
self.metric = metric
self.__dict__.update(**kwargs)
View File
+72
View File
@@ -0,0 +1,72 @@
from collections import namedtuple
from typing import NamedTuple, Tuple, List
import logging
import os
import numpy as np
import subprocess
import requests
import wget
from .youtube import download_youtube
from .flicker import download_flickr
from .ffmpeg import ffmpeg_load
logger = logging.getLogger(__name__)
# DownloadStatus = namedtuple("DownloadStatus", ["status_code", "msg"])
status_code = {0: "download: succ",
-1: "download: failed",
-2: "clip: failed",
-3: "directory not exists",
-4: "skip task",
- 404: "param error"}
def download_with_request(url, path):
res = requests.get(url)
if res.status_code == '200' or res.status_code == 200:
with open(path, "wb") as f:
f.write(res.content)
else:
print('request failed')
return path
def download_video(url, save_path:str=None, save_dir:str=None, basename:str=None, filename:str=None, format:str=None, data_type: str="wget", **kwargs) -> Tuple[int, str]:
if save_path is None:
if basename is None:
basename = f"{filename}.{format}"
save_path = os.path.join(save_dir, basename)
if save_dir is None:
save_dir = os.path.dirname(save_path)
if basename is None:
basename = os.path.basename(save_path)
if filename is None:
filename, format = os.path.splitext(basename)
os.makedirs(save_dir, exist_ok=True)
if os.path.exists(save_path):
return (-4, save_path)
try:
if data_type == "requests":
save_path = download_with_request(url=url, path=save_path)
elif data_type == "wget":
save_path = wget.download(url=url, out=save_path)
elif data_type == "youtube":
save_path = download_youtube(url, format=format, save_dir=save_dir, filename=basename)
elif data_type == "flickr":
save_path = download_flickr(url, save_path)
elif data_type == "ffmpeg":
code = ffmpeg_load(url=url, save_path=save_path)
else:
raise ValueError(f"data_type shoulbe one of [wget, youtube, flickr, ffmpeg], but given {data_type}")
except Exception as e:
logger.error("failed download file {} to {} failed!".format(url, save_path))
logger.exception(e)
return (-1, None)
return (0, save_path)
+20
View File
@@ -0,0 +1,20 @@
class SubprocessError(Exception):
"""
Exception object that contains information about an error that occurred
when running a command line command with a subprocess.
"""
def __init__(self, cmd, return_code, stdout, stderr, *args):
msg = 'Got non-zero exit code ({1}) from command "{0}": {2}'
if stderr.strip():
err_msg = stderr
else:
err_msg = stdout
msg = msg.format(cmd[0], return_code, err_msg)
self.cmd = cmd
self.cmd_return_code = return_code
self.cmd_stdout = stdout
self.cmd_stderr = stderr
super(SubprocessError, self).__init__(msg, *args)
+39
View File
@@ -0,0 +1,39 @@
import subprocess
from .error import SubprocessError
class FfmpegInvalidURLError(Exception):
"""
Exception raised when a 4XX or 5XX error is returned when making a request
"""
def __init__(self, url, error, *args):
self.url = url
self.error = error
msg = 'Got error when making request to "{}": {}'.format(url, error)
super(FfmpegInvalidURLError, self).__init__(msg, *args)
def ffmpeg_load(url: str, save_path: str) -> str:
def run(cmd):
proc = subprocess.Popen(
cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
stdout, stderr = proc.communicate()
return_code = proc.returncode
if return_code != 0:
raise SubprocessError(
cmd, return_code, stdout.decode(), stderr.decode())
return return_code
command = ['ffmpeg', '-n', '-i', url, '-t', '10', '-f', 'mp4',
'-r', '30', '-vcodec', 'h264', save_path, '-loglevel', 'error']
code = run(command)
return code
+22
View File
@@ -0,0 +1,22 @@
import os
from .ffmpeg import ffmpeg_load
def extract_flickr_id(url):
return url.strip('/').split('/')[-4]
def download_flickr(url: str, save_path: str) -> str:
code = -1
code = ffmpeg_load(url=url,
save_path=save_path)
if code == 0:
return (code, save_path)
# only retry when failed!
flickr_id = extract_flickr_id(url)
url = 'https://www.flickr.com/video_download.gne?id={}'.format(
flickr_id)
code = ffmpeg_load(url=url,
save_path=save_path)
return save_path
+13
View File
@@ -0,0 +1,13 @@
import os
from pytube import YouTube
def download_youtube(url, format, save_dir, filename):
youtube = YouTube(url)
streams = youtube.streams.filter(progressive=True,
file_extension=format)
save_path = streams.get_highest_resolution().download(output_path=save_dir,
filename=filename)
return save_path
+2
View File
@@ -0,0 +1,2 @@
from .emb import *
from .h5py_emb import H5pyMediaMapEmb, H5pyMediaMapEmbProxy
+104
View File
@@ -0,0 +1,104 @@
"""用于将 mediamap中的emb存储独立出去,仍处于开发中
"""
import logging
import numpy as np
logger = logging.getLogger(__name__) # pylint: disable=invalid-name
__all__ = ["MediaMapEmb"]
class MediaMapEmb(object):
def __init__(self, path: str) -> None:
"""
OfflineEmb = {
"overall_algo": Emb, # 整个文件的Emb
# 整个文件的多维度 Emb
"theme": np.array, # 主题,
"emotion_algo": np.array, # 情绪,
"semantic_algo": np.array, # 语义
"clips_overall_algo": np.array, n_clip x clip_emb
"clips_emotion_algo": np.array, n_clip x clip_emb
"clips_semantic_algo": np.array, n_clip x clip_emb
"clips_theme_algo": np.array, n_clip x clip_emb
"scenes_overall_algo": np.array, n_scenes x scene_emb
"scenes_emotion_algo": np.array, n_scenes x scene_emb
"scenes_semantic_algo": np.array, n_scenes x scene_emb
"scenes_theme_algo": E np.arraymb, n_scenes x scene_emb
# 片段可以是转场切分、MusicStage等, clips目前属于转场切分片段
# 若后续需要新增段落分割,可以和clips同级新增 stage字段。
"frames_overall_algo": np.array, n_frames x frame_emb
"frames_emotion_algo": np.array, n_frames x frame_emb
"frames_semantic_algo": np.array, n_frames x frame_emb
"frames_theme_algo": np.array, n_frames x frame_emb
"frames_objs": {
"frame_id": { #
"overall_algo": np.array, n_objs x obj_emb
"emotion_algo": np.array, n_objs x obj_emb
"semantic_algo": np.array, n_objs x obj_emb
"theme_algo": np.array, n_objs x obj_emb
}
}
"roles_algo": {
"roleid": np.array, n x obj_emb
}
}
Args:
path (str): hdf5 存储路径
"""
self.path = path
def get_value(self, key, idx=None):
raise NotImplementedError
def __getitem__(self, key):
return self.get_value(key)
def get_media(self, factor, algo):
return self.get_value(f"{factor}_{algo}")
def get_clips(self, factor, algo, idx=None):
return self.get_value(f"clips_{factor}_{algo}", idx=idx)
def get_frames(self, factor, algo, idx=None):
return self.get_value(f"frames_{factor}_{algo}", idx=idx)
def get_frame_objs(self, frame_idx, factor, algo, idx=None):
return self.get_value(["frames_objs", frame_idx, f"{factor}_{algo}"], idx=idx)
def set_value(self, key, value, idx=None):
raise NotImplementedError
def set_media(self, factor, value, algo):
self.set_value([f"{factor}_{algo}"], value)
def set_clips(self, factor, value, algo, idx=None):
self.set_value([f"clips_{factor}_{algo}"], value, idx=idx)
def set_frames(self, factor, value, algo, idx=None):
self.set_value([f"frames_{factor}_{algo}"], value)
def set_frame_objs(self, frame_idx, factor, value, algo, idx=None):
return self.set_value(
["frames_objs", frame_idx, f"{factor}_{algo}"], value, idx=idx
)
def set_roles(self, algo, value, idx=None):
return self.set_value(f"roles_{algo}", value, idx=idx)
def get_roles(self, algo, idx=None):
return self.get_value(f"roles_{algo}", idx=idx)
def __setitem__(self, key, value):
self.set_value(self, key, value)
class MediaMapEmbProxy(MediaMapEmb):
pass
+119
View File
@@ -0,0 +1,119 @@
from typing import Union, List
import logging
import h5py
import numpy as np
from .emb import MediaMapEmb
logger = logging.getLogger(__name__) # pylint: disable=invalid-name
__all__ = ["H5pyMediaMapEmb", "save_value_with_h5py"]
def save_value_with_h5py(
path: str,
value: Union[np.ndarray, None],
key: str,
idx: Union[int, List[int]] = None,
dtype=None,
shape=None,
overwrite: bool = False,
):
with h5py.File(path, "a") as f:
if dtype is None:
dtype = value.dtype
if shape is None:
shape = value.shape
del_key = False
if key in f:
if overwrite:
del_key = True
if f[key].dtype != h5py.special_dtype(vlen=str):
if f[key].shape != value.shape:
del_key = True
if del_key:
del f[key]
if key not in f:
f.create_dataset(key, shape=shape, dtype=dtype)
if idx is None:
f[key][...] = value
else:
f[key][idx] = value
class H5pyMediaMapEmb(MediaMapEmb):
def __init__(self, path: str) -> None:
"""
OfflineEmb = {
"overall_algo": Emb, # 整个文件的Emb
# 整个文件的多维度 Emb
"theme": np.array, # 主题,
"emotion_algo": np.array, # 情绪,
"semantic_algo": np.array, # 语义
"clips_overall_algo": np.array, n_clip x clip_emb
"clips_emotion_algo": np.array, n_clip x clip_emb
"clips_semantic_algo": np.array, n_clip x clip_emb
"clips_theme_algo": np.array, n_clip x clip_emb
"scenes_overall_algo": np.array, n_scenes x scene_emb
"scenes_emotion_algo": np.array, n_scenes x scene_emb
"scenes_semantic_algo": np.array, n_scenes x scene_emb
"scenes_theme_algo": E np.arraymb, n_scenes x scene_emb
# 片段可以是转场切分、MusicStage等, clips目前属于转场切分片段
# 若后续需要新增段落分割,可以和clips同级新增 stage字段。
"frames_overall_algo": np.array, n_frames x frame_emb
"frames_emotion_algo": np.array, n_frames x frame_emb
"frames_semantic_algo": np.array, n_frames x frame_emb
"frames_theme_algo": np.array, n_frames x frame_emb
"frames_objs_algo": {
"frame_id_algo": { #
"overall_algo": np.array, n_objs x obj_emb
"emotion_algo": np.array, n_objs x obj_emb
"semantic_algo": np.array, n_objs x obj_emb
"theme_algo": np.array, n_objs x obj_emb
}
}
"roles_algo": {
"roleid": np.array, n x obj_emb
}
}
Args:
path (str): hdf5 存储路径
"""
super().__init__(path)
# 待优化支持 with open 的方式来读写
self.f = h5py.File(path, "a")
def _keys_index(self, key):
if not isinstance(key, list):
key = [key]
key = "/".join([str(x) for x in key if x is not None])
return key
def get_value(self, key, idx=None):
new_key = self._keys_index(key)
if idx is None:
data = np.array(self.f[new_key])
else:
data = np.array(self.f[new_key][idx])
return data
def set_value(self, key, value, idx=None):
new_key = self._keys_index(key)
if new_key not in self.f:
self.f.create_dataset(new_key, shape=value.shape, dtype=value.dtype)
if idx is None:
self.f[new_key][...] = value
else:
self.f[new_key][idx] = value
def close(self):
self.f.close()
class H5pyMediaMapEmbProxy(H5pyMediaMapEmb):
pass
View File
View File
@@ -0,0 +1,28 @@
from typing import List, Union, Any
import torch
from torch import nn
import numpy as np
import h5py
class BaseFeatureExtractor(nn.Module):
def __init__(self, device: str = "cpu", dtype=torch.float32, name: str = None):
super().__init__()
self.device = device
self.dtype = dtype
self.name = name
def extract(
self, data: Any, return_type: Union[str, str] = "numpy"
) -> Union[np.ndarray, torch.tensor]:
raise NotADirectoryError
def __call__(self, *args: Any, **kwds: Any) -> Any:
return self.extract(*args, **kwds)
def save_with_h5py(self, f: Union[h5py.File, str], *args, **kwds):
raise NotImplementedError
def forward(self, *args: Any, **kwds: Any) -> Any:
return self.extract(*args, **kwds)
+1
View File
@@ -0,0 +1 @@
from .items import Items
+69
View File
@@ -0,0 +1,69 @@
from collections import UserList
from collections.abc import Iterable
from typing import Iterator, Any, List
from ...utils.util import convert_class_attr_to_dict
__all__ = ["Item", "Items"]
class Item(object):
def __init__(self) -> None:
pass
def to_dct(self, target_keys: List[str] = None, ignored_keys: List[str] = None):
base_ignored_keys = [
"kwargs",
]
if isinstance(ignored_keys, list):
base_ignored_keys.extend(ignored_keys)
elif isinstance(ignored_keys, str):
base_ignored_keys.append(ignored_keys)
else:
pass
return convert_class_attr_to_dict(
self, target_keys=target_keys, ignored_keys=base_ignored_keys
)
def preprocess(self):
pass
class Items(UserList):
def __init__(
self,
data: Any = None,
):
if data is None:
data = list()
if not isinstance(data, list):
data = [data]
super().__init__(data)
def __len__(self):
return len(self.data)
def __getitem__(self, i):
return self.data[i]
def __delitem__(self, i):
del self.data[i]
def __setitem__(self, i, v):
self.data[i] = v
def insert(self, i, v):
self.data.insert(i, v)
def __str__(self):
return str(self.data)
def to_dct(self, target_keys: List[str] = None, ignored_keys: List[str] = None):
items = [item.to_dct(target_keys, ignored_keys) for item in self.data]
return items
def __iter__(self) -> Iterator:
return iter(self.data)
def preprocess(self):
pass
+1
View File
@@ -0,0 +1 @@
from .media_map import MetaInfo, MediaMap, MetaInfoList
+393
View File
@@ -0,0 +1,393 @@
from __future__ import annotations
import bisect
import logging
from copy import deepcopy
from functools import partial
from typing import Any, Callable, Iterable, List, Union, Tuple, Dict
import numpy as np
from ..clip.clip_process import get_subseq_by_time
from ..clip.clip_stat import stat_clipseq_duration
from ..clip import Clip, ClipSeq, ClipIds, MatchedClipIds, MatchedClipIdsSeq
from .media_map_process import get_sub_mediamap_by_time
from ..emb import MediaMapEmb, H5pyMediaMapEmb
from ..general.items import Item, Items
from ...utils.data_util import pick_subdct
from ...utils.util import convert_class_attr_to_dict, load_dct_from_file
logger = logging.getLogger(__name__) # pylint: disable=invalid-name
__all__ = ["MetaInfo", "MetaInfoList", "MediaMap", "MediaMapSeq"]
class MetaInfo(Item):
"""歌曲、视频等媒体文件级别的元信息"""
def __init__(
self,
mediaid=None,
media_name=None,
media_duration=None,
signature=None,
media_path: str = None,
media_map_path: str = None,
start: float = None,
end: float = None,
ext=None,
**kwargs,
):
super(MetaInfo).__init__()
self.mediaid = mediaid
self.media_name = media_name
self.media_duration = media_duration
self.signature = signature
self.media_path = media_path
self.media_map_path = media_map_path
self.start = start
self.end = end
self.ext = ext
self.__dict__.update(**kwargs)
self.preprocess()
def preprocess(self):
self.set_start_end()
def set_start_end(self):
if self.start is None:
self.start = 0
elif self.start >= 0 and self.start <= 1:
self.start = self.start * self.media_duration
if self.end is None:
self.end = self.media_duration
elif self.end >= 0 and self.end <= 1:
self.end = self.end * self.media_duration
class MetaInfoList(Items):
"""媒体元数据列表,主要用于多歌曲、多视频剪辑时存储原单一媒体文件的元信息"""
def __init__(self, items: Union[MetaInfo, List[MetaInfo]] = None):
"""
Args:
meta_info_list (list, optional): MetaInfo 列表. Defaults to None.
"""
if items is None:
items = []
else:
items = items if isinstance(items, list) else [items]
super().__init__(items)
self.meta_info_list = self.items
if len(self.items) > 1:
self.reset()
def __len__(self):
return len(self.meta_info_list)
def __getitem__(self, i) -> MetaInfo:
return self.meta_info_list[i]
@property
def groupnum(self) -> int:
return len(self.meta_info_list)
class MediaMap(object):
"""媒体信息基类,也可以理解为音乐谱面、视觉谱面、音游谱面基类。主要有 MetaInfo、MetaInfoList、ClipSeq 属性。
不同的媒体信息的 属性 类会有不同,所以在类变量里做定义。如有变化,可以定义自己的属性类。
"""
def __init__(
self,
meta_info: MetaInfo = None,
clipseq: ClipSeq = None,
stageseq: ClipSeq = None,
frameseq: ClipSeq = None,
emb: H5pyMediaMapEmb = None,
**kwargs,
):
"""用于存储media的相关信息,media_info是json或直接字典
Args:
meta_info (MetaInfo): 当sub_meta_info不为None时, meta_info由sub_meta_info整合而成
sub_meta_info (None or [MetaInfo]): 当多个MediaInfo拼在一起时,用于保留子MediaInfo的信息
clipseq (ClipSeq): # 按照clipidx排序;
stageseq (ClipSeq): # 比 clipseq 更高纬度的片段划分,例如clips是镜头分割,stages是scenes分割;clips是关键点分割,stages是结构分割;
frameseq (ClipSeq): # 比 clipseq 更低纬度的片段划分
kwargs (dict, optional): 所有相关信息都会作为 meta_info 的补充,赋值到 meta_info 中
"""
self.meta_info = meta_info
self.clipseq = clipseq
self.frameseq = frameseq
self.stageseq = stageseq
self.emb = emb
self.meta_info.__dict__.update(**kwargs)
self.preprocess()
def preprocess(
self,
):
if (self.meta_info.start != 0 and self.meta_info.start is not None) or (
self.meta_info.end is not None and self.meta_info.end == 1
):
self.drop_head_and_tail()
self.meta_info.preprocess()
if self.clipseq is not None:
self.clipseq.preprocess()
if self.frameseq is not None:
self.frameseq.preprocess()
if self.stageseq is not None:
self.stageseq.preprocess()
self.clip_start_idx = self.clipseq[0].clipid
self.clip_end_idx = self.clipseq[-1].clipid
def drop_head_and_tail(self) -> MediaMap:
self.clipseq = get_subseq_by_time(
self.clipseq,
start=self.meta_info.start,
end=self.meta_info.end,
duration=self.meta_info.media_duration,
)
if self.stageseq is not None:
self.stageseq = get_subseq_by_time(
self.clipseq,
start=self.meta_info.start,
end=self.meta_info.end,
duration=self.meta_info.media_duration,
)
def set_clip_value(self, k, v):
"""为clipseq中的每个clip赋值,
Args:
k (str): Clip中字段名
v (any): Clip中字段值
"""
self.clipseq.set_clip_value(k, v)
def spread_metainfo_2_clip(
self, target_keys: List = None, ignored_keys: List = None
) -> None:
"""将metainfo中的信息赋值到clip中,便于clip后面做相关处理。
Args:
target_keys ([str]): 待赋值的目标字段
"""
dst = pick_subdct(
self.meta_info.__dict__, target_keys=target_keys, ignored_keys=ignored_keys
)
for k, v in dst.items():
self.set_clip_value(k, v)
def spread_parameters(self, target_keys: list, ignored_keys) -> None:
"""元数据广播,将 media_info 的元数据广播到 clip 中,以及调用 clip 自己的参数传播。"""
self.spread_metainfo_2_clip(target_keys=target_keys, ignored_keys=ignored_keys)
for clip in self.clipseq:
clip.spread_parameters()
def stat(
self,
):
"""统计 media_info 相关信息,便于了解,目前统计内容有
1. 片段长度
"""
self.stat_clipseq_duration()
def stat_clipseq_duration(
self,
):
hist, bin_edges = stat_clipseq_duration(self.clipseq)
print(self.media_name, "bin_edges", bin_edges)
print(self.media_name, "hist", hist)
def to_dct(self, target_keys: list = None, ignored_keys: list = None):
raise NotImplementedError
@property
def duration(
self,
):
return self.clipseq.duration
@property
def mediaid(
self,
):
return self.meta_info.mediaid
@property
def media_name(
self,
):
return self.meta_info.media_name
@property
def duration_seq_emb(self):
return self.clipseq.duration_seq_emb
@property
def timestamp_seq_emb(self):
return self.clipseq.timestamp_seq_emb
@property
def rela_timestamp_seq_emb(self):
return self.clipseq.rela_timestamp_seq_emb
def get_emb(self, key, idx=None):
# TODO: 待修改为更通用的形式
if idx is None:
idx = range(self.clip_start_idx, self.clip_end_idx + 1)
elif isinstance(idx, int):
idx += self.clip_start_idx
elif isinstance(idx, Iterable):
idx = [x + self.clip_start_idx for x in idx]
else:
raise ValueError(
f"idx only support None, int, Iterable, but given {idx},type is {type(idx)}"
)
return self.emb.get_value(key, idx=idx)
def get_meta_info_attr(self, key: str) -> Any:
return getattr(self.meta_info, key)
@classmethod
def from_json_path(
cls, path: Dict, emb_path: str, media_path: str = None, **kwargs
) -> MediaMap:
media_map = load_dct_from_file(path)
emb = H5pyMediaMapEmb(emb_path)
return cls.from_data(media_map, emb=emb, media_path=media_path, **kwargs)
class MediaMapSeq(Items):
def __init__(self, maps: List[MediaMap]) -> None:
super().__init__(maps)
self.maps = self.data
self.preprocess()
self.each_map_clipseq_num = [len(m.clipseq) for m in self.maps]
self.each_map_clipseq_num_cumsum = np.cumsum([0] + self.each_map_clipseq_num)
@property
def clipseq(self):
clipseq = []
for m in self.maps:
clipseq.extend(m.clipseq.data)
return type(self.maps[0].clipseq)(clipseq)
@property
def stagesseq(self):
stagesseq = []
for m in self.maps:
stagesseq.extend(m.stagesseq.data)
return type(self.maps[0].stagesseq)(stagesseq)
@property
def frameseq(self):
frameseq = []
for m in self.maps:
frameseq.extend(m.frameseq.data)
return type(self.maps[0].frameseq)(frameseq)
def preprocess(self):
for m in self.maps:
m.preprocess()
def _combine_str(
self,
attrs: List[str],
sep: str = "|",
single_maxlen: int = 10,
total_max_length: int = 60,
) -> str:
return sep.join([str(attr)[:single_maxlen] for attr in attrs])[
:total_max_length
]
def get_meta_info_attr(self, key: str, func: Callable) -> Any:
attrs = [m.get_meta_info_attr(key) for m in self.maps]
return func(attrs)
@property
def mediaid(self) -> str:
return self.get_meta_info_attr(key="mediaid", func=self._combine_str)
@property
def media_name(self) -> str:
return self.get_meta_info_attr(key="media_name", func=self._combine_str)
@property
def duration(self) -> float:
return sum([m.duration for m in self.maps])
@property
def media_duration(self) -> float:
return self.get_meta_info_attr(key="media_duration", func=sum)
@classmethod
def from_json_paths(
cls,
media_map_class: MediaMap,
media_paths: str,
media_map_paths: str,
emb_paths: str,
**kwargs,
) -> MediaMapSeq:
map_seq = [
media_map_class.from_json_path(
path=media_map_paths[i],
emb_path=emb_paths[i],
media_path=media_paths[i],
**kwargs,
)
for i in range(len(media_map_paths))
]
return cls(map_seq)
# TODO: implement mapseq stat func
def stat(self):
for m in self.maps:
m.stat()
def _combine_embs(self, embs):
return np.concatenate(embs, axis=0)
@property
def duration_seq_emb(self):
embs = [m.duration_seq_emb for m in self.maps]
return self._combine_embs(embs)
@property
def timestamp_seq_emb(self):
embs = [m.timestamp_seq_emb for m in self.maps]
return self._combine_embs(embs)
@property
def rela_timestamp_seq_emb(self):
embs = [m.rela_timestamp_seq_emb for m in self.maps]
return self._combine_embs(embs)
def clip_idx_2_map_idx(self, idx):
target_map_idx = bisect.bisect_right(self.each_map_clipseq_num_cumsum, idx)
target_map_idx = min(max(0, target_map_idx - 1), len(self.maps) - 1)
target_map_clip_idx = idx - self.each_map_clipseq_num_cumsum[target_map_idx]
return target_map_idx, target_map_clip_idx
def get_emb(self, key: str, idx: Union[None, int, List[int]] = None) -> np.array:
if idx is None:
embs = [m.get_emb(key, idx=idx) for m in self.maps]
else:
if not isinstance(idx, list):
idx = [idx]
embs = []
for c_idx in idx:
target_map_idx, target_map_clip_idx = self.clip_idx_2_map_idx(c_idx)
embs.append(
self.maps[target_map_idx].get_emb(key, int(target_map_clip_idx))
)
if len(embs) == 1:
return embs[0]
else:
return self._combine_embs(embs)
@@ -0,0 +1,72 @@
from __future__ import annotations
from typing import List, Union, TYPE_CHECKING
from ..clip.clip_process import (
get_subseq_by_time,
find_time_by_stage,
)
if TYPE_CHECKING:
from ..media_map.media_map import MediaMap
from ..clip import Clip, ClipSeq
__all__ =[
"get_sub_mediamap_by_clip_idx",
"get_sub_mediamap_by_stage",
"get_sub_mediamap_by_time",
]
def get_sub_mediamap_by_time(media_map:MediaMap, start: int=0, end:int=1, eps=1e-2) -> MediaMap:
"""获取子片段序列,同时更新media_map中的相关信息
Args:
media_map (MediaInfo): _description_
start (float): 开始时间
end (float): 结束时间
Returns:
_type_: _description_
"""
if start < 1:
start = media_map.duration * start
if end is None:
end = media_map.meta_info.media_duration
elif end <= 1:
end = media_map.duration * end
media_map.meta_info.start = start
media_map.meta_info.end = end
media_map.clipseq = get_subseq_by_time(
media_map.clipseq,
start=start,
end=end,
)
if media_map.stageseq is not None:
media_map.stageseq = get_subseq_by_time(media_map.stageseq, start=start, end=end)
return media_map
def get_sub_mediamap_by_clip_idx(media_map: MediaMap, start: int=None, end: int=None) -> MediaMap:
"""不仅获取子片段序列,还要更新media_map中的相关信息
Args:
media_map (_type_): _description_
"""
if start is None:
start = 0
if end is None:
end = -1
start = media_map.clipseq[start].time_start
end = media_map.clipseq[end].time_end
media_map = get_sub_mediamap_by_time(media_map=media_map, start=start, end=end)
return media_map
def get_sub_mediamap_by_stage(media_map: MediaMap, stages: Union[str, List[str]]) -> MediaMap:
if isinstance(stages, List):
stages = [stages]
start, _ = find_time_by_stage(media_map.stageseq, stages[0])
_, end = find_time_by_stage(media_map.stageseq, stages[-1])
media_map = get_sub_mediamap_by_time(media_map=media_map, start=start, end=end)
return media_map
+6
View File
@@ -0,0 +1,6 @@
from .music_map.music_map import MusicMap, MusicMapSeq
from .music_map.music_clip import MusicClip, MusicClipSeq
from .music_map.meta_info import MusicMetaInfo
from .music_map.load_music_map import load_music_map
from .utils.path_util import get_audio_path_dct
+82
View File
@@ -0,0 +1,82 @@
import numpy as np
from librosa.core.audio import get_duration
from ...data.clip.clip_process import insert_endclip, insert_startclip
from .clip_process import filter_clipseq_target_point
from .music_clip import MusicClip, MusicClipSeq
def beatnet2TMEType(beat: np.array, duration: float) -> MusicClipSeq:
"""conver beatnet beat to tme beat type
Args:
beat (np.array): Nx2,
1st column is time,
2rd is type,
0, end point
1, strong beat
2,3,4 weak beat
-1 lyric
duration (float): audio time length
Returns:
MusicClipSeq:
"""
n = len(beat)
beat = np.insert(beat, 0, 0, axis=0)
beat = np.insert(beat, n + 1, [duration, 0], axis=0)
clips = []
for i in range(n + 1):
beat_type = int(beat[i + 1, 1])
clip = MusicClip(
time_start=beat[i, 0], # 开始时间
duration=round(beat[i + 1, 0] - beat[i, 0], 3), # 片段持续时间
clipid=i, # 片段序号,
timepoint_type=beat_type,
)
clips.append(clip)
clipseq = MusicClipSeq(clips=clips)
return clipseq
def generate_beatseq_with_beatnet(audio_path: str) -> np.array:
"""使用beatnet生成beat序列
Args:
audio_path (str):
Returns:
np.array: beat序列 Nx2,
1st column is time,
2rd is type,
0, end point
1, strong beat
2,3,4 weak beat
"""
from BeatNet.BeatNet import BeatNet
estimator = BeatNet(1, mode="offline", inference_model="DBN", plot=[], thread=False)
output = estimator.process(audio_path=audio_path)
return output
def generate_music_map_with_beatnet(
audio_path: str, target: list = [0, 1]
) -> MusicClipSeq:
"""使用beatnet生成beat MusicClipseq
Args:
audio_path (str):
target (list, optional): 只保留相应的拍点. Defaults to [0, 1].
Returns:
MusicClipSeq: 返回的beat序列
beat: np.array, 原始的beat检测结果
"""
output = generate_beatseq_with_beatnet(audio_path)
duration = get_duration(filename=audio_path)
clipseq = beatnet2TMEType(output, duration)
clipseq = insert_startclip(clipseq)
clipseq = insert_endclip(clipseq, duration)
clipseq = filter_clipseq_target_point(clipseq, target=target)
return clipseq, output
+196
View File
@@ -0,0 +1,196 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Dict, List
import numpy as np
from ...data.clip.clip_process import find_idx_by_time, reset_clipseq_id
from ...data.clip.clip_fusion import fuse_clips
from ...utils.util import merge_list_continuous_same_element
if TYPE_CHECKING:
from .music_clip import MusicClip, MusicClipSeq
from .music_map import MusicMap, MusicMapSeq
# TODO: 待和clip操作做整合
def music_clip_is_short(clip: MusicClip, th: float = 3) -> bool:
"""判断音乐片段是否过短
Args:
clip (MusicClip): 待判断的音乐片段
th (float, optional): 短篇的参数. Defaults to 3.
Returns:
bool: 是或不是 短片段
"""
if clip.duration < th:
return False
else:
return True
def music_clip_timepoint_is_target(clip: MusicClip, target: list = [-1, 1, 0]) -> bool:
"""音乐片段的关键点类型是否是目标关键点
关键点类型暂时参考:VideoMashup/videomashup/data_structure/music_data_structure.py
Args:
clip (MusicClip): 待判断的音乐片段
target (list, optional): 目标关键点类别. Defaults to [-1, 1, 0].
Returns:
bool: 是还是不是
"""
timepoint = clip.timepoint_type
if isinstance(timepoint, int):
timepoint = {timepoint}
else:
timepoint = {int(x) for x in timepoint.split("_")}
if timepoint & set(target):
return True
else:
return False
def filter_clipseq_target_point(
clipseq: MusicClipSeq, target: list = [-1, 1, 0]
) -> MusicClipSeq:
"""删除目标关键点之外的点,对相应的片段做融合
Args:
clipseq (MusicClipSeq): 待处理的音乐片段序列
target (list, optional): 保留的目标关键点. Defaults to [-1, 1, 0].
Returns:
MusicClipSeq: 处理后的音乐片段序列
"""
n_clipseq = len(clipseq)
if n_clipseq == 1:
return clipseq
newclipseq = []
start_clip = clipseq[0]
if music_clip_timepoint_is_target(start_clip, target=target):
has_start_clip = True
else:
has_start_clip = False
i = 1
while i <= n_clipseq - 1:
clip = clipseq[i]
start_clip_is_target = music_clip_timepoint_is_target(start_clip, target=target)
next_clip_is_target = music_clip_timepoint_is_target(clip, target=target)
# logger.debug("filter_clipseq_target_point: i={},start={}, clip={}".format(i, start_clip["timepoint_type"], clip["timepoint_type"]))
# logger.debug("start_clip_is_target: {}, next_clip_is_target {}".format(start_clip_is_target, next_clip_is_target))
if not has_start_clip:
start_clip = clip
has_start_clip = next_clip_is_target
else:
if start_clip_is_target:
has_start_clip = True
if next_clip_is_target:
newclipseq.append(start_clip)
start_clip = clip
if i == n_clipseq - 1:
newclipseq.append(clip)
else:
start_clip = fuse_clips(start_clip, clip)
if i == n_clipseq - 1:
newclipseq.append(start_clip)
# logger.debug("filter_clipseq_target_point: fuse {}, {}".format(i, clip["timepoint_type"]))
else:
start_clip = clip
i += 1
newclipseq = reset_clipseq_id(newclipseq)
return newclipseq
def merge_musicclip_into_clipseq(
clip: MusicClipSeq, clipseq: MusicClip, th: float = 1
) -> MusicClipSeq:
"""给clipseq插入一个新的音乐片段,会根据插入后片段是否过短来判断。
Args:
clip (MusicClipSeq): 要插入的音乐片段
clipseq (MusicClip): 待插入的音乐片段序列
th (float, optional): 插入后如果受影响的片段长度过短,则放弃插入. Defaults to 1.
Returns:
MusicClipSeq: _description_
"""
n_clipseq = len(clipseq)
clip_time = clip.time_start
idx = find_idx_by_time(clipseq, clip_time)
last_clip_time_start = clipseq[idx].time_start
next_clip_time_start = clipseq[idx].time_start + clipseq[idx].duration
last_clip_time_delta = clip_time - last_clip_time_start
clip_duration = next_clip_time_start - clip_time
# TODO: 副歌片段改变th参数来提升音符密度,暂不使用,等待音游谱面
# TODO: 待抽离独立的业务逻辑为单独的函数
# 只针对副歌片段插入关键点
if clipseq[idx].text is None or (
clipseq[idx].text is not None
and clipseq[idx].stage is not None
and "C" in clipseq[idx].stage
):
if (last_clip_time_delta > th) and (clip_duration > th):
clip.duration = clip_duration
clipseq[idx].duration = last_clip_time_delta
clipseq.insert(idx + 1, clip)
clipseq = reset_clipseq_id(clipseq)
return clipseq
def merge_music_clipseq(clipseq1: MusicClipSeq, clipseq2: MusicClipSeq) -> MusicClipSeq:
"""将片段序列clipseq2融合到音乐片段序列clipseq1中。融合过程也会判断新片段长度。
Args:
clipseq1 (MusicClipSeq): 要融合的目标音乐片段序列
clipseq2 (MusicClipSeq): 待融合的音乐片段序列
Returns:
MusicClipSeq: 融合后的音乐片段序列
"""
while len(clipseq2) > 0:
clip = clipseq2[0]
clipseq1 = merge_musicclip_into_clipseq(clip, clipseq1)
del clipseq2[0]
return clipseq1
def merge_lyricseq_beatseq(
lyric_clipseq: MusicClipSeq, beat_clipseq: MusicClipSeq
) -> MusicClipSeq:
"""将beat序列融合到歌词序列中
Args:
lyric_clipseq (MusicClipSeq): 歌词序列
beat_clipseq (MusicClipSeq): beat序列
Returns:
MusicClipSeq: 融合后的音乐片段序列
"""
newclipseq = merge_music_clipseq(lyric_clipseq, beat_clipseq)
# for i, clip in enumerate(newclipseq):
# logger.debug("i={}, time_start={}, duration={}".format(i, clip.time_start, clip.duration))
return newclipseq
def get_stageseq_from_clipseq(clipseq: MusicClipSeq) -> List[Dict]:
"""对clip.stage做近邻融合,返回总时间
Returns:
List[Dict]: 根据音乐结构进行分割的片段序列
"""
stages = [clip.stage for clip in clipseq]
merge_stages_idx = merge_list_continuous_same_element(stages)
merge_stages = []
for n, stages_idx in enumerate(merge_stages_idx):
dct = {
"clipid": n,
"time_start": clipseq[stages_idx["start"]].time_start,
"time_end": clipseq[stages_idx["end"]].time_end,
"stage": stages_idx["element"],
"original_clipid": list(
range(stages_idx["start"], stages_idx["end"] + 1)
), # mss都是左闭、 右闭的方式
}
dct["duration"] = dct["time_end"] - dct["time_start"]
merge_stages.append(dct)
return merge_stages
+57
View File
@@ -0,0 +1,57 @@
from ...data.clip.clip_process import (
insert_startclip,
insert_endclip,
reset_clipseq_id,
)
from .music_clip import MusicClip, MusicClipSeq
def read_osu_hitobjs(path: str) -> list:
"""读取osu的音游谱面
Args:
path (str): 谱面低质
Returns:
list: 只包含HitObjects的行字符串信息
"""
lines = []
is_hit_info_start = False
with open(path, "r") as f:
for line in f:
if is_hit_info_start:
lines.append(line.strip())
if "[HitObjects]" in line:
is_hit_info_start = True
return lines
def osu2itech(src: list, duration: float = None) -> MusicClipSeq:
"""将osu的音游谱面转换为我们的目标格式
Args:
src (list): 音游谱面路径或者是读取的目标行字符串列表
duration (float, optional): 歌曲长度. Defaults to None.
Returns:
MusicClipSeq: 音乐片段序列
"""
if isinstance(src, str):
src = read_osu_hitobjs(src)
timepoints = [float(line.split(",")[2]) for line in src]
clips = []
for i in range(len(timepoints) - 1):
clip = MusicClip(
time_start=round(timepoints[i] / 1000, 3),
timepoint_type=0,
duration=round((timepoints[i + 1] - timepoints[i]) / 1000, 3),
clipid=i,
)
clips.append(clip)
if len(clips) > 0:
clips = insert_startclip(clips)
if duration is not None:
clips = insert_endclip(clips, duration=duration)
clips = reset_clipseq_id(clips)
return MusicClipSeq(clips)
@@ -0,0 +1,38 @@
from typing import List
from .music_map import MusicMap, MusicMapSeq
def load_music_map(
music_map_paths,
music_paths,
emb_paths,
start: float=None,
end: None=None,
target_stages: List[str] = None,
**kwargs,
):
"""读取视频谱面,转化成MusicInfo。当 musicinfo_path_lst 为列表时,表示多歌曲
Args:
musicinfo_path_lst (str or [str]): 视频谱面路径文件列表
music_path_lst (str or [str]): 视频文件路径文件列表,须与musicinfo_path_lst等长度
Returns:
MusicInfo: 视频谱面信息
"""
dct ={
"start": start,
"end": end,
"target_stages": target_stages,
}
if isinstance(music_map_paths, list):
music_map = MusicMapSeq.from_json_paths(media_map_class=MusicMapSeq, media_paths=music_paths, media_map_paths=music_map_paths, emb_paths=emb_paths, **dct, **kwargs)
if len(music_map) == 1:
music_map = music_map[0]
else:
music_map = MusicMap.from_json_path(path=music_map_paths, emb_path=emb_paths, media_path=music_paths, **dct, **kwargs)
return music_map
+149
View File
@@ -0,0 +1,149 @@
import numpy as np
from sklearn.preprocessing import normalize, minmax_scale
from scipy.signal import savgol_filter
# TODO:待更新音乐谱面的类信息
from ...data.clip.clip_process import (
complete_clipseq,
find_idx_by_clip,
insert_endclip,
insert_startclip,
reset_clipseq_id,
)
from .music_clip import Clip, ClipSeq
from .music_clip import MusicClipSeq
from .music_map import MusicMap
def generate_lyric_map(
path: str, duration: float = None, gap_th: float = 2
) -> MusicClipSeq:
"""从歌词文件中生成音乐谱面
Args:
path (str): 歌词文件路径
duration (float, optional): 歌词对应音频的总时长. Defaults to None.
gap_th (float, optional): 歌词中间的空白部分是否融合到上一个片段中. Defaults to 3.
Returns:
MusicClipSeq: 以歌词文件生成的音乐谱面
"""
from ..music_map.lyric_process import lyricfile2musicinfo
lyric_info = lyricfile2musicinfo(path)
lyric_info = MusicMap(lyric_info, duration=duration)
clipseq = lyric_info.clipseq
lyric_info.meta_info.duration = duration
# set part of nonlyric as clip whose timepoint is 0
for i in range(len(clipseq)):
clipseq[i].timepoint_type = -1
lyric_info.clipseq = complete_clipseq(
clipseq=clipseq, duration=duration, gap_th=gap_th
)
return lyric_info
def insert_field_2_clipseq(clipseq: ClipSeq, reference: ClipSeq, field: str) -> ClipSeq:
"""将reference中每个clip的字段信息根据赋给clipseq中最近的clip
Args:
clipseq (ClipSeq): 目标clip序列
reference (ClipSeq): 参考clip序列
field (str): 目标字段
Returns:
ClipSeq: 更新目标字段新值后的clip序列
"""
for i, clip in enumerate(clipseq):
idx = find_idx_by_clip(reference, clip=clip)
if idx is not None:
if getattr(reference[idx], field) is not None:
clipseq[i].__dict__[field] = getattr(reference[idx], field)
return clipseq
def insert_rythm_2_clipseq(clipseq, reference):
"""参考MSS字段的结构信息设置rythm信息。目前策略非常简单,主歌(Vx)0.25,副歌(Cx)0.75,其他为None
Args:
clipseq (ClipSeq): 目标clip序列,设置rythm字段
reference (ClipSeq): 参考clip序列,参考stage字段
Returns:
ClipSeq: 更新rythm字段新值后的clip序列
"""
def stage2rythm(stage):
if "V" in stage:
return 0.25
elif "C" in stage:
return 0.75
else:
return None
for i, clip in enumerate(clipseq):
idx = find_idx_by_clip(reference, clip=clip)
if idx is not None:
if reference[idx].rythm is not None:
clipseq[i].rythm = stage2rythm(reference[idx].stage)
return clipseq
def insert_rythm_from_clip(clipseq: MusicClipSeq, beat: np.array) -> MusicClipSeq:
"""给MusicClipSeq中的每个Clip新增节奏信息。目前使用
1. 单位时间内的歌词数量特征, 使用 min-max 归一化到 0 - 1 之间
2. 单位时间内的关键点数量,目前使用beatnet,使用 min-max 归一化到 0 - 1 之间
3. 对1、2中的特征相加,并根据歌曲结构不同进行加权
Args:
clipseq (MusicClipSeq): 待处理的 MusicClipSeq
beat (np.array): beat检测结果,Nx2,,用于结算单位时间内的关键点数。
1st column is time,
2rd is type,
0, end point
1, strong beat
2,3,4 weak beat
Returns:
MusicClipSeq: 新增 rythm 的 MusicClipSeq
"""
mss_cofficient = {
"intro": 1.0,
"bridge": 1.0,
"end": 0.8,
"VA": 1.0,
"VB": 1.0,
"CA": 1.6,
"CB": 1.6,
}
# text_num_per_second
text_num_per_second_lst = [clip.tnps for clip in clipseq if clip.tnps != 0]
common_tnps = np.min(text_num_per_second_lst)
tnps = np.array([clip.tnps if clip.tnps != 0 else common_tnps for clip in clipseq])
tnps = minmax_scale(tnps)
# beat point _num_per_second
beat_pnps = np.zeros(len(clipseq))
for i, clip in enumerate(clipseq):
time_start = clip.time_start
time_end = clip.time_end
target_beat = beat[(beat[:, 0] >= time_start) & (beat[:, 0] < time_end)]
beat_pnps[i] = len(target_beat) / clip.duration
beat_pnps = minmax_scale(beat_pnps)
# cofficient
cofficients = np.array(
[
mss_cofficient[clip.stage]
if clip.stage in mss_cofficient and clip.stage is not None
else 1.0
for clip in clipseq
]
)
rythm = cofficients * (tnps + beat_pnps)
rythm = minmax_scale(rythm)
rythm = savgol_filter(rythm, window_length=5, polyorder=3)
rythm = minmax_scale(rythm)
for i, clip in enumerate(clipseq):
clip.dynamic = rythm[i]
return clipseq
+515
View File
@@ -0,0 +1,515 @@
from genericpath import isfile
import re
import os
from ...text.utils.read_text import read_xml2json
# 一个正则表达式非常好用的网站
# https://regex101.com/r/cW8jA6/2
CHINESE_PATTERN = r"[\u4e00-\u9fff]+"
NOT_CHINESE_PATTERN = r"[^\u4e00-\u9fa5]"
ENGLISH_CHARACHTER_PATTERN = r"[a-zA-Z]+"
WORD_PATTERN = r"\w+" # equal to [a-zA-Z0-9_].
NOT_WORD_PATTERN = r"\W+"
def has_target_string(lyric: str, pattern: str) -> bool:
"""本句歌词是否有目标字符串
Args:
lyric (str):
pattern (str): 目标字符串的正则表达式式patteren
Returns:
bool: 有没有目标字符串
"""
matched = re.findall(pattern, lyric)
flag = len(matched) > 0
return flag
def has_chinese_char(lyric: str) -> bool:
"""是否有中文字符
Args:
lyric (str):
Returns:
bool: 是否有中文字符
"""
return has_target_string(lyric, CHINESE_PATTERN)
def has_non_chinese_char(lyric: str) -> bool:
"""是否有非中文字符,参考https://git.woa.com/innovative_tech/CopyrightGroup/LyricTools/blob/master/lyric_tools/dataProcess.py#L53
Args:
lyric (str):
Returns:
bool: 是否有中文字符
"""
return has_target_string(lyric, NOT_CHINESE_PATTERN)
def has_english_alphabet_char(lyric: str) -> bool:
"""是否有英文字母表字符
Args:
lyric (str):
Returns:
bool:
"""
return has_target_string(lyric, ENGLISH_CHARACHTER_PATTERN)
def check_is_lyric_row(lyric: str) -> bool:
"""该字符串是否是歌词
Args:
lyric (str): 待判断的字符串
Returns:
bool: 该字符串是否是歌词
"""
is_not_lyric = [
re.search(r"\[ti[::]?", lyric),
re.search(r"\[ar[::]?", lyric),
re.search(r"\[al[::]?", lyric),
re.search(r"\[by[::]?", lyric),
re.search(r"\[offset[::]?", lyric),
re.search(r"词[::]?\(\d+,\d+\)[::]?", lyric),
re.search(r"曲[::]?\(\d+,\d+\)[::]?", lyric),
re.search(r"作\(\d+,\d+\)词[::]?", lyric),
re.search(r"作\(\d+,\d+\)曲[::]?", lyric),
re.search(r"演\(\d+,\d+\)唱[::]?", lyric),
re.search(r"编\(\d+,\d+\)曲[::]?", lyric),
re.search(r"吉\(\d+,\d+\)他[::]", lyric),
re.search(r"人\(\d+,\d+\)声\(\d+,\d+\)录\(\d+,\d+\)音\(\d+,\d+\)师[::]?", lyric),
re.search(r"人\(\d+,\d+\)声\(\d+,\d+\)录\(\d+,\d+\)音\(\d+,\d+\)棚[::]?", lyric),
re.search(r"Vocal\s+\(\d+,\d+\)edite[::]?", lyric),
re.search(r"混\(\d+,\d+\)音\(\d+,\d+\)/\(\d+,\d+\)母\(\d+,\d+\)带[::]?", lyric),
re.search(r"混\(\d+,\d+\)音", lyric),
re.search(r"和\(\d+,\d+\)声\(\d+,\d+\)编\(\d+,\d+\)写[::]?", lyric),
re.search(
r"词\(\d+,\d+\)版\(\d+,\d+\)权\(\d+,\d+\)管\(\d+,\d+\)理\(\d+,\d+\)方[::]?", lyric
),
re.search(
r"曲\(\d+,\d+\)版\(\d+,\d+\)权\(\d+,\d+\)管\(\d+,\d+\)理\(\d+,\d+\)方[::]?", lyric
),
re.search(r"联\(\d+,\d+\)合\(\d+,\d+\)出\(\d+,\d+\)品[::]?", lyric),
re.search(r"录\(\d+,\d+\)音\(\d+,\d+\)作\(\d+,\d+\)品", lyric),
re.search(
r"录\(\d+,\d+\)音\(\d+,\d+\)作\(\d+,\d+\)品\(\d+,\d+\)监\(\d+,\d+\)制[::]?", lyric
),
re.search(r"制\(\d+,\d+\)作\(\d+,\d+\)人[::]?", lyric),
re.search(r"制\(\d+,\d+\)作\(\d+,\d+\)人[::]?", lyric),
re.search(r"不\(\d+,\d+\)得\(\d+,\d+\)翻\(\d+,\d+\)唱", lyric),
re.search(r"未\(\d+,\d+\)经\(\d+,\d+\)许\(\d+,\d+\)可", lyric),
re.search(r"酷\(\d+,\d+\)狗\(\d+,\d+\)音\(\d+,\d+\)乐", lyric),
re.search(r"[::]", lyric),
]
is_not_lyric = [x is not None for x in is_not_lyric]
is_not_lyric = any(is_not_lyric)
is_lyric = not is_not_lyric
return is_lyric
def lyric2clip(lyric: str) -> dict:
"""convert a line of lyric into a clip
Clip定义可以参考 https://git.woa.com/innovative_tech/VideoMashup/blob/master/videomashup/media/clip.py
Args:
lyric (str): _description_
Returns:
dict: 转化成Clip 字典
"""
time_str_groups = re.findall(r"\d+,\d+", lyric)
line_time_start = round(int(time_str_groups[0].split(",")[0]) / 1000, 3)
line_duration = round(int(time_str_groups[0].split(",")[-1]) / 1000, 3)
line_end_time = line_time_start + line_duration
last_word_time_start = round(int(time_str_groups[-1].split(",")[0]) / 1000, 3)
last_word_duration = round(int(time_str_groups[-1].split(",")[-1]) / 1000, 3)
last_word_end_time = last_word_time_start + last_word_duration
actual_duration = min(line_end_time, last_word_end_time) - line_time_start
lyric = re.sub(r"\[\d+,\d+\]", "", lyric)
# by yuuhong: 把每个字的起始时间点、结束时间点、具体的字拆分出来
words_with_timestamp = get_words_with_timestamp(lyric)
lyric = re.sub(r"\(\d+,\d+\)", "", lyric)
dct = {
"time_start": line_time_start,
"duration": actual_duration,
"text": lyric,
"original_text": lyric,
"timepoint_type": -1,
"clips": words_with_timestamp,
}
return dct
# by yuuhong
# 把一句QRC中的每个字拆分出来
# lyric示例:漫(17316,178)步(17494,174)走(17668,193)在(17861,183) (18044,0)莎(18044,153)玛(18197,159)丽(18356,176)丹(18532,200)
def get_words_with_timestamp(lyric):
words_with_timestamp = []
elements = lyric.split(")")
for element in elements:
sub_elements = element.split("(")
if len(sub_elements) != 2:
continue
text = sub_elements[0]
timestamp = sub_elements[1]
if re.match(r"\d+,\d+", timestamp):
# 有效时间戳
time_start_str = timestamp.split(",")[0]
time_start = round(int(time_start_str) / 1000, 3)
duration_str = timestamp.split(",")[1]
duration = round(int(duration_str) / 1000, 3)
clip = {"text": text, "time_start": time_start, "duration": duration}
words_with_timestamp.append(clip)
return words_with_timestamp
def lyric2clips(lyric: str, th: float = 0.75) -> list:
"""将一句歌词转换为至少1个的clip。拆分主要是针对中文空格拆分,如果拆分后片段过短,也会整句处理。
Args:
lyric (str): such as [173247,3275]去(173247,403)吗(173649,677) 配(174326,189)吗(174516,593) 这(175108,279)
th (float, optional): 后面如果拆分后片段过短,也会整句处理. Defaults to 1.0.
Returns:
list: 歌词Clip序列
"""
# 目前只对中文的一句歌词按照空格拆分,如果是英文空格则整句处理
# 后面如果拆分后片段过短,也会整句处理
if has_english_alphabet_char(lyric):
return [lyric2clip(lyric)]
splited_lyric = lyric.split(" ")
if len(splited_lyric) == 1:
return [lyric2clip(splited_lyric[0])]
line_time_str, sub_lyric = re.split(r"]", splited_lyric[0])
line_time_groups = re.findall(r"\d+,\d+", line_time_str)
line_time_start = round(int(line_time_groups[0].split(",")[0]) / 1000, 3)
line_duration = round(int(line_time_groups[0].split(",")[-1]) / 1000, 3)
splited_lyric[0] = sub_lyric
# 歌词xml都是歌词仅跟着时间,如果有空格 空格也应该是在时间后面,但有时候空格却在字后面、在时间前,因此需要修正
# 错误的:[173247,3275]去(173247,403)吗 (173649,677)配(174326,189)吗 (174516,593)这(175108,279)
# 错误的:[46122,2082]以(46122,213)身(46335,260)淬(46595,209)炼(46804,268)天(47072,250)地(47322,370)造(47692,341)化 (48033,172)
# 修正成:[173247,3275]去(173247,403)吗(173649,677) 配(174326,189)吗(174516,593) 这(175108,279)
for i in range(len(splited_lyric)):
if splited_lyric[i] == "":
del splited_lyric[i]
break
if splited_lyric[i][-1] != ")":
next_lyric_time_start = re.search(
r"\(\d+,\d+\)", splited_lyric[i + 1]
).group(0)
splited_lyric[i] += next_lyric_time_start
splited_lyric[i + 1] = re.sub(
next_lyric_time_start, "", splited_lyric[i + 1]
)
splited_lyric[i + 1] = re.sub("\(\)", "", splited_lyric[i + 1])
lyric_text = re.sub(r"\[\d+,\d+\]", "", lyric)
lyric_text = re.sub(r"\(\d+,\d+\)", "", lyric_text)
clips = []
has_short_clip = False
for sub_lyric in splited_lyric:
sub_lyric_groups = re.findall(r"\d+,\d+", sub_lyric)
sub_lyric_1st_word_time_start = round(
int(sub_lyric_groups[0].split(",")[0]) / 1000, 3
)
sub_lyric_last_word_time_start = round(
int(sub_lyric_groups[-1].split(",")[0]) / 1000, 3
)
sub_lyric_last_word_duration = round(
int(sub_lyric_groups[-1].split(",")[-1]) / 1000, 3
)
sub_lyric_last_word_time_end = (
sub_lyric_last_word_time_start + sub_lyric_last_word_duration
)
sub_lyric_duration = (
sub_lyric_last_word_time_end - sub_lyric_1st_word_time_start
)
if sub_lyric_duration <= th:
has_short_clip = True
break
sub_lyric_text = re.sub(r"\[\d+,\d+\]", "", sub_lyric)
sub_lyric_text = re.sub(r"\(\d+,\d+\)", "", sub_lyric_text)
# 使用原始lyric,而不是sub_lyric_text 主要是保留相关clip的歌词信息,便于语义连续
dct = {
"time_start": sub_lyric_1st_word_time_start,
"duration": sub_lyric_duration,
"text": sub_lyric_text,
"original_text": lyric_text,
"timepoint_type": -1,
}
clips.append(dct)
if has_short_clip:
clips = [lyric2clip(lyric)]
return clips
def is_songname(lyric: str) -> bool:
"""是否是歌名,歌名文本含有ti, 如[ti:霍元甲 (《霍元甲》电影主题曲)]
Args:
lyric (str):
Returns:
bool:
"""
return has_target_string(lyric, r"\[ti[::]?")
def get_songname(lyric: str) -> str:
"""获取文本中的歌名,输入必须类似[ti:霍元甲 (《霍元甲》电影主题曲)]
Args:
lyric (str): 含有歌名的QRC文本行
Returns:
str: 歌名
"""
return lyric.split("(")[0][4:-1]
def is_album(lyric: str) -> bool:
"""是否含有专辑名,文本必须类似[al:霍元甲]
Args:
lyric (str): _description_
Returns:
bool: _description_
"""
return has_target_string(lyric, r"\[al[::]?")
def get_album(lyric: str) -> str:
"""提取专辑名,文本必须类似[al:霍元甲]
Args:
lyric (str): 含有专辑名的QRC文本行
Returns:
str: 专辑名
"""
return lyric[4:-1]
def is_singer(lyric: str) -> bool:
"""是否有歌手名,目标文本类似 [ar:周杰伦]
Args:
lyric (str): _description_
Returns:
bool: _description_
"""
return has_target_string(lyric, r"\[ar[::]?")
def get_singer(lyric: str) -> str:
"""提取歌手信息,文本必须类似[ar:周杰伦]
Args:
lyric (str): 含有歌手名的QRC文本行
Returns:
str: 歌手名
"""
return lyric[4:-1]
def lyric2musicinfo(lyric: str) -> dict:
"""convert lyric content from str into musicinfo, a dict
参考https://git.woa.com/innovative_tech/VideoMashup/blob/master/videomashup/media/media_info.py#L19
{
"meta_info": {},
"sub_meta_info": [],
"clips": [
clip
]
}
Args:
lyric (str): 来自QRC的歌词字符串
Returns:
musicinfo: 音乐谱面字典,https://git.woa.com/innovative_tech/VideoMashup/blob/master/videomashup/media/media_info.py#L19
"""
lyrics = lyric["QrcInfos"]["LyricInfo"]["Lyric_1"]["@LyricContent"]
musicinfo = {
"meta_info": {
"mediaid": None,
"media_name": None,
"singer": None,
},
"sub_meata_info": {},
"clips": [],
}
# lyrics = [line.strip() for line in re.split(r"[\t\n\s+]", lyrics)]
lyrics = ["[" + line.strip() for line in re.split(r"\[", lyrics)]
next_is_title_row = False
lyric_clips = []
for line in lyrics:
if is_songname(line):
musicinfo["meta_info"]["media_name"] = get_songname(line)
continue
if is_singer(line):
musicinfo["meta_info"]["singer"] = get_singer(line)
continue
if is_album(line):
musicinfo["meta_info"]["album"] = get_album(line)
continue
is_lyric_row = check_is_lyric_row(line)
if next_is_title_row:
next_is_title_row = False
continue
# remove tille row
if not next_is_title_row and re.search(r"\[offset[::]", line):
next_is_title_row = True
if is_lyric_row and re.match(r"\[\d+,\d+\]", line):
lyric_clip = lyric2clip(line)
lyric_clips.append(lyric_clip)
clips = lyric2clips(line)
musicinfo["clips"].extend(clips)
musicinfo["meta_info"]["lyric"] = lyric_clips
return musicinfo
def lrc_timestr2time(time_str: str) -> float:
"""提取lrc中的时间戳文本,类似[00:00.00],转化成秒的浮点数
Args:
time_str (str):
Returns:
float: 时间浮点数
"""
m, s, ms = (float(x) for x in re.split(r"[:.]", time_str))
return round((m * 60 + s + ms / 1000), 3)
def get_lrc_line_time(text: str, time_pattern: str) -> str:
"""提取lrc中的时间字符串, 类似 \"[00:00.00]本字幕由天琴实验室独家AI字幕技术生成\"
Args:
text (str): 输入文本
time_pattern (str): 时间字符串正则表达式
Returns:
str: 符合正则表达式的时间信息文本
"""
time_str = re.search(time_pattern, text).group(0)
return lrc_timestr2time(time_str)
def lrc_lyric2clip(lyric: str, time_pattern: str, duration: float) -> dict:
"""将一行lrc文本字符串转化为Clip 字典
Args:
lyric (str): 类似 \"[00:00.00]本字幕由天琴实验室独家AI字幕技术生成\"
time_pattern (str): 时间字符串正则表达式,类似 r"\d+:\d+\.\d+"
duration (float): clip的时长信息,
Returns:
dict: 转化后Clip
Clip定义可以参考 https://git.woa.com/innovative_tech/VideoMashup/blob/master/videomashup/media/clip.py
"""
time_str = get_lrc_line_time(lyric, time_pattern=time_pattern)
text = re.sub(time_pattern, "", lyric)
text = text[2:]
clip = {
"time_start": time_str,
"duration": duration,
"text": text,
"timepoint_type": -1,
}
return clip
def lrc2musicinfo(lyric: str, time_pattern: str = "\d+:\d+\.\d+") -> dict:
"""将lrc转化为音乐谱面
Args:
lyric (str): lrc文本路径
time_pattern (str, optional): lrc时间戳字符串正则表达式. Defaults to "\d+:\d+\.\d+".
Returns:
dict: 生成的音乐谱面字典,定义可参考 https://git.woa.com/innovative_tech/VideoMashup/blob/master/videomashup/music/music_info.py
"""
if isinstance(lyric, str):
if os.path.isfile(lyric):
with open(lyric, "r") as f:
lyric = [line.strip() for line in f.readlines()]
return lrc2musicinfo(lyric)
else:
lyric = lyric.split("\n")
return lrc2musicinfo(lyric)
else:
musicinfo = {
"meta_info": {
"mediaid": None,
"media_name": None,
"singer": None,
},
"sub_meata_info": {},
"clips": [],
}
# lyrics = [line.strip() for line in re.split(r"[\t\n\s+]", lyrics)]
lyric_clips = []
rows = len(lyric)
for i, line in enumerate(lyric):
if is_songname(line):
musicinfo["meta_info"]["media_name"] = line[4:-1]
continue
if is_singer(line):
musicinfo["meta_info"]["singer"] = line[4:-1]
continue
if is_album(line):
musicinfo["meta_info"]["album"] = line[4:-1]
continue
if len(re.findall(time_pattern, line)) > 0:
if i < rows - 1:
time_start = get_lrc_line_time(line, time_pattern=time_pattern)
next_line_time_start = get_lrc_line_time(
lyric[i + 1], time_pattern=time_pattern
)
duration = next_line_time_start - time_start
else:
duration = 1
clip = lrc_lyric2clip(
line, duration=duration, time_pattern=time_pattern
)
musicinfo["clips"].append(clip)
musicinfo["meta_info"]["lyric"] = lyric_clips
return musicinfo
def lyricfile2musicinfo(path: str) -> dict:
"""将歌词文件转化为音乐谱面,歌词文件可以是QRC的xml文件、也可以是lrc对应的lrc文件
TODO: 待支持osu
Args:
path (str): 歌词文件路径
Returns:
dict: 音乐谱面字典,定义可参考 https://git.woa.com/innovative_tech/VideoMashup/blob/master/videomashup/music/music_info.py
"""
filename, ext = os.path.basename(path).split(".")
if ext == "xml":
lyric = read_xml2json(path)
musicinfo = lyric2musicinfo(lyric)
elif ext == "lrc":
musicinfo = lrc2musicinfo(path)
musicinfo["meta_info"]["mediaid"] = filename
return musicinfo
+21
View File
@@ -0,0 +1,21 @@
from __future__ import annotations
from ...data import MetaInfo
class MusicMetaInfo(MetaInfo):
def __init__(self, mediaid=None, media_name=None, media_duration=None, signature=None, media_path: str = None, media_map_path: str = None,
singer=None,
lyric_path=None,
genre=None,
language=None,
start: float = None, end: float = None, ext=None, **kwargs):
super().__init__(mediaid, media_name, media_duration, signature, media_path, media_map_path, start, end, ext, **kwargs)
self.singer = singer
self.genre = genre
self.language = language
self.lyric_path = lyric_path
@classmethod
def from_data(cls, data) -> MusicMetaInfo:
return MusicMetaInfo(**data)
+185
View File
@@ -0,0 +1,185 @@
import logging
from .music_clip import MusicClip, MusicClipSeq
from .music_map import MusicMap
from ...data.clip.clip_process import find_idx_by_time
logger = logging.getLogger(__name__) # pylint: disable=invalid-name
def insert_mss_2_clipseq(
clipseq: MusicClipSeq, mss_clipseq: MusicClipSeq
) -> MusicClipSeq:
"""将mss中的结构字段信息赋予到目标clipseq中的最近clip
Args:
clipseq (ClipSeq): 目标clip序列
reference (ClipSeq): 参考clip序列
field (str): 目标字段
Returns:
ClipSeq: 更新目标字段新值后的clip序列
"""
for i, clip in enumerate(clipseq):
idx = find_idx_by_time(mss_clipseq, clip.time_start)
if idx is not None:
clipseq[i].stage = mss_clipseq[idx].stage
else:
clipseq[i].stage = "unknow"
return clipseq
def get_mss_musicinfo(songid: str) -> MusicMap:
"""通过调用media_data中的接口 获取天琴实验室的歌曲结构信息
Args:
songid (str): 歌词id
Returns:
MusicMap: mss结构信息生成的音乐谱面
"""
try:
from media_data.oi.tianqin_database import get_mss
mss = get_mss(songid=songid)
except Exception as e:
logger.warning("get mss failed, mss={}".format(songid))
logger.exception(e)
mss = None
mss_musicinfo = MusicMap(mss) if mss is not None else None
return mss_musicinfo
def merge_mss(musicinfo: MusicMap, mss: MusicMap) -> MusicMap:
"""融合mss音乐谱面到目标音乐谱面
Args:
musicinfo (MusicMap): 目标音乐谱面
mss (MusicMap): 待融合的mss音乐谱面
Returns:
MusicMap: 融合后的音乐谱面
"""
musicinfo.meta_info.bpm = mss.meta_info.bpm
if len(mss.clipseq) > 0:
musicinfo.clipseq = insert_mss_2_clipseq(musicinfo.clipseq, mss.clipseq)
return musicinfo
def generate_mss_from_lyric(lyrics: list, audio_duration: float, th=8) -> MusicClipSeq:
# "intro", "VA", "CA", "bridge", "VB", "CB", "end"]
mss = []
n_lyric = len(lyrics)
for lyric_idx, line_lyric_dct in enumerate(lyrics):
time_start = line_lyric_dct["time_start"]
duration = line_lyric_dct["duration"]
time_end = time_start + duration
# text = line_lyric_dct["text"]
if lyric_idx == 0:
sub_mss = {
"stage": "intro",
"time_start": 0,
"duration": time_start,
}
mss.append(sub_mss)
continue
if lyric_idx == n_lyric - 1:
sub_mss = {
"stage": "end",
"time_start": time_end,
"duration": audio_duration - time_end,
}
mss.append(sub_mss)
continue
if lyrics[lyric_idx + 1]["time_start"] - time_end >= th:
sub_mss = {
"stage": "bridge",
"time_start": time_end,
"duration": lyrics[lyric_idx + 1]["time_start"] - time_end,
}
mss.append(sub_mss)
mss_lyric = []
for sub_idx, sub_mss in enumerate(mss):
if sub_idx == len(mss) - 1:
continue
time_end = sub_mss["time_start"] + sub_mss["duration"]
next_time_start = mss[sub_idx + 1]["time_start"]
if next_time_start - time_end > 0.1:
mss_lyric.append(
{
"stage": "lyric",
"time_start": time_end,
"duration": next_time_start - time_end,
}
)
mss.extend(mss_lyric)
mss = sorted(mss, key=lambda x: x["time_start"])
mss = MusicClipSeq(mss)
return mss
def refine_mss_info_from_tianqin(
mss_info: MusicMap, lyricseq: MusicClipSeq
) -> MusicMap:
"""优化天琴的歌曲结信息,
优化前:天琴歌曲结构里面只有每句歌词和结构信息,时间前后不连续,对于整首歌去时间结构不完备。
优化后:增加intro,bridge,end,将相近的结构信息合并,时间前后连续,时间完备
Args:
mss_info (MusicMap): 天琴歌曲结构
lyricseq (ClipSeq): 原始歌曲信息,用于计算Intro,bridge,end。其实也可以从mss_info中获取。
Returns:
MusicMap: 优化后的歌曲结构信息
"""
lyric_mss_clipseq = generate_mss_from_lyric(
lyricseq, audio_duration=mss_info.meta_info.duration
)
new_mss_clipseq = []
# lyric_mss_dct = lyric_mss_clipseq.to_dct()
# mss_dct = mss_info.clipseq.to_dct()
for l_clip_idx, lyric_clip in enumerate(lyric_mss_clipseq):
if lyric_clip.stage != "lyric":
new_mss_clipseq.append(lyric_clip)
else:
new_clip_time_start = lyric_clip.time_start
last_stage = "ANewClipStart"
for clip_idx, clip in enumerate(mss_info.clipseq):
if clip.time_start < new_clip_time_start:
continue
if (
clip.time_start >= lyric_mss_clipseq[l_clip_idx + 1].time_start
or clip_idx == len(mss_info.clipseq) - 1
):
if clip.time_start >= lyric_mss_clipseq[l_clip_idx + 1].time_start:
stage = last_stage
# 像偶阵雨这首歌最后一个歌词段落 只有一句歌词
if clip_idx == len(mss_info.clipseq) - 1:
stage = clip.stage
new_clip_time_end = lyric_mss_clipseq[l_clip_idx + 1].time_start
new_stage_clip = {
"time_start": new_clip_time_start,
"duration": new_clip_time_end - new_clip_time_start,
"stage": stage,
}
new_mss_clipseq.append(MusicClip(**new_stage_clip))
new_clip_time_start = new_clip_time_end
last_stage = clip.stage
break
if clip.stage != last_stage:
if last_stage == "ANewClipStart":
last_stage = clip.stage
continue
new_clip_time_end = mss_info.clipseq[clip_idx].time_start
new_stage_clip = {
"time_start": new_clip_time_start,
"duration": new_clip_time_end - new_clip_time_start,
"stage": last_stage,
}
new_mss_clipseq.append(MusicClip(**new_stage_clip))
new_clip_time_start = new_clip_time_end
last_stage = clip.stage
new_mss_clipseq = MusicClipSeq(sorted(new_mss_clipseq, key=lambda x: x.time_start))
mss_info.clipseq = new_mss_clipseq
return mss_info
+83
View File
@@ -0,0 +1,83 @@
from __future__ import annotations
from typing import Dict, List
from ...data.clip import Clip, ClipSeq
class MusicClip(Clip):
def __init__(self, time_start: float, duration: float, clipid: int = None, media_type: str = None, mediaid: str = None, timepoint_type: str = None, text: str = None, stage: str = None, path: str = None, duration_num: int = None, similar_clipseq: MatchedClipIds = None, dynamic: float = None, **kwargs):
super().__init__(time_start, duration, clipid, media_type, mediaid, timepoint_type, text, stage, path, duration_num, similar_clipseq, dynamic, **kwargs)
@property
def text_num(self):
return self._cal_text_num()
@property
def original_text_num(self):
return self._cal_text_num(text_mode=1)
def _cal_text_num(self, text_mode: int = 0) -> int:
"""计算 文本 字的数量
Args:
text_mode (int, optional): 0选text, 其他选original_text. Defaults to 0.
Returns:
int: _description_
"""
if text_mode == 0:
text = self.text
else:
text = self.original_text
if text is None:
n_text = 0
else:
text = text.strip().split(" ")
n_text = len(text)
return n_text
@property
def text_num_per_second(self):
"""单位时间内的text数量"""
return self._cal_text_num_per_second(mode=0)
@property
def original_text_num_per_second(self):
"""单位时间内的original_text数量"""
return self._cal_text_num_per_second(mode=1)
@property
def tnps(self):
"""单位时间内的text数量"""
return self.text_num_per_second
@property
def original_tnps(self):
"""单位时间内的original_text数量"""
return self.original_text_num_per_second
def _cal_text_num_per_second(self, mode=0):
"""计算单位时间内的文本数量"""
text_num = self.text_num if mode == 0 else self.original_text_num
return text_num / self.duration
@classmethod
def from_data(cls, data: Dict):
return MusicClip(**data)
class MusicClipSeq(ClipSeq):
def __init__(self, items: List[Clip] = None):
super().__init__(items)
self.clipseq = self.data
@classmethod
def from_data(cls, clipseq: List[Dict]) -> MusicClipSeq:
new_clipseq = []
for clip in clipseq:
video_clip = MusicClip.from_data(clip)
new_clipseq.append(video_clip)
video_clipseq = MusicClipSeq(new_clipseq)
return video_clipseq
+140
View File
@@ -0,0 +1,140 @@
from __future__ import annotations
from typing import List, Dict
from moviepy.editor import concatenate_audioclips, AudioClip, AudioFileClip
from ...data import MediaMap, MediaMapEmb, MetaInfo, MediaMapSeq
from ...data.clip.clip_process import find_time_by_stage
from ...data.emb.h5py_emb import H5pyMediaMapEmb
from ...utils.util import load_dct_from_file
from .clip_process import get_stageseq_from_clipseq
from .music_clip import MusicClip, MusicClipSeq
from .meta_info import MusicMetaInfo
class MusicMap(MediaMap):
def __init__(
self,
meta_info: MetaInfo,
clipseq: MusicClipSeq,
lyricseq: MusicClipSeq = None,
stageseq: MusicClipSeq = None,
frameseq: MusicClipSeq = None,
emb: MediaMapEmb = None,
**kwargs,
):
self.lyricseq = lyricseq
super().__init__(meta_info, clipseq, stageseq, frameseq, emb, **kwargs)
if self.stageseq is None:
self.stageseq = MusicClipSeq.from_data(
get_stageseq_from_clipseq(self.clipseq)
)
self.stageseq.preprocess()
def preprocess(self):
if (
hasattr(self.meta_info, "target_stages")
and self.meta_info.target_stages is not None
):
self.set_start_end_by_target_stages()
super().preprocess()
self.spread_metainfo_2_clip(
target_keys=[
"media_path",
"media_map_path",
"emb_path",
"media_duration",
"mediaid",
"media_name",
"emb",
]
)
def set_start_end_by_target_stages(self):
target_stages = self.meta_info.target_stages
if not isinstance(target_stages, List):
target_stages = [target_stages]
start, _ = find_time_by_stage(self.stageseq, target_stages[0])
_, end = find_time_by_stage(self.stageseq, target_stages[-1])
self.meta_info.start = start
self.meta_info.end = end
@property
def audio_clip(self) -> AudioFileClip:
"""读取实际ClipSeq中的音频
Returns:
AudioClip: Moviepy中的audio_clip
"""
audio_clip = AudioFileClip(self.meta_info.media_path)
audio_clip = audio_clip.subclip(self.meta_info.start, self.meta_info.end)
return audio_clip
@classmethod
def from_json_path(
cls, path: Dict, emb_path: str, media_path: str = None, **kwargs
) -> MusicMap:
media_map = load_dct_from_file(path)
emb = H5pyMediaMapEmb(emb_path)
return cls.from_data(media_map, emb=emb, media_path=media_path, **kwargs)
@classmethod
def from_data(
cls, data: Dict, emb: H5pyMediaMapEmb, media_path: str = None, **kwargs
) -> MusicMap:
meta_info = MusicMetaInfo.from_data(data.get("meta_info", {}))
meta_info.media_path = media_path
clipseq = MusicClipSeq.from_data(data.get("clipseq", []))
stageseq = MusicClipSeq.from_data(data.get("stageseq", []))
lyricseq = MusicClipSeq.from_data(data.get("lyricseq", []))
target_keys = ["meta_info", "clipseq", "frameseq", "stageseq", "lyricseq"]
dct = {k: data[k] for k in data.keys() if k not in target_keys}
dct.update(**kwargs)
video_map = MusicMap(
meta_info=meta_info,
clipseq=clipseq,
stageseq=stageseq,
lyricseq=lyricseq,
emb=emb,
**dct,
)
return video_map
def to_dct(
self, target_keys: List[str] = None, ignored_keys: List[str] = None
) -> Dict:
dct = {}
dct["meta_info"] = self.meta_info.to_dct(
target_keys=target_keys, ignored_keys=ignored_keys
)
dct["clipseq"] = self.clipseq.to_dct(
target_keys=target_keys, ignored_keys=ignored_keys
)
if self.frameseq is not None:
dct["frameseq"] = self.frameseq.to_dct(
target_keys=target_keys, ignored_keys=ignored_keys
)
else:
dct["frameseq"] = None
if self.stageseq is not None:
dct["stageseq"] = self.stageseq.to_dct(
target_keys=target_keys, ignored_keys=ignored_keys
)
else:
dct["stageseq"] = None
dct["lyricseq"] = self.lyricseq.to_dct(
target_keys=target_keys, ignored_keys=ignored_keys
)
return dct
class MusicMapSeq(MediaMapSeq):
def __init__(self, maps: List[MusicMap]) -> None:
super().__init__(maps)
@property
def audio_clip(self) -> AudioFileClip:
audio_clip_lst = [m.audi_clip for m in self.maps]
audio_clip = concatenate_audioclips(audio_clip_lst)
return audio_clip
@@ -0,0 +1,58 @@
from moviepy.editor import (
ColorClip,
concatenate_videoclips,
AudioFileClip,
CompositeVideoClip,
)
from ...vision.video_map.video_lyric import render_lyric2video
from ...vision.video_map.video_writer import write_videoclip
from .music_map import MusicMap
def generate_music_map_videodemo(
music_map: MusicMap,
path: str,
audio_path: str,
render_lyric: bool = True,
width: int = 360,
height: int = 240,
fps: int = 25,
n_thread: int = 8,
colors: list = [[51, 161, 201], [46, 139, 87]],
) -> None:
"""输入音乐谱面,生成对应的转场视频Demo,视频内容只是简单的颜色切换
Args:
music_map (MusicInfo): 待可视化的音乐谱面
path (str): 可视化视频的存储路径
audio_path (str): 音乐谱面对应的音频路径
render_lyric (bool, optional): 是否渲染歌词,歌词在音乐谱面中. Defaults to True.
width (int, optional): 可视化视频的宽. Defaults to 360.
height (int, optional): 可视化视频的高. Defaults to 240.
fps (int, optional): 可视化视频的fps. Defaults to 25.
n_thread (int, optional): 可视化视频的写入线程数. Defaults to 8.
colors (list, optional): 可视化的视频颜色. Defaults to [[51, 161, 201], [46, 139, 87]].
"""
audio_clip = AudioFileClip(audio_path)
video_clips = []
size = (width, height)
for i, clip in enumerate(music_map.clipseq):
clip = ColorClip(
size=size, color=colors[i % len(colors)], duration=clip.duration
)
video_clips.append(clip)
video_clips = concatenate_videoclips(video_clips, method="compose")
if render_lyric:
video_clips = render_lyric2video(
videoclip=video_clips,
lyric=music_map,
lyric_info_type="music_map",
)
video_clips = video_clips.set_audio(audio_clip)
write_videoclip(
video_clips,
path=path,
fps=fps,
n_thread=n_thread,
)
View File
+9
View File
@@ -0,0 +1,9 @@
import os
from typing import Dict, Tuple
from ...utils.path_util import get_dir_file_map
def get_audio_path_dct(path, exts=["mp3", "flac", "wav"]) -> Dict[str, str]:
"""遍历目标文件夹及子文件夹下所有音频文件,生成字典。"""
return get_dir_file_map(path, exts=exts)
+158
View File
@@ -0,0 +1,158 @@
# Byte-compiled / optimized / DLL files
__pycache__/
*.py[cod]
*$py.class
# C extensions
*.so
# Distribution / packaging
.Python
build/
develop-eggs/
dist/
downloads/
eggs/
.eggs/
lib/
lib64/
parts/
sdist/
var/
wheels/
pip-wheel-metadata/
share/python-wheels/
*.egg-info/
.installed.cfg
*.egg
MANIFEST
# PyInstaller
# Usually these files are written by a python script from a template
# before PyInstaller builds the exe, so as to inject date/other infos into it.
*.manifest
*.spec
# Installer logs
pip-log.txt
pip-delete-this-directory.txt
# Unit test / coverage reports
htmlcov/
.tox/
.nox/
.coverage
.coverage.*
.cache
nosetests.xml
coverage.xml
*.cover
*.py,cover
.hypothesis/
.pytest_cache/
# Translations
*.mo
*.pot
# Django stuff:
*.log
local_settings.py
db.sqlite3
db.sqlite3-journal
# Flask stuff:
instance/
.webassets-cache
# Scrapy stuff:
.scrapy
# Sphinx documentation
docs/_build/
# PyBuilder
target/
# Jupyter Notebook
.ipynb_checkpoints
# IPython
profile_default/
ipython_config.py
# pyenv
.python-version
# pipenv
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
# However, in case of collaboration, if having platform-specific dependencies or dependencies
# having no cross-platform support, pipenv may install dependencies that don't work, or not
# install all needed dependencies.
#Pipfile.lock
# PEP 582; used by e.g. github.com/David-OConnor/pyflow
__pypackages__/
# Celery stuff
celerybeat-schedule
celerybeat.pid
# SageMath parsed files
*.sage.py
# Environments
.env
.venv
env/
venv/
ENV/
env.bak/
venv.bak/
# Spyder project settings
.spyderproject
.spyproject
# Rope project settings
.ropeproject
# mkdocs documentation
/site
# mypy
.mypy_cache/
.dmypy.json
dmypy.json
# Pyre type checker
.pyre/
.vscode
dataset/dataset_TM_train_cb1_temp.py
train_gpt_cnn_temp.py
train_gpt_cnn_mask.py
start.sh
start_eval.sh
config.json
output_GPT_Final
output_vqfinal
output_transformer
glove
checkpoints
dataset/HumanML3D
dataset/KIT-ML
output
matrix_multi.py
body_models
render_final_diffuse.py
render_final_mdm.py
pretrained
MDM
Motiondiffusion
Visualize_temp.py
new.sh
T2M_render
render_final_t2m.py
pose
+121
View File
@@ -0,0 +1,121 @@
import os
import torch
import numpy as np
from torch.utils.tensorboard import SummaryWriter
import json
import clip
import options.option_transformer as option_trans
import models.vqvae as vqvae
import utils.utils_model as utils_model
import utils.eval_trans as eval_trans
from dataset import dataset_TM_eval
import models.t2m_trans as trans
from options.get_eval_option import get_opt
from models.evaluator_wrapper import EvaluatorModelWrapper
import warnings
warnings.filterwarnings('ignore')
##### ---- Exp dirs ---- #####
args = option_trans.get_args_parser()
torch.manual_seed(args.seed)
args.out_dir = os.path.join(args.out_dir, f'{args.exp_name}')
os.makedirs(args.out_dir, exist_ok = True)
##### ---- Logger ---- #####
logger = utils_model.get_logger(args.out_dir)
writer = SummaryWriter(args.out_dir)
logger.info(json.dumps(vars(args), indent=4, sort_keys=True))
from utils.word_vectorizer import WordVectorizer
w_vectorizer = WordVectorizer('./glove', 'our_vab')
val_loader = dataset_TM_eval.DATALoader(args.dataname, True, 32, w_vectorizer)
dataset_opt_path = 'checkpoints/kit/Comp_v6_KLD005/opt.txt' if args.dataname == 'kit' else 'checkpoints/t2m/Comp_v6_KLD005/opt.txt'
wrapper_opt = get_opt(dataset_opt_path, torch.device('cuda'))
eval_wrapper = EvaluatorModelWrapper(wrapper_opt)
##### ---- Network ---- #####
## load clip model and datasets
clip_model, clip_preprocess = clip.load("ViT-B/32", device=torch.device('cuda'), 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
net = vqvae.HumanVQVAE(args, ## use args to define different parameters in different quantizers
args.nb_code,
args.code_dim,
args.output_emb_width,
args.down_t,
args.stride_t,
args.width,
args.depth,
args.dilation_growth_rate)
trans_encoder = trans.Text2Motion_Transformer(num_vq=args.nb_code,
embed_dim=args.embed_dim_gpt,
clip_dim=args.clip_dim,
block_size=args.block_size,
num_layers=args.num_layers,
n_head=args.n_head_gpt,
drop_out_rate=args.drop_out_rate,
fc_rate=args.ff_rate)
print ('loading checkpoint from {}'.format(args.resume_pth))
ckpt = torch.load(args.resume_pth, map_location='cpu')
net.load_state_dict(ckpt['net'], strict=True)
net.eval()
net.cuda()
if args.resume_trans is not None:
print ('loading transformer checkpoint from {}'.format(args.resume_trans))
ckpt = torch.load(args.resume_trans, map_location='cpu')
trans_encoder.load_state_dict(ckpt['trans'], strict=True)
trans_encoder.train()
trans_encoder.cuda()
fid = []
div = []
top1 = []
top2 = []
top3 = []
matching = []
multi = []
repeat_time = 20
for i in range(repeat_time):
best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, best_multi, writer, logger = eval_trans.evaluation_transformer_test(args.out_dir, val_loader, net, trans_encoder, logger, writer, 0, best_fid=1000, best_iter=0, best_div=100, best_top1=0, best_top2=0, best_top3=0, best_matching=100, best_multi=0, clip_model=clip_model, eval_wrapper=eval_wrapper, draw=False, savegif=False, save=False, savenpy=(i==0))
fid.append(best_fid)
div.append(best_div)
top1.append(best_top1)
top2.append(best_top2)
top3.append(best_top3)
matching.append(best_matching)
multi.append(best_multi)
print('final result:')
print('fid: ', sum(fid)/repeat_time)
print('div: ', sum(div)/repeat_time)
print('top1: ', sum(top1)/repeat_time)
print('top2: ', sum(top2)/repeat_time)
print('top3: ', sum(top3)/repeat_time)
print('matching: ', sum(matching)/repeat_time)
print('multi: ', sum(multi)/repeat_time)
fid = np.array(fid)
div = np.array(div)
top1 = np.array(top1)
top2 = np.array(top2)
top3 = np.array(top3)
matching = np.array(matching)
multi = np.array(multi)
msg_final = f"FID. {np.mean(fid):.3f}, conf. {np.std(fid)*1.96/np.sqrt(repeat_time):.3f}, Diversity. {np.mean(div):.3f}, conf. {np.std(div)*1.96/np.sqrt(repeat_time):.3f}, TOP1. {np.mean(top1):.3f}, conf. {np.std(top1)*1.96/np.sqrt(repeat_time):.3f}, TOP2. {np.mean(top2):.3f}, conf. {np.std(top2)*1.96/np.sqrt(repeat_time):.3f}, TOP3. {np.mean(top3):.3f}, conf. {np.std(top3)*1.96/np.sqrt(repeat_time):.3f}, Matching. {np.mean(matching):.3f}, conf. {np.std(matching)*1.96/np.sqrt(repeat_time):.3f}, Multi. {np.mean(multi):.3f}, conf. {np.std(multi)*1.96/np.sqrt(repeat_time):.3f}"
logger.info(msg_final)
+201
View File
@@ -0,0 +1,201 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright 2023 tencent
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
+329
View File
@@ -0,0 +1,329 @@
# (CVPR 2023) T2M-GPT
Pytorch implementation of paper "T2M-GPT: Generating Human Motion from Textual Descriptions with Discrete Representations"
[[Project Page]](https://mael-zys.github.io/T2M-GPT/) [[Paper]](https://arxiv.org/abs/2301.06052) [[Notebook Demo]](https://colab.research.google.com/drive/1Vy69w2q2d-Hg19F-KibqG0FRdpSj3L4O?usp=sharing) [[HuggingFace]](https://huggingface.co/vumichien/T2M-GPT) [[Space Demo]](https://huggingface.co/spaces/vumichien/generate_human_motion)
<p align="center">
<img src="img/Teaser.png" width="600px" alt="teaser">
</p>
If our project is helpful for your research, please consider citing :
```
@inproceedings{zhang2023generating,
title={T2M-GPT: Generating Human Motion from Textual Descriptions with Discrete Representations},
author={Zhang, Jianrong and Zhang, Yangsong and Cun, Xiaodong and Huang, Shaoli and Zhang, Yong and Zhao, Hongwei and Lu, Hongtao and Shen, Xi},
booktitle={Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition (CVPR)},
year={2023},
}
```
## Table of Content
* [1. Visual Results](#1-visual-results)
* [2. Installation](#2-installation)
* [3. Quick Start](#3-quick-start)
* [4. Train](#4-train)
* [5. Evaluation](#5-evaluation)
* [6. SMPL Mesh Rendering](#6-smpl-mesh-rendering)
* [7. Acknowledgement](#7-acknowledgement)
* [8. ChangLog](#8-changlog)
## 1. Visual Results (More results can be found in our [project page](https://mael-zys.github.io/T2M-GPT/))
<!-- ![visualization](img/ALLvis_new.png) -->
<p align="center">
<table>
<tr>
<th colspan="5">Text: a man steps forward and does a handstand.</th>
</tr>
<tr>
<th>GT</th>
<th><u><a href="https://ericguo5513.github.io/text-to-motion/"><nobr>T2M</nobr> </a></u></th>
<th><u><a href="https://guytevet.github.io/mdm-page/"><nobr>MDM</nobr> </a></u></th>
<th><u><a href="https://mingyuan-zhang.github.io/projects/MotionDiffuse.html"><nobr>MotionDiffuse</nobr> </a></u></th>
<th>Ours</th>
</tr>
<tr>
<td><img src="img/002103_gt_16.gif" width="140px" alt="gif"></td>
<td><img src="img/002103_pred_t2m_16.gif" width="140px" alt="gif"></td>
<td><img src="img/002103_pred_mdm_16.gif" width="140px" alt="gif"></td>
<td><img src="img/002103_pred_MotionDiffuse_16.gif" width="140px" alt="gif"></td>
<td><img src="img/002103_pred_16.gif" width="140px" alt="gif"></td>
</tr>
<tr>
<th colspan="5">Text: A man rises from the ground, walks in a circle and sits back down on the ground.</th>
</tr>
<tr>
<th>GT</th>
<th><u><a href="https://ericguo5513.github.io/text-to-motion/"><nobr>T2M</nobr> </a></u></th>
<th><u><a href="https://guytevet.github.io/mdm-page/"><nobr>MDM</nobr> </a></u></th>
<th><u><a href="https://mingyuan-zhang.github.io/projects/MotionDiffuse.html"><nobr>MotionDiffuse</nobr> </a></u></th>
<th>Ours</th>
</tr>
<tr>
<td><img src="img/000066_gt_16.gif" width="140px" alt="gif"></td>
<td><img src="img/000066_pred_t2m_16.gif" width="140px" alt="gif"></td>
<td><img src="img/000066_pred_mdm_16.gif" width="140px" alt="gif"></td>
<td><img src="img/000066_pred_MotionDiffuse_16.gif" width="140px" alt="gif"></td>
<td><img src="img/000066_pred_16.gif" width="140px" alt="gif"></td>
</tr>
</table>
</p>
## 2. Installation
### 2.1. Environment
Our model can be learnt in a **single GPU V100-32G**
```bash
conda env create -f environment.yml
conda activate T2M-GPT
```
The code was tested on Python 3.8 and PyTorch 1.8.1.
### 2.2. Dependencies
```bash
bash dataset/prepare/download_glove.sh
```
### 2.3. Datasets
We are using two 3D human motion-language dataset: HumanML3D and KIT-ML. For both datasets, you could find the details as well as download link [[here]](https://github.com/EricGuo5513/HumanML3D).
Take HumanML3D for an example, the file directory should look like this:
```
./dataset/HumanML3D/
├── new_joint_vecs/
├── texts/
├── Mean.npy # same as in [HumanML3D](https://github.com/EricGuo5513/HumanML3D)
├── Std.npy # same as in [HumanML3D](https://github.com/EricGuo5513/HumanML3D)
├── train.txt
├── val.txt
├── test.txt
├── train_val.txt
└── all.txt
```
### 2.4. Motion & text feature extractors:
We use the same extractors provided by [t2m](https://github.com/EricGuo5513/text-to-motion) to evaluate our generated motions. Please download the extractors.
```bash
bash dataset/prepare/download_extractor.sh
```
### 2.5. Pre-trained models
The pretrained model files will be stored in the 'pretrained' folder:
```bash
bash dataset/prepare/download_model.sh
```
### 2.6. Render SMPL mesh (optional)
If you want to render the generated motion, you need to install:
```bash
sudo sh dataset/prepare/download_smpl.sh
conda install -c menpo osmesa
conda install h5py
conda install -c conda-forge shapely pyrender trimesh mapbox_earcut
```
## 3. Quick Start
A quick start guide of how to use our code is available in [demo.ipynb](https://colab.research.google.com/drive/1Vy69w2q2d-Hg19F-KibqG0FRdpSj3L4O?usp=sharing)
<p align="center">
<img src="img/demo.png" width="400px" alt="demo">
</p>
## 4. Train
Note that, for kit dataset, just need to set '--dataname kit'.
### 4.1. VQ-VAE
The results are saved in the folder output.
<details>
<summary>
VQ training
</summary>
```bash
python3 train_vq.py \
--batch-size 256 \
--lr 2e-4 \
--total-iter 300000 \
--lr-scheduler 200000 \
--nb-code 512 \
--down-t 2 \
--depth 3 \
--dilation-growth-rate 3 \
--out-dir output \
--dataname t2m \
--vq-act relu \
--quantizer ema_reset \
--loss-vel 0.5 \
--recons-loss l1_smooth \
--exp-name VQVAE
```
</details>
### 4.2. GPT
The results are saved in the folder output.
<details>
<summary>
GPT training
</summary>
```bash
python3 train_t2m_trans.py \
--exp-name GPT \
--batch-size 128 \
--num-layers 9 \
--embed-dim-gpt 1024 \
--nb-code 512 \
--n-head-gpt 16 \
--block-size 51 \
--ff-rate 4 \
--drop-out-rate 0.1 \
--resume-pth output/VQVAE/net_last.pth \
--vq-name VQVAE \
--out-dir output \
--total-iter 300000 \
--lr-scheduler 150000 \
--lr 0.0001 \
--dataname t2m \
--down-t 2 \
--depth 3 \
--quantizer ema_reset \
--eval-iter 10000 \
--pkeep 0.5 \
--dilation-growth-rate 3 \
--vq-act relu
```
</details>
## 5. Evaluation
### 5.1. VQ-VAE
<details>
<summary>
VQ eval
</summary>
```bash
python3 VQ_eval.py \
--batch-size 256 \
--lr 2e-4 \
--total-iter 300000 \
--lr-scheduler 200000 \
--nb-code 512 \
--down-t 2 \
--depth 3 \
--dilation-growth-rate 3 \
--out-dir output \
--dataname t2m \
--vq-act relu \
--quantizer ema_reset \
--loss-vel 0.5 \
--recons-loss l1_smooth \
--exp-name TEST_VQVAE \
--resume-pth output/VQVAE/net_last.pth
```
</details>
### 5.2. GPT
<details>
<summary>
GPT eval
</summary>
Follow the evaluation setting of [text-to-motion](https://github.com/EricGuo5513/text-to-motion), we evaluate our model 20 times and report the average result. Due to the multimodality part where we should generate 30 motions from the same text, the evaluation takes a long time.
```bash
python3 GPT_eval_multi.py \
--exp-name TEST_GPT \
--batch-size 128 \
--num-layers 9 \
--embed-dim-gpt 1024 \
--nb-code 512 \
--n-head-gpt 16 \
--block-size 51 \
--ff-rate 4 \
--drop-out-rate 0.1 \
--resume-pth output/VQVAE/net_last.pth \
--vq-name VQVAE \
--out-dir output \
--total-iter 300000 \
--lr-scheduler 150000 \
--lr 0.0001 \
--dataname t2m \
--down-t 2 \
--depth 3 \
--quantizer ema_reset \
--eval-iter 10000 \
--pkeep 0.5 \
--dilation-growth-rate 3 \
--vq-act relu \
--resume-trans output/GPT/net_best_fid.pth
```
</details>
## 6. SMPL Mesh Rendering
<details>
<summary>
SMPL Mesh Rendering
</summary>
You should input the npy folder address and the motion names. Here is an example:
```bash
python3 render_final.py --filedir output/TEST_GPT/ --motion-list 000019 005485
```
</details>
### 7. Acknowledgement
We appreciate helps from :
* public code like [text-to-motion](https://github.com/EricGuo5513/text-to-motion), [TM2T](https://github.com/EricGuo5513/TM2T), [MDM](https://github.com/GuyTevet/motion-diffusion-model), [MotionDiffuse](https://github.com/mingyuan-zhang/MotionDiffuse) etc.
* <a href='https://mathis.petrovich.fr/'>Mathis Petrovich</a>, <a href='https://dulucas.github.io/'>Yuming Du</a>, <a href='https://github.com/yingyichen-cyy'>Yingyi Chen</a>, <a href='https://dexiong.me/'>Dexiong Chen</a> and <a href='https://xuelin-chen.github.io/'>Xuelin Chen</a> for inspiring discussions and valuable feedback.
* <a href='https://github.com/vumichien'>Minh Chien Vu</a> for the hugging face space demo.
### 8. ChangLog
* 2023/02/19 add the hugging face space demo for both skelton and SMPL mesh visualization.
+95
View File
@@ -0,0 +1,95 @@
import os
import json
import torch
from torch.utils.tensorboard import SummaryWriter
import numpy as np
import models.vqvae as vqvae
import options.option_vq as option_vq
import utils.utils_model as utils_model
from dataset import dataset_TM_eval
import utils.eval_trans as eval_trans
from options.get_eval_option import get_opt
from models.evaluator_wrapper import EvaluatorModelWrapper
import warnings
warnings.filterwarnings('ignore')
import numpy as np
##### ---- Exp dirs ---- #####
args = option_vq.get_args_parser()
torch.manual_seed(args.seed)
args.out_dir = os.path.join(args.out_dir, f'{args.exp_name}')
os.makedirs(args.out_dir, exist_ok = True)
##### ---- Logger ---- #####
logger = utils_model.get_logger(args.out_dir)
writer = SummaryWriter(args.out_dir)
logger.info(json.dumps(vars(args), indent=4, sort_keys=True))
from utils.word_vectorizer import WordVectorizer
w_vectorizer = WordVectorizer('./glove', 'our_vab')
dataset_opt_path = 'checkpoints/kit/Comp_v6_KLD005/opt.txt' if args.dataname == 'kit' else 'checkpoints/t2m/Comp_v6_KLD005/opt.txt'
wrapper_opt = get_opt(dataset_opt_path, torch.device('cuda'))
eval_wrapper = EvaluatorModelWrapper(wrapper_opt)
##### ---- Dataloader ---- #####
args.nb_joints = 21 if args.dataname == 'kit' else 22
val_loader = dataset_TM_eval.DATALoader(args.dataname, True, 32, w_vectorizer, unit_length=2**args.down_t)
##### ---- Network ---- #####
net = vqvae.HumanVQVAE(args, ## use args to define different parameters in different quantizers
args.nb_code,
args.code_dim,
args.output_emb_width,
args.down_t,
args.stride_t,
args.width,
args.depth,
args.dilation_growth_rate,
args.vq_act,
args.vq_norm)
if args.resume_pth :
logger.info('loading checkpoint from {}'.format(args.resume_pth))
ckpt = torch.load(args.resume_pth, map_location='cpu')
net.load_state_dict(ckpt['net'], strict=True)
net.train()
net.cuda()
fid = []
div = []
top1 = []
top2 = []
top3 = []
matching = []
repeat_time = 20
for i in range(repeat_time):
best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, writer, logger = eval_trans.evaluation_vqvae(args.out_dir, val_loader, net, logger, writer, 0, best_fid=1000, best_iter=0, best_div=100, best_top1=0, best_top2=0, best_top3=0, best_matching=100, eval_wrapper=eval_wrapper, draw=False, save=False, savenpy=(i==0))
fid.append(best_fid)
div.append(best_div)
top1.append(best_top1)
top2.append(best_top2)
top3.append(best_top3)
matching.append(best_matching)
print('final result:')
print('fid: ', sum(fid)/repeat_time)
print('div: ', sum(div)/repeat_time)
print('top1: ', sum(top1)/repeat_time)
print('top2: ', sum(top2)/repeat_time)
print('top3: ', sum(top3)/repeat_time)
print('matching: ', sum(matching)/repeat_time)
fid = np.array(fid)
div = np.array(div)
top1 = np.array(top1)
top2 = np.array(top2)
top3 = np.array(top3)
matching = np.array(matching)
msg_final = f"FID. {np.mean(fid):.3f}, conf. {np.std(fid)*1.96/np.sqrt(repeat_time):.3f}, Diversity. {np.mean(div):.3f}, conf. {np.std(div)*1.96/np.sqrt(repeat_time):.3f}, TOP1. {np.mean(top1):.3f}, conf. {np.std(top1)*1.96/np.sqrt(repeat_time):.3f}, TOP2. {np.mean(top2):.3f}, conf. {np.std(top2)*1.96/np.sqrt(repeat_time):.3f}, TOP3. {np.mean(top3):.3f}, conf. {np.std(top3)*1.96/np.sqrt(repeat_time):.3f}, Matching. {np.mean(matching):.3f}, conf. {np.std(matching)*1.96/np.sqrt(repeat_time):.3f}"
logger.info(msg_final)
View File
+217
View File
@@ -0,0 +1,217 @@
import torch
from torch.utils import data
import numpy as np
from os.path import join as pjoin
import random
import codecs as cs
from tqdm import tqdm
import utils.paramUtil as paramUtil
from torch.utils.data._utils.collate import default_collate
def collate_fn(batch):
batch.sort(key=lambda x: x[3], reverse=True)
return default_collate(batch)
'''For use of training text-2-motion generative model'''
class Text2MotionDataset(data.Dataset):
def __init__(self, dataset_name, is_test, w_vectorizer, feat_bias = 5, max_text_len = 20, unit_length = 4):
self.max_length = 20
self.pointer = 0
self.dataset_name = dataset_name
self.is_test = is_test
self.max_text_len = max_text_len
self.unit_length = unit_length
self.w_vectorizer = w_vectorizer
if dataset_name == 't2m':
self.data_root = './dataset/HumanML3D'
self.motion_dir = pjoin(self.data_root, 'new_joint_vecs')
self.text_dir = pjoin(self.data_root, 'texts')
self.joints_num = 22
radius = 4
fps = 20
self.max_motion_length = 196
dim_pose = 263
kinematic_chain = paramUtil.t2m_kinematic_chain
self.meta_dir = 'checkpoints/t2m/VQVAEV3_CB1024_CMT_H1024_NRES3/meta'
elif dataset_name == 'kit':
self.data_root = './dataset/KIT-ML'
self.motion_dir = pjoin(self.data_root, 'new_joint_vecs')
self.text_dir = pjoin(self.data_root, 'texts')
self.joints_num = 21
radius = 240 * 8
fps = 12.5
dim_pose = 251
self.max_motion_length = 196
kinematic_chain = paramUtil.kit_kinematic_chain
self.meta_dir = 'checkpoints/kit/VQVAEV3_CB1024_CMT_H1024_NRES3/meta'
mean = np.load(pjoin(self.meta_dir, 'mean.npy'))
std = np.load(pjoin(self.meta_dir, 'std.npy'))
if is_test:
split_file = pjoin(self.data_root, 'test.txt')
else:
split_file = pjoin(self.data_root, 'val.txt')
min_motion_len = 40 if self.dataset_name =='t2m' else 24
# min_motion_len = 64
joints_num = self.joints_num
data_dict = {}
id_list = []
with cs.open(split_file, 'r') as f:
for line in f.readlines():
id_list.append(line.strip())
new_name_list = []
length_list = []
for name in tqdm(id_list):
try:
motion = np.load(pjoin(self.motion_dir, name + '.npy'))
if (len(motion)) < min_motion_len or (len(motion) >= 200):
continue
text_data = []
flag = False
with cs.open(pjoin(self.text_dir, name + '.txt')) as f:
for line in f.readlines():
text_dict = {}
line_split = line.strip().split('#')
caption = line_split[0]
tokens = line_split[1].split(' ')
f_tag = float(line_split[2])
to_tag = float(line_split[3])
f_tag = 0.0 if np.isnan(f_tag) else f_tag
to_tag = 0.0 if np.isnan(to_tag) else to_tag
text_dict['caption'] = caption
text_dict['tokens'] = tokens
if f_tag == 0.0 and to_tag == 0.0:
flag = True
text_data.append(text_dict)
else:
try:
n_motion = motion[int(f_tag*fps) : int(to_tag*fps)]
if (len(n_motion)) < min_motion_len or (len(n_motion) >= 200):
continue
new_name = random.choice('ABCDEFGHIJKLMNOPQRSTUVW') + '_' + name
while new_name in data_dict:
new_name = random.choice('ABCDEFGHIJKLMNOPQRSTUVW') + '_' + name
data_dict[new_name] = {'motion': n_motion,
'length': len(n_motion),
'text':[text_dict]}
new_name_list.append(new_name)
length_list.append(len(n_motion))
except:
print(line_split)
print(line_split[2], line_split[3], f_tag, to_tag, name)
# break
if flag:
data_dict[name] = {'motion': motion,
'length': len(motion),
'text': text_data}
new_name_list.append(name)
length_list.append(len(motion))
except Exception as e:
# print(e)
pass
name_list, length_list = zip(*sorted(zip(new_name_list, length_list), key=lambda x: x[1]))
self.mean = mean
self.std = std
self.length_arr = np.array(length_list)
self.data_dict = data_dict
self.name_list = name_list
self.reset_max_len(self.max_length)
def reset_max_len(self, length):
assert length <= self.max_motion_length
self.pointer = np.searchsorted(self.length_arr, length)
print("Pointer Pointing at %d"%self.pointer)
self.max_length = length
def inv_transform(self, data):
return data * self.std + self.mean
def forward_transform(self, data):
return (data - self.mean) / self.std
def __len__(self):
return len(self.data_dict) - self.pointer
def __getitem__(self, item):
idx = self.pointer + item
name = self.name_list[idx]
data = self.data_dict[name]
# data = self.data_dict[self.name_list[idx]]
motion, m_length, text_list = data['motion'], data['length'], data['text']
# Randomly select a caption
text_data = random.choice(text_list)
caption, tokens = text_data['caption'], text_data['tokens']
if len(tokens) < self.max_text_len:
# pad with "unk"
tokens = ['sos/OTHER'] + tokens + ['eos/OTHER']
sent_len = len(tokens)
tokens = tokens + ['unk/OTHER'] * (self.max_text_len + 2 - sent_len)
else:
# crop
tokens = tokens[:self.max_text_len]
tokens = ['sos/OTHER'] + tokens + ['eos/OTHER']
sent_len = len(tokens)
pos_one_hots = []
word_embeddings = []
for token in tokens:
word_emb, pos_oh = self.w_vectorizer[token]
pos_one_hots.append(pos_oh[None, :])
word_embeddings.append(word_emb[None, :])
pos_one_hots = np.concatenate(pos_one_hots, axis=0)
word_embeddings = np.concatenate(word_embeddings, axis=0)
if self.unit_length < 10:
coin2 = np.random.choice(['single', 'single', 'double'])
else:
coin2 = 'single'
if coin2 == 'double':
m_length = (m_length // self.unit_length - 1) * self.unit_length
elif coin2 == 'single':
m_length = (m_length // self.unit_length) * self.unit_length
idx = random.randint(0, len(motion) - m_length)
motion = motion[idx:idx+m_length]
"Z Normalization"
motion = (motion - self.mean) / self.std
if m_length < self.max_motion_length:
motion = np.concatenate([motion,
np.zeros((self.max_motion_length - m_length, motion.shape[1]))
], axis=0)
return word_embeddings, pos_one_hots, caption, sent_len, motion, m_length, '_'.join(tokens), name
def DATALoader(dataset_name, is_test,
batch_size, w_vectorizer,
num_workers = 8, unit_length = 4) :
val_loader = torch.utils.data.DataLoader(Text2MotionDataset(dataset_name, is_test, w_vectorizer, unit_length=unit_length),
batch_size,
shuffle = True,
num_workers=num_workers,
collate_fn=collate_fn,
drop_last = True)
return val_loader
def cycle(iterable):
while True:
for x in iterable:
yield x
+161
View File
@@ -0,0 +1,161 @@
import torch
from torch.utils import data
import numpy as np
from os.path import join as pjoin
import random
import codecs as cs
from tqdm import tqdm
import utils.paramUtil as paramUtil
from torch.utils.data._utils.collate import default_collate
def collate_fn(batch):
batch.sort(key=lambda x: x[3], reverse=True)
return default_collate(batch)
'''For use of training text-2-motion generative model'''
class Text2MotionDataset(data.Dataset):
def __init__(self, dataset_name, feat_bias = 5, unit_length = 4, codebook_size = 1024, tokenizer_name=None):
self.max_length = 64
self.pointer = 0
self.dataset_name = dataset_name
self.unit_length = unit_length
# self.mot_start_idx = codebook_size
self.mot_end_idx = codebook_size
self.mot_pad_idx = codebook_size + 1
if dataset_name == 't2m':
self.data_root = './dataset/HumanML3D'
self.motion_dir = pjoin(self.data_root, 'new_joint_vecs')
self.text_dir = pjoin(self.data_root, 'texts')
self.joints_num = 22
radius = 4
fps = 20
self.max_motion_length = 26 if unit_length == 8 else 51
dim_pose = 263
kinematic_chain = paramUtil.t2m_kinematic_chain
elif dataset_name == 'kit':
self.data_root = './dataset/KIT-ML'
self.motion_dir = pjoin(self.data_root, 'new_joint_vecs')
self.text_dir = pjoin(self.data_root, 'texts')
self.joints_num = 21
radius = 240 * 8
fps = 12.5
dim_pose = 251
self.max_motion_length = 26 if unit_length == 8 else 51
kinematic_chain = paramUtil.kit_kinematic_chain
split_file = pjoin(self.data_root, 'train.txt')
id_list = []
with cs.open(split_file, 'r') as f:
for line in f.readlines():
id_list.append(line.strip())
new_name_list = []
data_dict = {}
for name in tqdm(id_list):
try:
m_token_list = np.load(pjoin(self.data_root, tokenizer_name, '%s.npy'%name))
# Read text
with cs.open(pjoin(self.text_dir, name + '.txt')) as f:
text_data = []
flag = False
lines = f.readlines()
for line in lines:
try:
text_dict = {}
line_split = line.strip().split('#')
caption = line_split[0]
t_tokens = line_split[1].split(' ')
f_tag = float(line_split[2])
to_tag = float(line_split[3])
f_tag = 0.0 if np.isnan(f_tag) else f_tag
to_tag = 0.0 if np.isnan(to_tag) else to_tag
text_dict['caption'] = caption
text_dict['tokens'] = t_tokens
if f_tag == 0.0 and to_tag == 0.0:
flag = True
text_data.append(text_dict)
else:
m_token_list_new = [tokens[int(f_tag*fps/unit_length) : int(to_tag*fps/unit_length)] for tokens in m_token_list if int(f_tag*fps/unit_length) < int(to_tag*fps/unit_length)]
if len(m_token_list_new) == 0:
continue
new_name = '%s_%f_%f'%(name, f_tag, to_tag)
data_dict[new_name] = {'m_token_list': m_token_list_new,
'text':[text_dict]}
new_name_list.append(new_name)
except:
pass
if flag:
data_dict[name] = {'m_token_list': m_token_list,
'text':text_data}
new_name_list.append(name)
except:
pass
self.data_dict = data_dict
self.name_list = new_name_list
def __len__(self):
return len(self.data_dict)
def __getitem__(self, item):
data = self.data_dict[self.name_list[item]]
m_token_list, text_list = data['m_token_list'], data['text']
m_tokens = random.choice(m_token_list)
text_data = random.choice(text_list)
caption= text_data['caption']
coin = np.random.choice([False, False, True])
# print(len(m_tokens))
if coin:
# drop one token at the head or tail
coin2 = np.random.choice([True, False])
if coin2:
m_tokens = m_tokens[:-1]
else:
m_tokens = m_tokens[1:]
m_tokens_len = m_tokens.shape[0]
if m_tokens_len+1 < self.max_motion_length:
m_tokens = np.concatenate([m_tokens, np.ones((1), dtype=int) * self.mot_end_idx, np.ones((self.max_motion_length-1-m_tokens_len), dtype=int) * self.mot_pad_idx], axis=0)
else:
m_tokens = np.concatenate([m_tokens, np.ones((1), dtype=int) * self.mot_end_idx], axis=0)
return caption, m_tokens.reshape(-1), m_tokens_len
def DATALoader(dataset_name,
batch_size, codebook_size, tokenizer_name, unit_length=4,
num_workers = 8) :
train_loader = torch.utils.data.DataLoader(Text2MotionDataset(dataset_name, codebook_size = codebook_size, tokenizer_name = tokenizer_name, unit_length=unit_length),
batch_size,
shuffle=True,
num_workers=num_workers,
#collate_fn=collate_fn,
drop_last = True)
return train_loader
def cycle(iterable):
while True:
for x in iterable:
yield x
+109
View File
@@ -0,0 +1,109 @@
import torch
from torch.utils import data
import numpy as np
from os.path import join as pjoin
import random
import codecs as cs
from tqdm import tqdm
class VQMotionDataset(data.Dataset):
def __init__(self, dataset_name, window_size = 64, unit_length = 4):
self.window_size = window_size
self.unit_length = unit_length
self.dataset_name = dataset_name
if dataset_name == 't2m':
self.data_root = './dataset/HumanML3D'
self.motion_dir = pjoin(self.data_root, 'new_joint_vecs')
self.text_dir = pjoin(self.data_root, 'texts')
self.joints_num = 22
self.max_motion_length = 196
self.meta_dir = 'checkpoints/t2m/VQVAEV3_CB1024_CMT_H1024_NRES3/meta'
elif dataset_name == 'kit':
self.data_root = './dataset/KIT-ML'
self.motion_dir = pjoin(self.data_root, 'new_joint_vecs')
self.text_dir = pjoin(self.data_root, 'texts')
self.joints_num = 21
self.max_motion_length = 196
self.meta_dir = 'checkpoints/kit/VQVAEV3_CB1024_CMT_H1024_NRES3/meta'
joints_num = self.joints_num
mean = np.load(pjoin(self.meta_dir, 'mean.npy'))
std = np.load(pjoin(self.meta_dir, 'std.npy'))
split_file = pjoin(self.data_root, 'train.txt')
self.data = []
self.lengths = []
id_list = []
with cs.open(split_file, 'r') as f:
for line in f.readlines():
id_list.append(line.strip())
for name in tqdm(id_list):
try:
motion = np.load(pjoin(self.motion_dir, name + '.npy'))
if motion.shape[0] < self.window_size:
continue
self.lengths.append(motion.shape[0] - self.window_size)
self.data.append(motion)
except:
# Some motion may not exist in KIT dataset
pass
self.mean = mean
self.std = std
print("Total number of motions {}".format(len(self.data)))
def inv_transform(self, data):
return data * self.std + self.mean
def compute_sampling_prob(self) :
prob = np.array(self.lengths, dtype=np.float32)
prob /= np.sum(prob)
return prob
def __len__(self):
return len(self.data)
def __getitem__(self, item):
motion = self.data[item]
idx = random.randint(0, len(motion) - self.window_size)
motion = motion[idx:idx+self.window_size]
"Z Normalization"
motion = (motion - self.mean) / self.std
return motion
def DATALoader(dataset_name,
batch_size,
num_workers = 8,
window_size = 64,
unit_length = 4):
trainSet = VQMotionDataset(dataset_name, window_size=window_size, unit_length=unit_length)
prob = trainSet.compute_sampling_prob()
sampler = torch.utils.data.WeightedRandomSampler(prob, num_samples = len(trainSet) * 1000, replacement=True)
train_loader = torch.utils.data.DataLoader(trainSet,
batch_size,
shuffle=True,
#sampler=sampler,
num_workers=num_workers,
#collate_fn=collate_fn,
drop_last = True)
return train_loader
def cycle(iterable):
while True:
for x in iterable:
yield x
+117
View File
@@ -0,0 +1,117 @@
import torch
from torch.utils import data
import numpy as np
from os.path import join as pjoin
import random
import codecs as cs
from tqdm import tqdm
class VQMotionDataset(data.Dataset):
def __init__(self, dataset_name, feat_bias = 5, window_size = 64, unit_length = 8):
self.window_size = window_size
self.unit_length = unit_length
self.feat_bias = feat_bias
self.dataset_name = dataset_name
min_motion_len = 40 if dataset_name =='t2m' else 24
if dataset_name == 't2m':
self.data_root = './dataset/HumanML3D'
self.motion_dir = pjoin(self.data_root, 'new_joint_vecs')
self.text_dir = pjoin(self.data_root, 'texts')
self.joints_num = 22
radius = 4
fps = 20
self.max_motion_length = 196
dim_pose = 263
self.meta_dir = 'checkpoints/t2m/VQVAEV3_CB1024_CMT_H1024_NRES3/meta'
#kinematic_chain = paramUtil.t2m_kinematic_chain
elif dataset_name == 'kit':
self.data_root = './dataset/KIT-ML'
self.motion_dir = pjoin(self.data_root, 'new_joint_vecs')
self.text_dir = pjoin(self.data_root, 'texts')
self.joints_num = 21
radius = 240 * 8
fps = 12.5
dim_pose = 251
self.max_motion_length = 196
self.meta_dir = 'checkpoints/kit/VQVAEV3_CB1024_CMT_H1024_NRES3/meta'
#kinematic_chain = paramUtil.kit_kinematic_chain
joints_num = self.joints_num
mean = np.load(pjoin(self.meta_dir, 'mean.npy'))
std = np.load(pjoin(self.meta_dir, 'std.npy'))
split_file = pjoin(self.data_root, 'train.txt')
data_dict = {}
id_list = []
with cs.open(split_file, 'r') as f:
for line in f.readlines():
id_list.append(line.strip())
new_name_list = []
length_list = []
for name in tqdm(id_list):
try:
motion = np.load(pjoin(self.motion_dir, name + '.npy'))
if (len(motion)) < min_motion_len or (len(motion) >= 200):
continue
data_dict[name] = {'motion': motion,
'length': len(motion),
'name': name}
new_name_list.append(name)
length_list.append(len(motion))
except:
# Some motion may not exist in KIT dataset
pass
self.mean = mean
self.std = std
self.length_arr = np.array(length_list)
self.data_dict = data_dict
self.name_list = new_name_list
def inv_transform(self, data):
return data * self.std + self.mean
def __len__(self):
return len(self.data_dict)
def __getitem__(self, item):
name = self.name_list[item]
data = self.data_dict[name]
motion, m_length = data['motion'], data['length']
m_length = (m_length // self.unit_length) * self.unit_length
idx = random.randint(0, len(motion) - m_length)
motion = motion[idx:idx+m_length]
"Z Normalization"
motion = (motion - self.mean) / self.std
return motion, name
def DATALoader(dataset_name,
batch_size = 1,
num_workers = 8, unit_length = 4) :
train_loader = torch.utils.data.DataLoader(VQMotionDataset(dataset_name, unit_length=unit_length),
batch_size,
shuffle=True,
num_workers=num_workers,
#collate_fn=collate_fn,
drop_last = True)
return train_loader
def cycle(iterable):
while True:
for x in iterable:
yield x
@@ -0,0 +1,15 @@
rm -rf checkpoints
mkdir checkpoints
cd checkpoints
echo -e "Downloading extractors"
gdown --fuzzy https://drive.google.com/file/d/1o7RTDQcToJjTm9_mNWTyzvZvjTWpZfug/view
gdown --fuzzy https://drive.google.com/file/d/1KNU8CsMAnxFrwopKBBkC8jEULGLPBHQp/view
unzip t2m.zip
unzip kit.zip
echo -e "Cleaning\n"
rm t2m.zip
rm kit.zip
echo -e "Downloading done!"
@@ -0,0 +1,9 @@
echo -e "Downloading glove (in use by the evaluators)"
gdown --fuzzy https://drive.google.com/file/d/1bCeS6Sh_mLVTebxIgiUHgdPrroW06mb6/view?usp=sharing
rm -rf glove
unzip glove.zip
echo -e "Cleaning\n"
rm glove.zip
echo -e "Downloading done!"
@@ -0,0 +1,12 @@
mkdir -p pretrained
cd pretrained/
echo -e "The pretrained model files will be stored in the 'pretrained' folder\n"
gdown 1LaOvwypF-jM2Axnq5dc-Iuvv3w_G-WDE
unzip VQTrans_pretrained.zip
echo -e "Cleaning\n"
rm VQTrans_pretrained.zip
echo -e "Downloading done!"
@@ -0,0 +1,13 @@
mkdir -p body_models
cd body_models/
echo -e "The smpl files will be stored in the 'body_models/smpl/' folder\n"
gdown 1INYlGA76ak_cKGzvpOV2Pe6RkYTlXTW2
rm -rf smpl
unzip smpl.zip
echo -e "Cleaning\n"
rm smpl.zip
echo -e "Downloading done!"
+121
View File
@@ -0,0 +1,121 @@
name: T2M-GPT
channels:
- pytorch
- defaults
dependencies:
- _libgcc_mutex=0.1=main
- _openmp_mutex=4.5=1_gnu
- blas=1.0=mkl
- bzip2=1.0.8=h7b6447c_0
- ca-certificates=2021.7.5=h06a4308_1
- certifi=2021.5.30=py38h06a4308_0
- cudatoolkit=10.1.243=h6bb024c_0
- ffmpeg=4.3=hf484d3e_0
- freetype=2.10.4=h5ab3b9f_0
- gmp=6.2.1=h2531618_2
- gnutls=3.6.15=he1e5248_0
- intel-openmp=2021.3.0=h06a4308_3350
- jpeg=9b=h024ee3a_2
- lame=3.100=h7b6447c_0
- lcms2=2.12=h3be6417_0
- ld_impl_linux-64=2.35.1=h7274673_9
- libffi=3.3=he6710b0_2
- libgcc-ng=9.3.0=h5101ec6_17
- libgomp=9.3.0=h5101ec6_17
- libiconv=1.15=h63c8f33_5
- libidn2=2.3.2=h7f8727e_0
- libpng=1.6.37=hbc83047_0
- libstdcxx-ng=9.3.0=hd4cf53a_17
- libtasn1=4.16.0=h27cfd23_0
- libtiff=4.2.0=h85742a9_0
- libunistring=0.9.10=h27cfd23_0
- libuv=1.40.0=h7b6447c_fxfi0
- libwebp-base=1.2.0=h27cfd23_0
- lz4-c=1.9.3=h295c915_1
- mkl=2021.3.0=h06a4308_520
- mkl-service=2.4.0=py38h7f8727e_0
- mkl_fft=1.3.0=py38h42c9631_2
- mkl_random=1.2.2=py38h51133e4_0
- ncurses=6.2=he6710b0_1
- nettle=3.7.3=hbbd107a_1
- ninja=1.10.2=hff7bd54_1
- numpy=1.20.3=py38hf144106_0
- numpy-base=1.20.3=py38h74d4b33_0
- olefile=0.46=py_0
- openh264=2.1.0=hd408876_0
- openjpeg=2.3.0=h05c96fa_1
- openssl=1.1.1k=h27cfd23_0
- pillow=8.3.1=py38h2c7a002_0
- pip=21.0.1=py38h06a4308_0
- python=3.8.11=h12debd9_0_cpython
- pytorch=1.8.1=py3.8_cuda10.1_cudnn7.6.3_0
- readline=8.1=h27cfd23_0
- setuptools=52.0.0=py38h06a4308_0
- six=1.16.0=pyhd3eb1b0_0
- sqlite=3.36.0=hc218d9a_0
- tk=8.6.10=hbc83047_0
- torchaudio=0.8.1=py38
- torchvision=0.9.1=py38_cu101
- typing_extensions=3.10.0.0=pyh06a4308_0
- wheel=0.37.0=pyhd3eb1b0_0
- xz=5.2.5=h7b6447c_0
- zlib=1.2.11=h7b6447c_3
- zstd=1.4.9=haebb681_0
- pip:
- absl-py==0.13.0
- backcall==0.2.0
- cachetools==4.2.2
- charset-normalizer==2.0.4
- chumpy==0.70
- cycler==0.10.0
- decorator==5.0.9
- google-auth==1.35.0
- google-auth-oauthlib==0.4.5
- grpcio==1.39.0
- idna==3.2
- imageio==2.9.0
- ipdb==0.13.9
- ipython==7.26.0
- ipython-genutils==0.2.0
- jedi==0.18.0
- joblib==1.0.1
- kiwisolver==1.3.1
- markdown==3.3.4
- matplotlib==3.4.3
- matplotlib-inline==0.1.2
- oauthlib==3.1.1
- pandas==1.3.2
- parso==0.8.2
- pexpect==4.8.0
- pickleshare==0.7.5
- prompt-toolkit==3.0.20
- protobuf==3.17.3
- ptyprocess==0.7.0
- pyasn1==0.4.8
- pyasn1-modules==0.2.8
- pygments==2.10.0
- pyparsing==2.4.7
- python-dateutil==2.8.2
- pytz==2021.1
- pyyaml==5.4.1
- requests==2.26.0
- requests-oauthlib==1.3.0
- rsa==4.7.2
- scikit-learn==0.24.2
- scipy==1.7.1
- sklearn==0.0
- smplx==0.1.28
- tensorboard==2.6.0
- tensorboard-data-server==0.6.1
- tensorboard-plugin-wit==1.8.0
- threadpoolctl==2.2.0
- toml==0.10.2
- tqdm==4.62.2
- traitlets==5.0.5
- urllib3==1.26.6
- wcwidth==0.2.5
- werkzeug==2.0.1
- git+https://mirrors.tencent.com/github.com/openai/CLIP.git
- git+https://mirrors.tencent.com/github.com/nghorbani/human_body_prior
- gdown
- moviepy
Binary file not shown.

After

Width:  |  Height:  |  Size: 650 KiB

+67
View File
@@ -0,0 +1,67 @@
import torch.nn as nn
from .resnet import Resnet1D
class Encoder(nn.Module):
def __init__(self,
input_emb_width = 3,
output_emb_width = 512,
down_t = 3,
stride_t = 2,
width = 512,
depth = 3,
dilation_growth_rate = 3,
activation='relu',
norm=None):
super().__init__()
blocks = []
filter_t, pad_t = stride_t * 2, stride_t // 2
blocks.append(nn.Conv1d(input_emb_width, width, 3, 1, 1))
blocks.append(nn.ReLU())
for i in range(down_t):
input_dim = width
block = nn.Sequential(
nn.Conv1d(input_dim, width, filter_t, stride_t, pad_t),
Resnet1D(width, depth, dilation_growth_rate, activation=activation, norm=norm),
)
blocks.append(block)
blocks.append(nn.Conv1d(width, output_emb_width, 3, 1, 1))
self.model = nn.Sequential(*blocks)
def forward(self, x):
return self.model(x)
class Decoder(nn.Module):
def __init__(self,
input_emb_width = 3,
output_emb_width = 512,
down_t = 3,
stride_t = 2,
width = 512,
depth = 3,
dilation_growth_rate = 3,
activation='relu',
norm=None):
super().__init__()
blocks = []
filter_t, pad_t = stride_t * 2, stride_t // 2
blocks.append(nn.Conv1d(output_emb_width, width, 3, 1, 1))
blocks.append(nn.ReLU())
for i in range(down_t):
out_dim = width
block = nn.Sequential(
Resnet1D(width, depth, dilation_growth_rate, reverse_dilation=True, activation=activation, norm=norm),
nn.Upsample(scale_factor=2, mode='nearest'),
nn.Conv1d(width, out_dim, 3, 1, 1)
)
blocks.append(block)
blocks.append(nn.Conv1d(width, width, 3, 1, 1))
blocks.append(nn.ReLU())
blocks.append(nn.Conv1d(width, input_emb_width, 3, 1, 1))
self.model = nn.Sequential(*blocks)
def forward(self, x):
return self.model(x)
+92
View File
@@ -0,0 +1,92 @@
import torch
from os.path import join as pjoin
import numpy as np
from .modules import MovementConvEncoder, TextEncoderBiGRUCo, MotionEncoderBiGRUCo
from ..utils.word_vectorizer import POS_enumerator
def build_models(opt):
movement_enc = MovementConvEncoder(opt.dim_pose-4, opt.dim_movement_enc_hidden, opt.dim_movement_latent)
text_enc = TextEncoderBiGRUCo(word_size=opt.dim_word,
pos_size=opt.dim_pos_ohot,
hidden_size=opt.dim_text_hidden,
output_size=opt.dim_coemb_hidden,
device=opt.device)
motion_enc = MotionEncoderBiGRUCo(input_size=opt.dim_movement_latent,
hidden_size=opt.dim_motion_hidden,
output_size=opt.dim_coemb_hidden,
device=opt.device)
checkpoint = torch.load(pjoin(opt.checkpoints_dir, opt.dataset_name, 'text_mot_match', 'model', 'finest.tar'),
map_location=opt.device)
movement_enc.load_state_dict(checkpoint['movement_encoder'])
text_enc.load_state_dict(checkpoint['text_encoder'])
motion_enc.load_state_dict(checkpoint['motion_encoder'])
print('Loading Evaluation Model Wrapper (Epoch %d) Completed!!' % (checkpoint['epoch']))
return text_enc, motion_enc, movement_enc
class EvaluatorModelWrapper(object):
def __init__(self, opt):
if opt.dataset_name == 't2m':
opt.dim_pose = 263
elif opt.dataset_name == 'kit':
opt.dim_pose = 251
else:
raise KeyError('Dataset not Recognized!!!')
opt.dim_word = 300
opt.max_motion_length = 196
opt.dim_pos_ohot = len(POS_enumerator)
opt.dim_motion_hidden = 1024
opt.max_text_len = 20
opt.dim_text_hidden = 512
opt.dim_coemb_hidden = 512
# print(opt)
self.text_encoder, self.motion_encoder, self.movement_encoder = build_models(opt)
self.opt = opt
self.device = opt.device
self.text_encoder.to(opt.device)
self.motion_encoder.to(opt.device)
self.movement_encoder.to(opt.device)
self.text_encoder.eval()
self.motion_encoder.eval()
self.movement_encoder.eval()
# Please note that the results does not following the order of inputs
def get_co_embeddings(self, word_embs, pos_ohot, cap_lens, motions, m_lens):
with torch.no_grad():
word_embs = word_embs.detach().to(self.device).float()
pos_ohot = pos_ohot.detach().to(self.device).float()
motions = motions.detach().to(self.device).float()
'''Movement Encoding'''
movements = self.movement_encoder(motions[..., :-4]).detach()
m_lens = m_lens // self.opt.unit_length
motion_embedding = self.motion_encoder(movements, m_lens)
'''Text Encoding'''
text_embedding = self.text_encoder(word_embs, pos_ohot, cap_lens)
return text_embedding, motion_embedding
# Please note that the results does not following the order of inputs
def get_motion_embeddings(self, motions, m_lens):
with torch.no_grad():
motions = motions.detach().to(self.device).float()
align_idx = np.argsort(m_lens.data.tolist())[::-1].copy()
motions = motions[align_idx]
m_lens = m_lens[align_idx]
'''Movement Encoding'''
movements = self.movement_encoder(motions[..., :-4]).detach()
m_lens = m_lens // self.opt.unit_length
motion_embedding = self.motion_encoder(movements, m_lens)
return motion_embedding
+109
View File
@@ -0,0 +1,109 @@
import torch
import torch.nn as nn
from torch.nn.utils.rnn import pack_padded_sequence
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)
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 TextEncoderBiGRUCo(nn.Module):
def __init__(self, word_size, pos_size, hidden_size, output_size, device):
super(TextEncoderBiGRUCo, self).__init__()
self.device = device
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.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)
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 MotionEncoderBiGRUCo(nn.Module):
def __init__(self, input_size, hidden_size, output_size, device):
super(MotionEncoderBiGRUCo, self).__init__()
self.device = device
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, 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)
+43
View File
@@ -0,0 +1,43 @@
"""
Various positional encodings for the transformer.
"""
import math
import torch
from torch import nn
def PE1d_sincos(seq_length, dim):
"""
:param d_model: dimension of the model
:param length: length of positions
:return: length*d_model position matrix
"""
if dim % 2 != 0:
raise ValueError("Cannot use sin/cos positional encoding with "
"odd dim (got dim={:d})".format(dim))
pe = torch.zeros(seq_length, dim)
position = torch.arange(0, seq_length).unsqueeze(1)
div_term = torch.exp((torch.arange(0, dim, 2, dtype=torch.float) *
-(math.log(10000.0) / dim)))
pe[:, 0::2] = torch.sin(position.float() * div_term)
pe[:, 1::2] = torch.cos(position.float() * div_term)
return pe.unsqueeze(1)
class PositionEmbedding(nn.Module):
"""
Absolute pos embedding (standard), learned.
"""
def __init__(self, seq_length, dim, dropout, grad=False):
super().__init__()
self.embed = nn.Parameter(data=PE1d_sincos(seq_length, dim), requires_grad=grad)
self.dropout = nn.Dropout(p=dropout)
def forward(self, x):
# x.shape: bs, seq_len, feat_dim
l = x.shape[1]
x = x.permute(1, 0, 2) + self.embed[:l].expand(x.permute(1, 0, 2).shape)
x = self.dropout(x.permute(1, 0, 2))
return x
+413
View File
@@ -0,0 +1,413 @@
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
class QuantizeEMAReset(nn.Module):
def __init__(self, nb_code, code_dim, args):
super().__init__()
self.nb_code = nb_code
self.code_dim = code_dim
self.mu = args.mu
self.reset_codebook()
def reset_codebook(self):
self.init = False
self.code_sum = None
self.code_count = None
self.register_buffer('codebook', torch.zeros(self.nb_code, self.code_dim).cuda())
def _tile(self, x):
nb_code_x, code_dim = x.shape
if nb_code_x < self.nb_code:
n_repeats = (self.nb_code + nb_code_x - 1) // nb_code_x
std = 0.01 / np.sqrt(code_dim)
out = x.repeat(n_repeats, 1)
out = out + torch.randn_like(out) * std
else :
out = x
return out
def init_codebook(self, x):
out = self._tile(x)
self.codebook = out[:self.nb_code]
self.code_sum = self.codebook.clone()
self.code_count = torch.ones(self.nb_code, device=self.codebook.device)
self.init = True
@torch.no_grad()
def compute_perplexity(self, code_idx) :
# Calculate new centres
code_onehot = torch.zeros(self.nb_code, code_idx.shape[0], device=code_idx.device) # nb_code, N * L
code_onehot.scatter_(0, code_idx.view(1, code_idx.shape[0]), 1)
code_count = code_onehot.sum(dim=-1) # nb_code
prob = code_count / torch.sum(code_count)
perplexity = torch.exp(-torch.sum(prob * torch.log(prob + 1e-7)))
return perplexity
@torch.no_grad()
def update_codebook(self, x, code_idx):
code_onehot = torch.zeros(self.nb_code, x.shape[0], device=x.device) # nb_code, N * L
code_onehot.scatter_(0, code_idx.view(1, x.shape[0]), 1)
code_sum = torch.matmul(code_onehot, x) # nb_code, w
code_count = code_onehot.sum(dim=-1) # nb_code
out = self._tile(x)
code_rand = out[:self.nb_code]
# Update centres
self.code_sum = self.mu * self.code_sum + (1. - self.mu) * code_sum # w, nb_code
self.code_count = self.mu * self.code_count + (1. - self.mu) * code_count # nb_code
usage = (self.code_count.view(self.nb_code, 1) >= 1.0).float()
code_update = self.code_sum.view(self.nb_code, self.code_dim) / self.code_count.view(self.nb_code, 1)
self.codebook = usage * code_update + (1 - usage) * code_rand
prob = code_count / torch.sum(code_count)
perplexity = torch.exp(-torch.sum(prob * torch.log(prob + 1e-7)))
return perplexity
def preprocess(self, x):
# NCT -> NTC -> [NT, C]
x = x.permute(0, 2, 1).contiguous()
x = x.view(-1, x.shape[-1])
return x
def quantize(self, x):
# Calculate latent code x_l
k_w = self.codebook.t()
distance = torch.sum(x ** 2, dim=-1, keepdim=True) - 2 * torch.matmul(x, k_w) + torch.sum(k_w ** 2, dim=0,
keepdim=True) # (N * L, b)
_, code_idx = torch.min(distance, dim=-1)
return code_idx
def dequantize(self, code_idx):
x = F.embedding(code_idx, self.codebook)
return x
def forward(self, x):
N, width, T = x.shape
# Preprocess
x = self.preprocess(x)
# Init codebook if not inited
if self.training and not self.init:
self.init_codebook(x)
# quantize and dequantize through bottleneck
code_idx = self.quantize(x)
x_d = self.dequantize(code_idx)
# Update embeddings
if self.training:
perplexity = self.update_codebook(x, code_idx)
else :
perplexity = self.compute_perplexity(code_idx)
# Loss
commit_loss = F.mse_loss(x, x_d.detach())
# Passthrough
x_d = x + (x_d - x).detach()
# Postprocess
x_d = x_d.view(N, T, -1).permute(0, 2, 1).contiguous() #(N, DIM, T)
return x_d, commit_loss, perplexity
class Quantizer(nn.Module):
def __init__(self, n_e, e_dim, beta):
super(Quantizer, self).__init__()
self.e_dim = e_dim
self.n_e = n_e
self.beta = beta
self.embedding = nn.Embedding(self.n_e, self.e_dim)
self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e)
def forward(self, z):
N, width, T = z.shape
z = self.preprocess(z)
assert z.shape[-1] == self.e_dim
z_flattened = z.contiguous().view(-1, self.e_dim)
# B x V
d = torch.sum(z_flattened ** 2, dim=1, keepdim=True) + \
torch.sum(self.embedding.weight**2, dim=1) - 2 * \
torch.matmul(z_flattened, self.embedding.weight.t())
# B x 1
min_encoding_indices = torch.argmin(d, dim=1)
z_q = self.embedding(min_encoding_indices).view(z.shape)
# compute loss for embedding
loss = torch.mean((z_q - z.detach())**2) + self.beta * \
torch.mean((z_q.detach() - z)**2)
# preserve gradients
z_q = z + (z_q - z).detach()
z_q = z_q.view(N, T, -1).permute(0, 2, 1).contiguous() #(N, DIM, T)
min_encodings = F.one_hot(min_encoding_indices, self.n_e).type(z.dtype)
e_mean = torch.mean(min_encodings, dim=0)
perplexity = torch.exp(-torch.sum(e_mean*torch.log(e_mean + 1e-10)))
return z_q, loss, perplexity
def quantize(self, z):
assert z.shape[-1] == self.e_dim
# B x V
d = torch.sum(z ** 2, dim=1, keepdim=True) + \
torch.sum(self.embedding.weight ** 2, dim=1) - 2 * \
torch.matmul(z, self.embedding.weight.t())
# B x 1
min_encoding_indices = torch.argmin(d, dim=1)
return min_encoding_indices
def dequantize(self, indices):
index_flattened = indices.view(-1)
z_q = self.embedding(index_flattened)
z_q = z_q.view(indices.shape + (self.e_dim, )).contiguous()
return z_q
def preprocess(self, x):
# NCT -> NTC -> [NT, C]
x = x.permute(0, 2, 1).contiguous()
x = x.view(-1, x.shape[-1])
return x
class QuantizeReset(nn.Module):
def __init__(self, nb_code, code_dim, args):
super().__init__()
self.nb_code = nb_code
self.code_dim = code_dim
self.reset_codebook()
self.codebook = nn.Parameter(torch.randn(nb_code, code_dim))
def reset_codebook(self):
self.init = False
self.code_count = None
def _tile(self, x):
nb_code_x, code_dim = x.shape
if nb_code_x < self.nb_code:
n_repeats = (self.nb_code + nb_code_x - 1) // nb_code_x
std = 0.01 / np.sqrt(code_dim)
out = x.repeat(n_repeats, 1)
out = out + torch.randn_like(out) * std
else :
out = x
return out
def init_codebook(self, x):
out = self._tile(x)
self.codebook = nn.Parameter(out[:self.nb_code])
self.code_count = torch.ones(self.nb_code, device=self.codebook.device)
self.init = True
@torch.no_grad()
def compute_perplexity(self, code_idx) :
# Calculate new centres
code_onehot = torch.zeros(self.nb_code, code_idx.shape[0], device=code_idx.device) # nb_code, N * L
code_onehot.scatter_(0, code_idx.view(1, code_idx.shape[0]), 1)
code_count = code_onehot.sum(dim=-1) # nb_code
prob = code_count / torch.sum(code_count)
perplexity = torch.exp(-torch.sum(prob * torch.log(prob + 1e-7)))
return perplexity
def update_codebook(self, x, code_idx):
code_onehot = torch.zeros(self.nb_code, x.shape[0], device=x.device) # nb_code, N * L
code_onehot.scatter_(0, code_idx.view(1, x.shape[0]), 1)
code_count = code_onehot.sum(dim=-1) # nb_code
out = self._tile(x)
code_rand = out[:self.nb_code]
# Update centres
self.code_count = code_count # nb_code
usage = (self.code_count.view(self.nb_code, 1) >= 1.0).float()
self.codebook.data = usage * self.codebook.data + (1 - usage) * code_rand
prob = code_count / torch.sum(code_count)
perplexity = torch.exp(-torch.sum(prob * torch.log(prob + 1e-7)))
return perplexity
def preprocess(self, x):
# NCT -> NTC -> [NT, C]
x = x.permute(0, 2, 1).contiguous()
x = x.view(-1, x.shape[-1])
return x
def quantize(self, x):
# Calculate latent code x_l
k_w = self.codebook.t()
distance = torch.sum(x ** 2, dim=-1, keepdim=True) - 2 * torch.matmul(x, k_w) + torch.sum(k_w ** 2, dim=0,
keepdim=True) # (N * L, b)
_, code_idx = torch.min(distance, dim=-1)
return code_idx
def dequantize(self, code_idx):
x = F.embedding(code_idx, self.codebook)
return x
def forward(self, x):
N, width, T = x.shape
# Preprocess
x = self.preprocess(x)
# Init codebook if not inited
if self.training and not self.init:
self.init_codebook(x)
# quantize and dequantize through bottleneck
code_idx = self.quantize(x)
x_d = self.dequantize(code_idx)
# Update embeddings
if self.training:
perplexity = self.update_codebook(x, code_idx)
else :
perplexity = self.compute_perplexity(code_idx)
# Loss
commit_loss = F.mse_loss(x, x_d.detach())
# Passthrough
x_d = x + (x_d - x).detach()
# Postprocess
x_d = x_d.view(N, T, -1).permute(0, 2, 1).contiguous() #(N, DIM, T)
return x_d, commit_loss, perplexity
class QuantizeEMA(nn.Module):
def __init__(self, nb_code, code_dim, args):
super().__init__()
self.nb_code = nb_code
self.code_dim = code_dim
self.mu = 0.99
self.reset_codebook()
def reset_codebook(self):
self.init = False
self.code_sum = None
self.code_count = None
self.register_buffer('codebook', torch.zeros(self.nb_code, self.code_dim).cuda())
def _tile(self, x):
nb_code_x, code_dim = x.shape
if nb_code_x < self.nb_code:
n_repeats = (self.nb_code + nb_code_x - 1) // nb_code_x
std = 0.01 / np.sqrt(code_dim)
out = x.repeat(n_repeats, 1)
out = out + torch.randn_like(out) * std
else :
out = x
return out
def init_codebook(self, x):
out = self._tile(x)
self.codebook = out[:self.nb_code]
self.code_sum = self.codebook.clone()
self.code_count = torch.ones(self.nb_code, device=self.codebook.device)
self.init = True
@torch.no_grad()
def compute_perplexity(self, code_idx) :
# Calculate new centres
code_onehot = torch.zeros(self.nb_code, code_idx.shape[0], device=code_idx.device) # nb_code, N * L
code_onehot.scatter_(0, code_idx.view(1, code_idx.shape[0]), 1)
code_count = code_onehot.sum(dim=-1) # nb_code
prob = code_count / torch.sum(code_count)
perplexity = torch.exp(-torch.sum(prob * torch.log(prob + 1e-7)))
return perplexity
@torch.no_grad()
def update_codebook(self, x, code_idx):
code_onehot = torch.zeros(self.nb_code, x.shape[0], device=x.device) # nb_code, N * L
code_onehot.scatter_(0, code_idx.view(1, x.shape[0]), 1)
code_sum = torch.matmul(code_onehot, x) # nb_code, w
code_count = code_onehot.sum(dim=-1) # nb_code
# Update centres
self.code_sum = self.mu * self.code_sum + (1. - self.mu) * code_sum # w, nb_code
self.code_count = self.mu * self.code_count + (1. - self.mu) * code_count # nb_code
code_update = self.code_sum.view(self.nb_code, self.code_dim) / self.code_count.view(self.nb_code, 1)
self.codebook = code_update
prob = code_count / torch.sum(code_count)
perplexity = torch.exp(-torch.sum(prob * torch.log(prob + 1e-7)))
return perplexity
def preprocess(self, x):
# NCT -> NTC -> [NT, C]
x = x.permute(0, 2, 1).contiguous()
x = x.view(-1, x.shape[-1])
return x
def quantize(self, x):
# Calculate latent code x_l
k_w = self.codebook.t()
distance = torch.sum(x ** 2, dim=-1, keepdim=True) - 2 * torch.matmul(x, k_w) + torch.sum(k_w ** 2, dim=0,
keepdim=True) # (N * L, b)
_, code_idx = torch.min(distance, dim=-1)
return code_idx
def dequantize(self, code_idx):
x = F.embedding(code_idx, self.codebook)
return x
def forward(self, x):
N, width, T = x.shape
# Preprocess
x = self.preprocess(x)
# Init codebook if not inited
if self.training and not self.init:
self.init_codebook(x)
# quantize and dequantize through bottleneck
code_idx = self.quantize(x)
x_d = self.dequantize(code_idx)
# Update embeddings
if self.training:
perplexity = self.update_codebook(x, code_idx)
else :
perplexity = self.compute_perplexity(code_idx)
# Loss
commit_loss = F.mse_loss(x, x_d.detach())
# Passthrough
x_d = x + (x_d - x).detach()
# Postprocess
x_d = x_d.view(N, T, -1).permute(0, 2, 1).contiguous() #(N, DIM, T)
return x_d, commit_loss, perplexity
+82
View File
@@ -0,0 +1,82 @@
import torch.nn as nn
import torch
class nonlinearity(nn.Module):
def __init__(self):
super().__init__()
def forward(self, x):
# swish
return x * torch.sigmoid(x)
class ResConv1DBlock(nn.Module):
def __init__(self, n_in, n_state, dilation=1, activation='silu', norm=None, dropout=None):
super().__init__()
padding = dilation
self.norm = norm
if norm == "LN":
self.norm1 = nn.LayerNorm(n_in)
self.norm2 = nn.LayerNorm(n_in)
elif norm == "GN":
self.norm1 = nn.GroupNorm(num_groups=32, num_channels=n_in, eps=1e-6, affine=True)
self.norm2 = nn.GroupNorm(num_groups=32, num_channels=n_in, eps=1e-6, affine=True)
elif norm == "BN":
self.norm1 = nn.BatchNorm1d(num_features=n_in, eps=1e-6, affine=True)
self.norm2 = nn.BatchNorm1d(num_features=n_in, eps=1e-6, affine=True)
else:
self.norm1 = nn.Identity()
self.norm2 = nn.Identity()
if activation == "relu":
self.activation1 = nn.ReLU()
self.activation2 = nn.ReLU()
elif activation == "silu":
self.activation1 = nonlinearity()
self.activation2 = nonlinearity()
elif activation == "gelu":
self.activation1 = nn.GELU()
self.activation2 = nn.GELU()
self.conv1 = nn.Conv1d(n_in, n_state, 3, 1, padding, dilation)
self.conv2 = nn.Conv1d(n_state, n_in, 1, 1, 0,)
def forward(self, x):
x_orig = x
if self.norm == "LN":
x = self.norm1(x.transpose(-2, -1))
x = self.activation1(x.transpose(-2, -1))
else:
x = self.norm1(x)
x = self.activation1(x)
x = self.conv1(x)
if self.norm == "LN":
x = self.norm2(x.transpose(-2, -1))
x = self.activation2(x.transpose(-2, -1))
else:
x = self.norm2(x)
x = self.activation2(x)
x = self.conv2(x)
x = x + x_orig
return x
class Resnet1D(nn.Module):
def __init__(self, n_in, n_depth, dilation_growth_rate=1, reverse_dilation=True, activation='relu', norm=None):
super().__init__()
blocks = [ResConv1DBlock(n_in, n_in, dilation=dilation_growth_rate ** depth, activation=activation, norm=norm) for depth in range(n_depth)]
if reverse_dilation:
blocks = blocks[::-1]
self.model = nn.Sequential(*blocks)
def forward(self, x):
return self.model(x)
+92
View File
@@ -0,0 +1,92 @@
# This code is based on https://github.com/Mathux/ACTOR.git
import torch
from ..utils import rotation_conversions as geometry
from ..models.smpl import SMPL, JOINTSTYPE_ROOT
# from .get_model import JOINTSTYPES
JOINTSTYPES = ["a2m", "a2mpl", "smpl", "vibe", "vertices"]
class Rotation2xyz:
def __init__(self, device, dataset='amass'):
self.device = device
self.dataset = dataset
self.smpl_model = SMPL().eval().to(device)
def __call__(self, x, mask, pose_rep, translation, glob,
jointstype, vertstrans, betas=None, 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 betas is None:
betas = torch.zeros([rotations.shape[0], self.smpl_model.num_betas],
dtype=rotations.dtype, device=rotations.device)
betas[:, 1] = beta
# import ipdb; ipdb.set_trace()
out = self.smpl_model(body_pose=rotations, global_orient=global_orient, betas=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
+97
View File
@@ -0,0 +1,97 @@
# 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]
from ..utils.config import SMPL_MODEL_PATH, JOINT_REGRESSOR_TRAIN_EXTRA
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
+211
View File
@@ -0,0 +1,211 @@
import math
import torch
import torch.nn as nn
from torch.nn import functional as F
from torch.distributions import Categorical
from . import pos_encoding as pos_encoding
class Text2Motion_Transformer(nn.Module):
def __init__(self,
num_vq=1024,
embed_dim=512,
clip_dim=512,
block_size=16,
num_layers=2,
n_head=8,
drop_out_rate=0.1,
fc_rate=4):
super().__init__()
self.trans_base = CrossCondTransBase(num_vq, embed_dim, clip_dim, block_size, num_layers, n_head, drop_out_rate, fc_rate)
self.trans_head = CrossCondTransHead(num_vq, embed_dim, block_size, num_layers, n_head, drop_out_rate, fc_rate)
self.block_size = block_size
self.num_vq = num_vq
def get_block_size(self):
return self.block_size
def forward(self, idxs, clip_feature):
feat = self.trans_base(idxs, clip_feature)
logits = self.trans_head(feat)
return logits
def sample(self, clip_feature, if_categorial=False):
for k in range(self.block_size):
if k == 0:
x = []
else:
x = xs
logits = self.forward(x, clip_feature)
logits = logits[:, -1, :]
probs = F.softmax(logits, dim=-1)
if if_categorial:
dist = Categorical(probs)
idx = dist.sample()
if idx == self.num_vq:
break
idx = idx.unsqueeze(-1)
else:
_, idx = torch.topk(probs, k=1, dim=-1)
if idx[0] == self.num_vq:
break
# append to the sequence and continue
if k == 0:
xs = idx
else:
xs = torch.cat((xs, idx), dim=1)
if k == self.block_size - 1:
return xs[:, :-1]
return xs
class CausalCrossConditionalSelfAttention(nn.Module):
def __init__(self, embed_dim=512, block_size=16, n_head=8, drop_out_rate=0.1):
super().__init__()
assert embed_dim % 8 == 0
# key, query, value projections for all heads
self.key = nn.Linear(embed_dim, embed_dim)
self.query = nn.Linear(embed_dim, embed_dim)
self.value = nn.Linear(embed_dim, embed_dim)
self.attn_drop = nn.Dropout(drop_out_rate)
self.resid_drop = nn.Dropout(drop_out_rate)
self.proj = nn.Linear(embed_dim, embed_dim)
# causal mask to ensure that attention is only applied to the left in the input sequence
self.register_buffer("mask", torch.tril(torch.ones(block_size, block_size)).view(1, 1, block_size, block_size))
self.n_head = n_head
def forward(self, x):
B, T, C = x.size()
# calculate query, key, values for all heads in batch and move head forward to be the batch dim
k = self.key(x).view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs)
q = self.query(x).view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs)
v = self.value(x).view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs)
# causal self-attention; Self-attend: (B, nh, T, hs) x (B, nh, hs, T) -> (B, nh, T, T)
att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1)))
att = att.masked_fill(self.mask[:,:,:T,:T] == 0, float('-inf'))
att = F.softmax(att, dim=-1)
att = self.attn_drop(att)
y = att @ v # (B, nh, T, T) x (B, nh, T, hs) -> (B, nh, T, hs)
y = y.transpose(1, 2).contiguous().view(B, T, C) # re-assemble all head outputs side by side
# output projection
y = self.resid_drop(self.proj(y))
return y
class Block(nn.Module):
def __init__(self, embed_dim=512, block_size=16, n_head=8, drop_out_rate=0.1, fc_rate=4):
super().__init__()
self.ln1 = nn.LayerNorm(embed_dim)
self.ln2 = nn.LayerNorm(embed_dim)
self.attn = CausalCrossConditionalSelfAttention(embed_dim, block_size, n_head, drop_out_rate)
self.mlp = nn.Sequential(
nn.Linear(embed_dim, fc_rate * embed_dim),
nn.GELU(),
nn.Linear(fc_rate * embed_dim, embed_dim),
nn.Dropout(drop_out_rate),
)
def forward(self, x):
x = x + self.attn(self.ln1(x))
x = x + self.mlp(self.ln2(x))
return x
class CrossCondTransBase(nn.Module):
def __init__(self,
num_vq=1024,
embed_dim=512,
clip_dim=512,
block_size=16,
num_layers=2,
n_head=8,
drop_out_rate=0.1,
fc_rate=4):
super().__init__()
self.tok_emb = nn.Embedding(num_vq + 2, embed_dim)
self.cond_emb = nn.Linear(clip_dim, embed_dim)
self.pos_embedding = nn.Embedding(block_size, embed_dim)
self.drop = nn.Dropout(drop_out_rate)
# transformer block
self.blocks = nn.Sequential(*[Block(embed_dim, block_size, n_head, drop_out_rate, fc_rate) for _ in range(num_layers)])
self.pos_embed = pos_encoding.PositionEmbedding(block_size, embed_dim, 0.0, False)
self.block_size = block_size
self.apply(self._init_weights)
def get_block_size(self):
return self.block_size
def _init_weights(self, module):
if isinstance(module, (nn.Linear, nn.Embedding)):
module.weight.data.normal_(mean=0.0, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
module.bias.data.zero_()
elif isinstance(module, nn.LayerNorm):
module.bias.data.zero_()
module.weight.data.fill_(1.0)
def forward(self, idx, clip_feature):
if len(idx) == 0:
token_embeddings = self.cond_emb(clip_feature).unsqueeze(1)
else:
b, t = idx.size()
assert t <= self.block_size, "Cannot forward, model block size is exhausted."
# forward the Trans model
token_embeddings = self.tok_emb(idx)
token_embeddings = torch.cat([self.cond_emb(clip_feature).unsqueeze(1), token_embeddings], dim=1)
x = self.pos_embed(token_embeddings)
x = self.blocks(x)
return x
class CrossCondTransHead(nn.Module):
def __init__(self,
num_vq=1024,
embed_dim=512,
block_size=16,
num_layers=2,
n_head=8,
drop_out_rate=0.1,
fc_rate=4):
super().__init__()
self.blocks = nn.Sequential(*[Block(embed_dim, block_size, n_head, drop_out_rate, fc_rate) for _ in range(num_layers)])
self.ln_f = nn.LayerNorm(embed_dim)
self.head = nn.Linear(embed_dim, num_vq + 1, bias=False)
self.block_size = block_size
self.apply(self._init_weights)
def get_block_size(self):
return self.block_size
def _init_weights(self, module):
if isinstance(module, (nn.Linear, nn.Embedding)):
module.weight.data.normal_(mean=0.0, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
module.bias.data.zero_()
elif isinstance(module, nn.LayerNorm):
module.bias.data.zero_()
module.weight.data.fill_(1.0)
def forward(self, x):
x = self.blocks(x)
x = self.ln_f(x)
logits = self.head(x)
return logits
+118
View File
@@ -0,0 +1,118 @@
import torch.nn as nn
from .encdec import Encoder, Decoder
from .quantize_cnn import QuantizeEMAReset, Quantizer, QuantizeEMA, QuantizeReset
class VQVAE_251(nn.Module):
def __init__(self,
args,
nb_code=1024,
code_dim=512,
output_emb_width=512,
down_t=3,
stride_t=2,
width=512,
depth=3,
dilation_growth_rate=3,
activation='relu',
norm=None):
super().__init__()
self.code_dim = code_dim
self.num_code = nb_code
self.quant = args.quantizer
self.encoder = Encoder(251 if args.dataname == 'kit' else 263, output_emb_width, down_t, stride_t, width, depth, dilation_growth_rate, activation=activation, norm=norm)
self.decoder = Decoder(251 if args.dataname == 'kit' else 263, output_emb_width, down_t, stride_t, width, depth, dilation_growth_rate, activation=activation, norm=norm)
if args.quantizer == "ema_reset":
self.quantizer = QuantizeEMAReset(nb_code, code_dim, args)
elif args.quantizer == "orig":
self.quantizer = Quantizer(nb_code, code_dim, 1.0)
elif args.quantizer == "ema":
self.quantizer = QuantizeEMA(nb_code, code_dim, args)
elif args.quantizer == "reset":
self.quantizer = QuantizeReset(nb_code, code_dim, args)
def preprocess(self, x):
# (bs, T, Jx3) -> (bs, Jx3, T)
x = x.permute(0,2,1).float()
return x
def postprocess(self, x):
# (bs, Jx3, T) -> (bs, T, Jx3)
x = x.permute(0,2,1)
return x
def encode(self, x):
N, T, _ = x.shape
x_in = self.preprocess(x)
x_encoder = self.encoder(x_in)
x_encoder = self.postprocess(x_encoder)
x_encoder = x_encoder.contiguous().view(-1, x_encoder.shape[-1]) # (NT, C)
code_idx = self.quantizer.quantize(x_encoder)
code_idx = code_idx.view(N, -1)
return code_idx
def forward(self, x):
x_in = self.preprocess(x)
# Encode
x_encoder = self.encoder(x_in)
## quantization
x_quantized, loss, perplexity = self.quantizer(x_encoder)
## decoder
x_decoder = self.decoder(x_quantized)
x_out = self.postprocess(x_decoder)
return x_out, loss, perplexity
def forward_decoder(self, x):
x_d = self.quantizer.dequantize(x)
x_d = x_d.view(1, -1, self.code_dim).permute(0, 2, 1).contiguous()
# decoder
x_decoder = self.decoder(x_d)
x_out = self.postprocess(x_decoder)
return x_out
class HumanVQVAE(nn.Module):
def __init__(self,
args,
nb_code=512,
code_dim=512,
output_emb_width=512,
down_t=3,
stride_t=2,
width=512,
depth=3,
dilation_growth_rate=3,
activation='relu',
norm=None):
super().__init__()
self.nb_joints = 21 if args.dataname == 'kit' else 22
self.vqvae = VQVAE_251(args, nb_code, code_dim, output_emb_width, down_t, stride_t, width, depth, dilation_growth_rate, activation=activation, norm=norm)
def encode(self, x):
b, t, c = x.size()
quants = self.vqvae.encode(x) # (N, T)
return quants
def forward(self, x):
x_out, loss, perplexity = self.vqvae(x)
return x_out, loss, perplexity
def forward_decoder(self, x):
x_out = self.vqvae.forward_decoder(x)
return x_out
+83
View File
@@ -0,0 +1,83 @@
from argparse import Namespace
import re
from os.path import join as pjoin
def is_float(numStr):
flag = False
numStr = str(numStr).strip().lstrip('-').lstrip('+')
try:
reg = re.compile(r'^[-+]?[0-9]+\.[0-9]+$')
res = reg.match(str(numStr))
if res:
flag = True
except Exception as ex:
print("is_float() - error: " + str(ex))
return flag
def is_number(numStr):
flag = False
numStr = str(numStr).strip().lstrip('-').lstrip('+')
if str(numStr).isdigit():
flag = True
return flag
def get_opt(opt_path, device):
opt = Namespace()
opt_dict = vars(opt)
skip = ('-------------- End ----------------',
'------------ Options -------------',
'\n')
print('Reading', opt_path)
with open(opt_path) as f:
for line in f:
if line.strip() not in skip:
# print(line.strip())
key, value = line.strip().split(': ')
if value in ('True', 'False'):
opt_dict[key] = (value == 'True')
# print(key, value)
elif is_float(value):
opt_dict[key] = float(value)
elif is_number(value):
opt_dict[key] = int(value)
else:
opt_dict[key] = str(value)
# print(opt)
opt_dict['which_epoch'] = 'finest'
opt.save_root = pjoin(opt.checkpoints_dir, opt.dataset_name, opt.name)
opt.model_dir = pjoin(opt.save_root, 'model')
opt.meta_dir = pjoin(opt.save_root, 'meta')
if opt.dataset_name == 't2m':
opt.data_root = './dataset/HumanML3D/'
opt.motion_dir = pjoin(opt.data_root, 'new_joint_vecs')
opt.text_dir = pjoin(opt.data_root, 'texts')
opt.joints_num = 22
opt.dim_pose = 263
opt.max_motion_length = 196
opt.max_motion_frame = 196
opt.max_motion_token = 55
elif opt.dataset_name == 'kit':
opt.data_root = './dataset/KIT-ML/'
opt.motion_dir = pjoin(opt.data_root, 'new_joint_vecs')
opt.text_dir = pjoin(opt.data_root, 'texts')
opt.joints_num = 21
opt.dim_pose = 251
opt.max_motion_length = 196
opt.max_motion_frame = 196
opt.max_motion_token = 55
else:
raise KeyError('Dataset not recognized')
opt.dim_word = 300
opt.num_classes = 200 // opt.unit_length
opt.is_train = False
opt.is_continue = False
opt.device = device
return opt
@@ -0,0 +1,68 @@
import argparse
def get_args_parser():
parser = argparse.ArgumentParser(description='Optimal Transport AutoEncoder training for Amass',
add_help=True,
formatter_class=argparse.ArgumentDefaultsHelpFormatter)
## dataloader
parser.add_argument('--dataname', type=str, default='kit', help='dataset directory')
parser.add_argument('--batch-size', default=128, type=int, help='batch size')
parser.add_argument('--fps', default=[20], nargs="+", type=int, help='frames per second')
parser.add_argument('--seq-len', type=int, default=64, help='training motion length')
## optimization
parser.add_argument('--total-iter', default=100000, type=int, help='number of total iterations to run')
parser.add_argument('--warm-up-iter', default=1000, type=int, help='number of total iterations for warmup')
parser.add_argument('--lr', default=2e-4, type=float, help='max learning rate')
parser.add_argument('--lr-scheduler', default=[60000], nargs="+", type=int, help="learning rate schedule (iterations)")
parser.add_argument('--gamma', default=0.05, type=float, help="learning rate decay")
parser.add_argument('--weight-decay', default=1e-6, type=float, help='weight decay')
parser.add_argument('--decay-option',default='all', type=str, choices=['all', 'noVQ'], help='disable weight decay on codebook')
parser.add_argument('--optimizer',default='adamw', type=str, choices=['adam', 'adamw'], help='disable weight decay on codebook')
## vqvae arch
parser.add_argument("--code-dim", type=int, default=512, help="embedding dimension")
parser.add_argument("--nb-code", type=int, default=512, help="nb of embedding")
parser.add_argument("--mu", type=float, default=0.99, help="exponential moving average to update the codebook")
parser.add_argument("--down-t", type=int, default=3, help="downsampling rate")
parser.add_argument("--stride-t", type=int, default=2, help="stride size")
parser.add_argument("--width", type=int, default=512, help="width of the network")
parser.add_argument("--depth", type=int, default=3, help="depth of the network")
parser.add_argument("--dilation-growth-rate", type=int, default=3, help="dilation growth rate")
parser.add_argument("--output-emb-width", type=int, default=512, help="output embedding width")
parser.add_argument('--vq-act', type=str, default='relu', choices = ['relu', 'silu', 'gelu'], help='dataset directory')
## gpt arch
parser.add_argument("--block-size", type=int, default=25, help="seq len")
parser.add_argument("--embed-dim-gpt", type=int, default=512, help="embedding dimension")
parser.add_argument("--clip-dim", type=int, default=512, help="latent dimension in the clip feature")
parser.add_argument("--num-layers", type=int, default=2, help="nb of transformer layers")
parser.add_argument("--n-head-gpt", type=int, default=8, help="nb of heads")
parser.add_argument("--ff-rate", type=int, default=4, help="feedforward size")
parser.add_argument("--drop-out-rate", type=float, default=0.1, help="dropout ratio in the pos encoding")
## quantizer
parser.add_argument("--quantizer", type=str, default='ema_reset', choices = ['ema', 'orig', 'ema_reset', 'reset'], help="eps for optimal transport")
parser.add_argument('--quantbeta', type=float, default=1.0, help='dataset directory')
## resume
parser.add_argument("--resume-pth", type=str, default=None, help='resume vq pth')
parser.add_argument("--resume-trans", type=str, default=None, help='resume gpt pth')
## output directory
parser.add_argument('--out-dir', type=str, default='output_GPT_Final/', help='output directory')
parser.add_argument('--exp-name', type=str, default='exp_debug', help='name of the experiment, will create a file inside out-dir')
parser.add_argument('--vq-name', type=str, default='exp_debug', help='name of the generated dataset .npy, will create a file inside out-dir')
## other
parser.add_argument('--print-iter', default=200, type=int, help='print frequency')
parser.add_argument('--eval-iter', default=5000, type=int, help='evaluation frequency')
parser.add_argument('--seed', default=123, type=int, help='seed for initializing training. ')
parser.add_argument("--if-maxtest", action='store_true', help="test in max")
parser.add_argument('--pkeep', type=float, default=1.0, help='keep rate for gpt training')
return parser.parse_args()
+61
View File
@@ -0,0 +1,61 @@
import argparse
def get_args_parser():
parser = argparse.ArgumentParser(description='Optimal Transport AutoEncoder training for AIST',
add_help=True,
formatter_class=argparse.ArgumentDefaultsHelpFormatter)
## dataloader
parser.add_argument('--dataname', type=str, default='kit', help='dataset directory')
parser.add_argument('--batch-size', default=128, type=int, help='batch size')
parser.add_argument('--window-size', type=int, default=64, help='training motion length')
## optimization
parser.add_argument('--total-iter', default=200000, type=int, help='number of total iterations to run')
parser.add_argument('--warm-up-iter', default=1000, type=int, help='number of total iterations for warmup')
parser.add_argument('--lr', default=2e-4, type=float, help='max learning rate')
parser.add_argument('--lr-scheduler', default=[50000, 400000], nargs="+", type=int, help="learning rate schedule (iterations)")
parser.add_argument('--gamma', default=0.05, type=float, help="learning rate decay")
parser.add_argument('--weight-decay', default=0.0, type=float, help='weight decay')
parser.add_argument("--commit", type=float, default=0.02, help="hyper-parameter for the commitment loss")
parser.add_argument('--loss-vel', type=float, default=0.1, help='hyper-parameter for the velocity loss')
parser.add_argument('--recons-loss', type=str, default='l2', help='reconstruction loss')
## vqvae arch
parser.add_argument("--code-dim", type=int, default=512, help="embedding dimension")
parser.add_argument("--nb-code", type=int, default=512, help="nb of embedding")
parser.add_argument("--mu", type=float, default=0.99, help="exponential moving average to update the codebook")
parser.add_argument("--down-t", type=int, default=2, help="downsampling rate")
parser.add_argument("--stride-t", type=int, default=2, help="stride size")
parser.add_argument("--width", type=int, default=512, help="width of the network")
parser.add_argument("--depth", type=int, default=3, help="depth of the network")
parser.add_argument("--dilation-growth-rate", type=int, default=3, help="dilation growth rate")
parser.add_argument("--output-emb-width", type=int, default=512, help="output embedding width")
parser.add_argument('--vq-act', type=str, default='relu', choices = ['relu', 'silu', 'gelu'], help='dataset directory')
parser.add_argument('--vq-norm', type=str, default=None, help='dataset directory')
## quantizer
parser.add_argument("--quantizer", type=str, default='ema_reset', choices = ['ema', 'orig', 'ema_reset', 'reset'], help="eps for optimal transport")
parser.add_argument('--beta', type=float, default=1.0, help='commitment loss in standard VQ')
## resume
parser.add_argument("--resume-pth", type=str, default=None, help='resume pth for VQ')
parser.add_argument("--resume-gpt", type=str, default=None, help='resume pth for GPT')
## output directory
parser.add_argument('--out-dir', type=str, default='output_vqfinal/', help='output directory')
parser.add_argument('--results-dir', type=str, default='visual_results/', help='output directory')
parser.add_argument('--visual-name', type=str, default='baseline', help='output directory')
parser.add_argument('--exp-name', type=str, default='exp_debug', help='name of the experiment, will create a file inside out-dir')
## other
parser.add_argument('--print-iter', default=200, type=int, help='print frequency')
parser.add_argument('--eval-iter', default=1000, type=int, help='evaluation frequency')
parser.add_argument('--seed', default=123, type=int, help='seed for initializing training.')
parser.add_argument('--vis-gt', action='store_true', help='whether visualize GT motions')
parser.add_argument('--nb-vis', default=20, type=int, help='nb of visualizations')
return parser.parse_args()
+194
View File
@@ -0,0 +1,194 @@
from models.rotation2xyz import Rotation2xyz
import numpy as np
from trimesh import Trimesh
import os
os.environ['PYOPENGL_PLATFORM'] = "osmesa"
import torch
from visualize.simplify_loc2rot import joints2smpl
import pyrender
import matplotlib.pyplot as plt
import io
import imageio
from shapely import geometry
import trimesh
from pyrender.constants import RenderFlags
import math
# import ffmpeg
from PIL import Image
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, outdir='test_vis', device_id=0, name=None, pred=True):
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_id=0, cuda=True)
rot2xyz = Rotation2xyz(device=torch.device("cuda:0"))
faces = rot2xyz.smpl_model.faces
if (not os.path.exists(outdir + name+'_pred.pt') and pred) or (not os.path.exists(outdir + name+'_gt.pt') and not pred):
print(f'Running SMPLify, it may take a few minutes.')
motion_tensor, opt_dict = j2s.joint2smpl(motions) # [nframes, njoints, 3]
vertices = rot2xyz(torch.tensor(motion_tensor).clone(), mask=None,
pose_rep='rot6d', translation=True, glob=True,
jointstype='vertices',
vertstrans=True)
if pred:
torch.save(vertices, outdir + name+'_pred.pt')
else:
torch.save(vertices, outdir + name+'_gt.pt')
else:
if pred:
vertices = torch.load(outdir + name+'_pred.pt')
else:
vertices = torch.load(outdir + name+'_gt.pt')
frames = vertices.shape[3] # shape: 1, nb_frames, 3, nb_joints
print (vertices.shape)
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)
if pred:
imageio.mimsave(outdir + name+'_pred.gif', out, fps=20)
else:
imageio.mimsave(outdir + name+'_gt.gif', out, fps=20)
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--filedir", type=str, default=None, help='motion npy file dir')
parser.add_argument('--motion-list', default=None, nargs="+", type=str, help="motion name list")
args = parser.parse_args()
filename_list = args.motion_list
filedir = args.filedir
for filename in filename_list:
motions = np.load(filedir + filename+'_pred.npy')
print('pred', motions.shape, filename)
render(motions[0], outdir=filedir, device_id=0, name=filename, pred=True)
motions = np.load(filedir + filename+'_gt.npy')
print('gt', motions.shape, filename)
render(motions[0], outdir=filedir, device_id=0, name=filename, pred=False)
Binary file not shown.

After

Width:  |  Height:  |  Size: 380 KiB

+359
View File
@@ -0,0 +1,359 @@
#from __future__ import absolute_import
import sys
import io
import os
sys.argv = ['GPT_eval_multi.py']
# 将项目根目录添加到sys.path中
PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(1, PROJECT_ROOT)
CKPT_ROOT="/cfs-datasets/public_models/motion"
from .options import option_transformer as option_trans
import sys
print(sys.path[0])
import clip
import torch
import cv2
import numpy as np
from .models import vqvae as vqvae
from .models import t2m_trans as trans
import warnings
from .visualization import plot_3d_global as plot_3d
import matplotlib.pyplot as plt
import numpy as np
import matplotlib.colors as mcolors
from tqdm import tqdm
from mpl_toolkits.mplot3d import Axes3D
from PIL import Image
import time
import random
warnings.filterwarnings('ignore')
from matplotlib.axes._axes import _log as matplotlib_axes_logger
matplotlib_axes_logger.setLevel('ERROR')
from math import cos,sin,radians
args = option_trans.get_args_parser()
args.dataname = 't2m'
args.resume_pth = os.path.join(CKPT_ROOT,'pretrained/VQVAE/net_last.pth')
args.resume_trans = os.path.join(CKPT_ROOT,'pretrained/VQTransformer_corruption05/net_best_fid.pth')
args.down_t = 2
args.depth = 3
args.block_size = 51
def replace_space_with_underscore(s):
return s.replace(' ', '_')
def Rz(angle):
theta=radians(angle)
return np.array([[cos(theta), -sin(theta), 0],
[sin(theta), cos(theta), 0],
[0, 0, 1]])
def Rx(angle):
theta=radians(angle)
return np.array(
[[1, 0, 0],
[0 , cos(theta), -sin(theta)],
[0, sin(theta), cos(theta)]])
def generate_cuid():
timestamp = hex(int(time.time() * 1000))[2:]
random_str = hex(random.randint(0, 0xfffff))[2:]
return (timestamp + random_str).zfill(10)
def smpl_to_openpose18(smpl_keypoints):
'''
22关键点SMPL对应关系解释
[0, 2, 5, 8, 11]
这个列表表示SMPL模型中左腿的连接方式,从骨盆(0号关键点)开始,连接左大腿(2号关键点)、左小腿(5号关键点)、左脚(8号关键点)和左脚尖(11号关键点)。
[0, 1, 4, 7, 10]
这个列表表示SMPL模型中右腿的连接方式,从骨盆(0号关键点)开始,连接右大腿(1号关键点)、右小腿(4号关键点)、右脚(7号关键点)和右脚尖(10号关键点)。
[0, 3, 6, 9, 12, 15]
这个列表表示SMPL模型中躯干的连接方式,从骨盆(0号关键点)开始,连接脊柱(3号关键点)、颈部(6号关键点)、头部(9号关键点)、左肩膀(12号关键点)、右肩膀(15号关键点)。
[9, 14, 17, 19, 21]
这个列表表示SMPL模型中左臂的连接方式,从左肩膀(9号关键点)开始,连接左上臂(14号关键点)、左前臂(17号关键点)、左手腕(19号关键点)和左手(21号关键点)。
[9, 13, 16, 18, 20]
这个列表表示SMPL模型中右臂的连接方式,从右肩膀(9号关键点)开始,连接右上臂(13号关键点)、右前臂(16号关键点)、右手腕(18号关键点)和右手(20号关键点)。
目前转Openpose忽略掉了SMPL的肩膀关键点
'''
openpose_keypoints = np.zeros((18, 3))
openpose_keypoints[0] = smpl_keypoints[9] # nose
openpose_keypoints[0][1] = openpose_keypoints[0][1]+0.3 #
openpose_keypoints[1] = smpl_keypoints[6] # neck
openpose_keypoints[2] = smpl_keypoints[16] # right shoulder
openpose_keypoints[3] = smpl_keypoints[18] # right elbow
openpose_keypoints[4] = smpl_keypoints[20] # right wrist
openpose_keypoints[5] = smpl_keypoints[17] # left shoulder
openpose_keypoints[6] = smpl_keypoints[19] # left elbow
openpose_keypoints[7] = smpl_keypoints[21] # left wrist
#TODO: Experiment,将neck的关键点抬高&&将nose的关键点相对高度关系与neck保持一致
openpose_keypoints[1][0]=(openpose_keypoints[2][0]+openpose_keypoints[5][0])/2
openpose_keypoints[1][1]=(openpose_keypoints[2][1]+openpose_keypoints[5][1])/2
openpose_keypoints[1][2]=(openpose_keypoints[2][2]+openpose_keypoints[5][2])/2
openpose_keypoints[0][1] = openpose_keypoints[1][1]+0.3 #
openpose_keypoints[8] = smpl_keypoints[1] # right hip
openpose_keypoints[9] = smpl_keypoints[4] # right knee
openpose_keypoints[10] = smpl_keypoints[7] # right ankle
openpose_keypoints[11] = smpl_keypoints[2] # left hip
openpose_keypoints[12] = smpl_keypoints[5] # left knee
openpose_keypoints[13] = smpl_keypoints[8] # left ankle
#TODO: Experiment,手工指定脸部关键点测试是否能够指定身体朝向
#openpose_keypoints[0][0] = openpose_keypoints[0][0]+0.3#测试0坐标轴方向(水平向右)
#openpose_keypoints[0][2] = openpose_keypoints[0][2]#测试2坐标轴方向(向外
#openpose_keypoints[0][1] = openpose_keypoints[0][1]+0.5#测试1坐标轴方向(垂直向上
openpose_keypoints[14] = openpose_keypoints[0] # right eye
openpose_keypoints[14][1]=openpose_keypoints[14][1]+0.05
openpose_keypoints[14][0]=openpose_keypoints[14][0]+0.3*(openpose_keypoints[2][0]-openpose_keypoints[1][0])
openpose_keypoints[14][2]=openpose_keypoints[14][2]+0.3*(openpose_keypoints[2][2]-openpose_keypoints[1][2])
openpose_keypoints[15] = openpose_keypoints[0] # left eye
openpose_keypoints[15][1]=openpose_keypoints[15][1]+0.05
openpose_keypoints[15][0]=openpose_keypoints[15][0]+0.3*(openpose_keypoints[5][0]-openpose_keypoints[1][0])
openpose_keypoints[15][2]=openpose_keypoints[15][2]+0.3*(openpose_keypoints[5][2]-openpose_keypoints[1][2])
openpose_keypoints[16] = openpose_keypoints[0] # right ear
openpose_keypoints[16][0]=openpose_keypoints[16][0]+0.7*(openpose_keypoints[2][0]-openpose_keypoints[1][0])
openpose_keypoints[16][2]=openpose_keypoints[16][2]+0.7*(openpose_keypoints[2][2]-openpose_keypoints[1][2])
openpose_keypoints[17] = openpose_keypoints[0] # left ear
openpose_keypoints[17][0]=openpose_keypoints[17][0]+0.7*(openpose_keypoints[5][0]-openpose_keypoints[1][0])
openpose_keypoints[17][2]=openpose_keypoints[17][2]+0.7*(openpose_keypoints[5][2]-openpose_keypoints[1][2])
return openpose_keypoints
# TODO: debug only, need to be deleted before unload
## load clip model and datasets
clip_model, clip_preprocess = clip.load("ViT-B/32", device=torch.device('cuda'), jit=False, download_root=CKPT_ROOT) # 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
print("loaded CLIP model")
net = vqvae.HumanVQVAE(args, ## use args to define different parameters in different quantizers
args.nb_code,
args.code_dim,
args.output_emb_width,
args.down_t,
args.stride_t,
args.width,
args.depth,
args.dilation_growth_rate)
trans_encoder = trans.Text2Motion_Transformer(num_vq=args.nb_code,
embed_dim=1024,
clip_dim=args.clip_dim,
block_size=args.block_size,
num_layers=9,
n_head=16,
drop_out_rate=args.drop_out_rate,
fc_rate=args.ff_rate)
print ('loading checkpoint from {}'.format(args.resume_pth))
ckpt = torch.load(args.resume_pth, map_location='cpu')
net.load_state_dict(ckpt['net'], strict=True)
net.eval()
net.cuda()
print ('loading transformer checkpoint from {}'.format(args.resume_trans))
ckpt = torch.load(args.resume_trans, map_location='cpu')
trans_encoder.load_state_dict(ckpt['trans'], strict=True)
trans_encoder.eval()
trans_encoder.cuda()
mean = torch.from_numpy(np.load(os.path.join(CKPT_ROOT,'./checkpoints/t2m/VQVAEV3_CB1024_CMT_H1024_NRES3/meta/mean.npy'))).cuda()
std = torch.from_numpy(np.load(os.path.join(CKPT_ROOT,'./checkpoints/t2m/VQVAEV3_CB1024_CMT_H1024_NRES3/meta/std.npy'))).cuda()
def get_open_pose(text,height,width,save_path,video_length):
CKPT_ROOT = os.path.dirname(os.path.abspath(__file__))
clip_text=[text]
print(f"Motion Prompt: {text}")
# cuid=generate_cuid()
# print(f"Motion Generation cuid: {cuid}")
# clip_text = ["the person jump and spin twice,then running straght and sit down. "] #支持单个token的生成
# change the text here
text = clip.tokenize(clip_text, truncate=False).cuda()
feat_clip_text = clip_model.encode_text(text).float()
index_motion = trans_encoder.sample(feat_clip_text[0:1], False)
pred_pose = net.forward_decoder(index_motion)
from utils.motion_process import recover_from_ric
pred_xyz = recover_from_ric((pred_pose*std+mean).float(), 22)
xyz = pred_xyz.reshape(1, -1, 22, 3)
np.save('motion.npy', xyz.detach().cpu().numpy())
pose_vis = plot_3d.draw_to_batch(xyz.detach().cpu().numpy(),clip_text, ['smpl.gif'])
res=xyz.detach().cpu().numpy()
points_3d_list=res[0]
frame_num=points_3d_list.shape[0]
open_pose_list=np.array(points_3d_list)
print("The total SMPL sequence shape is : "+str(open_pose_list.shape))
max_val = np.max(open_pose_list, axis=(0, 1))
min_val = np.min(open_pose_list, axis=(0, 1))
print("三维坐标在坐标系上的最大值:", max_val)
print("三维坐标在坐标系上的最小值:", min_val)
check= smpl_to_openpose18(open_pose_list[0]) # 18个关键点
print("********SMPL_2_OpenPose_List(14/18)********")
print(check)
print("*************************")
print(f"Total Frame Number: {frame_num}")
img_list=[]
for step in tqdm(range(0,frame_num)):
# 生成图像
dpi=84
fig =plt.figure(figsize=(width/dpi, height/dpi), dpi=dpi)
ax = fig.add_subplot(111, projection='3d')
limits=2
ax.set_xlim(-limits*0.7, limits*0.7)
ax.set_ylim(0, limits*1.5)#上下
ax.set_zlim(0, limits*1.5)# 前后
ax.grid(b=False)
#ax.dist = 1
ax.set_box_aspect([1.4, 1.5, 1.5],zoom=3.5)# 坐标轴比例 TODO:这个比例可能有问题,会出现超出坐标范围的bug
# 关键点坐标,每行包含(x, y, z)
keypoints = smpl_to_openpose18(open_pose_list[step]) # 18个关键点
# 运动学链 目前只用到body部分
kinematic_chain = [(0, 1), (1, 2), (2, 3), (3, 4), (1, 5), (5, 6), (6, 7), (1, 8), (8, 9), (9, 10), (1, 11), (11, 12), (12, 13), (0, 14), (14, 16), (0, 15), (15, 17)]
#kinematic_chain = [(0, 1), (1, 2), (2, 3), (3, 4), (1, 5), (5, 6), (6, 7), (1, 8), (8, 9), (9, 10), (1, 11), (11, 12), (12, 13)]
# 颜色RGB
colors = [(0, 0, 255), (0, 255, 255), (0, 255, 0), (255, 0, 0), (255, 0, 255), (255, 192, 203), (0, 165, 255), (19, 69, 139), (173, 216, 230), (34, 139, 34), (0, 0, 128), (184, 134, 11), (139, 0, 139), (0, 100, 0), (0, 255, 255), (0, 255, 0), (216, 191, 216), (255, 255, 224)]
#colors=[(0, 0, 255), (0, 255, 255), (0, 255, 0), (255, 0, 0), (255, 0, 255), (255, 192, 203), (0, 165, 255), (19, 69, 139), (173, 216, 230), (34, 139, 34), (0, 0, 128), (184, 134, 11), (139, 0, 139), (0, 100, 0)]
#18点
joint_colors=[(255,0,0),(255,85,0),(255,170,0),(255,255,0),(170,255,0),(85,255,0),(0,255,0),(0,255,85),(0,255,170),(0,255,255),(0,170,255),(0,85,255),(0,0,255),(85,0,255),(170,0,255),(255,0,255),(255,0,170),(255,0,85),(255,0,0)]
#14点主干
#joint_colors=[(255,0,0),(255,85,0),(255,170,0),(255,255,0),(170,255,0),(85,255,0),(0,255,0),(0,255,85),(0,255,170),(0,255,255),(0,170,255),(0,85,255),(0,0,255),(85,0,255),(170,0,255)]
#运动链连线是joint颜色的60%
#plt颜色在0-1之间
rgb_color2=[]
joint_rgb_color2=[]
kinematic_chain_rgb_color2=[]
for color in joint_colors:
joint_rgb_color2.append(tuple([x/255 for x in color]))
kinematic_chain_rgb_color2.append(tuple([x*0.6/255 for x in color])) #运动链连线是joint颜色的60%
# 可视化结果
for i in range(0,18):
# 绘制关键点
ax.scatter(keypoints[i][0], keypoints[i][1], keypoints[i][2], s=50, c=joint_rgb_color2[i], marker='o')
# 绘制运动学链
for j in range(len(kinematic_chain)):
if kinematic_chain[j][1] == i:
ax.plot([keypoints[kinematic_chain[j][0]][0], keypoints[kinematic_chain[j][1]][0]], [keypoints[kinematic_chain[j][0]][1], keypoints[kinematic_chain[j][1]][1]], [keypoints[kinematic_chain[j][0]][2], keypoints[kinematic_chain[j][1]][2]], c=kinematic_chain_rgb_color2[i], linewidth=5)
# 调整视角
ax.view_init(elev=110, azim=-90)
plt.axis('off')
# 保存图片
# 将图像数据输出为图像数组
if not os.path.exists(save_path):
os.makedirs(save_path)
image_tmp_path=str(f"{save_path}/{str(step)}.jpg")
plt.savefig(os.path.join(CKPT_ROOT,image_tmp_path))#RGB
img=cv2.imread(os.path.join(CKPT_ROOT,image_tmp_path))
img=cv2.cvtColor(img,cv2.COLOR_BGR2RGB)
img_list.append(img)
res=[]
if len(img_list)>=video_length:
key_frame_sample_step=int(len(img_list)/video_length)
else:
print("ERROR: video length is too long")
key_frame_sample_step=1
for i in range(0,len(img_list),key_frame_sample_step):
res.append(img_list[i])
return res
def offline_get_open_pose(text,motion_text,height,width,save_path):
#motion_text=text
clip_text=[text]
print(f"Motion Prompt: {text}")
cuid=generate_cuid()
print(f"Motion Generation cuid: {cuid}")
# clip_text = ["the person jump and spin twice,then running straght and sit down. "] #支持单个token的生成
# change the text here
text = clip.tokenize(clip_text, truncate=False).cuda()
feat_clip_text = clip_model.encode_text(text).float()
index_motion = trans_encoder.sample(feat_clip_text[0:1], False)
pred_pose = net.forward_decoder(index_motion)
from utils.motion_process import recover_from_ric
pred_xyz = recover_from_ric((pred_pose*std+mean).float(), 22)
xyz = pred_xyz.reshape(1, -1, 22, 3)
res=xyz.detach().cpu().numpy()
np.save(f'{save_path}/{replace_space_with_underscore(motion_text)}.npy', res)
pose_vis = plot_3d.draw_to_batch(res,clip_text, ['smpl.gif'])
if __name__ == "__main__":
text="walk around, jump, run straght."
pose = get_open_pose(text,512,512)
#pdb.set_trace()
+191
View File
@@ -0,0 +1,191 @@
import os
import torch
import numpy as np
from torch.utils.tensorboard import SummaryWriter
from os.path import join as pjoin
from torch.distributions import Categorical
import json
import clip
import options.option_transformer as option_trans
import models.vqvae as vqvae
import utils.utils_model as utils_model
import utils.eval_trans as eval_trans
from dataset import dataset_TM_train
from dataset import dataset_TM_eval
from dataset import dataset_tokenize
import models.t2m_trans as trans
from options.get_eval_option import get_opt
from models.evaluator_wrapper import EvaluatorModelWrapper
import warnings
warnings.filterwarnings('ignore')
##### ---- Exp dirs ---- #####
args = option_trans.get_args_parser()
torch.manual_seed(args.seed)
args.out_dir = os.path.join(args.out_dir, f'{args.exp_name}')
args.vq_dir= os.path.join("./dataset/KIT-ML" if args.dataname == 'kit' else "./dataset/HumanML3D", f'{args.vq_name}')
os.makedirs(args.out_dir, exist_ok = True)
os.makedirs(args.vq_dir, exist_ok = True)
##### ---- Logger ---- #####
logger = utils_model.get_logger(args.out_dir)
writer = SummaryWriter(args.out_dir)
logger.info(json.dumps(vars(args), indent=4, sort_keys=True))
##### ---- Dataloader ---- #####
train_loader_token = dataset_tokenize.DATALoader(args.dataname, 1, unit_length=2**args.down_t)
from utils.word_vectorizer import WordVectorizer
w_vectorizer = WordVectorizer('./glove', 'our_vab')
val_loader = dataset_TM_eval.DATALoader(args.dataname, False, 32, w_vectorizer)
dataset_opt_path = 'checkpoints/kit/Comp_v6_KLD005/opt.txt' if args.dataname == 'kit' else 'checkpoints/t2m/Comp_v6_KLD005/opt.txt'
wrapper_opt = get_opt(dataset_opt_path, torch.device('cuda'))
eval_wrapper = EvaluatorModelWrapper(wrapper_opt)
##### ---- Network ---- #####
clip_model, clip_preprocess = clip.load("ViT-B/32", device=torch.device('cuda'), 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
net = vqvae.HumanVQVAE(args, ## use args to define different parameters in different quantizers
args.nb_code,
args.code_dim,
args.output_emb_width,
args.down_t,
args.stride_t,
args.width,
args.depth,
args.dilation_growth_rate)
trans_encoder = trans.Text2Motion_Transformer(num_vq=args.nb_code,
embed_dim=args.embed_dim_gpt,
clip_dim=args.clip_dim,
block_size=args.block_size,
num_layers=args.num_layers,
n_head=args.n_head_gpt,
drop_out_rate=args.drop_out_rate,
fc_rate=args.ff_rate)
print ('loading checkpoint from {}'.format(args.resume_pth))
ckpt = torch.load(args.resume_pth, map_location='cpu')
net.load_state_dict(ckpt['net'], strict=True)
net.eval()
net.cuda()
if args.resume_trans is not None:
print ('loading transformer checkpoint from {}'.format(args.resume_trans))
ckpt = torch.load(args.resume_trans, map_location='cpu')
trans_encoder.load_state_dict(ckpt['trans'], strict=True)
trans_encoder.train()
trans_encoder.cuda()
##### ---- Optimizer & Scheduler ---- #####
optimizer = utils_model.initial_optim(args.decay_option, args.lr, args.weight_decay, trans_encoder, args.optimizer)
scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, milestones=args.lr_scheduler, gamma=args.gamma)
##### ---- Optimization goals ---- #####
loss_ce = torch.nn.CrossEntropyLoss()
nb_iter, avg_loss_cls, avg_acc = 0, 0., 0.
right_num = 0
nb_sample_train = 0
##### ---- get code ---- #####
for batch in train_loader_token:
pose, name = batch
bs, seq = pose.shape[0], pose.shape[1]
pose = pose.cuda().float() # bs, nb_joints, joints_dim, seq_len
target = net.encode(pose)
target = target.cpu().numpy()
np.save(pjoin(args.vq_dir, name[0] +'.npy'), target)
train_loader = dataset_TM_train.DATALoader(args.dataname, args.batch_size, args.nb_code, args.vq_name, unit_length=2**args.down_t)
train_loader_iter = dataset_TM_train.cycle(train_loader)
##### ---- Training ---- #####
best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, writer, logger = eval_trans.evaluation_transformer(args.out_dir, val_loader, net, trans_encoder, logger, writer, 0, best_fid=1000, best_iter=0, best_div=100, best_top1=0, best_top2=0, best_top3=0, best_matching=100, clip_model=clip_model, eval_wrapper=eval_wrapper)
while nb_iter <= args.total_iter:
batch = next(train_loader_iter)
clip_text, m_tokens, m_tokens_len = batch
m_tokens, m_tokens_len = m_tokens.cuda(), m_tokens_len.cuda()
bs = m_tokens.shape[0]
target = m_tokens # (bs, 26)
target = target.cuda()
text = clip.tokenize(clip_text, truncate=True).cuda()
feat_clip_text = clip_model.encode_text(text).float()
input_index = target[:,:-1]
if args.pkeep == -1:
proba = np.random.rand(1)[0]
mask = torch.bernoulli(proba * torch.ones(input_index.shape,
device=input_index.device))
else:
mask = torch.bernoulli(args.pkeep * torch.ones(input_index.shape,
device=input_index.device))
mask = mask.round().to(dtype=torch.int64)
r_indices = torch.randint_like(input_index, args.nb_code)
a_indices = mask*input_index+(1-mask)*r_indices
cls_pred = trans_encoder(a_indices, feat_clip_text)
cls_pred = cls_pred.contiguous()
loss_cls = 0.0
for i in range(bs):
# loss function (26), (26, 513)
loss_cls += loss_ce(cls_pred[i][:m_tokens_len[i] + 1], target[i][:m_tokens_len[i] + 1]) / bs
# Accuracy
probs = torch.softmax(cls_pred[i][:m_tokens_len[i] + 1], dim=-1)
if args.if_maxtest:
_, cls_pred_index = torch.max(probs, dim=-1)
else:
dist = Categorical(probs)
cls_pred_index = dist.sample()
right_num += (cls_pred_index.flatten(0) == target[i][:m_tokens_len[i] + 1].flatten(0)).sum().item()
## global loss
optimizer.zero_grad()
loss_cls.backward()
optimizer.step()
scheduler.step()
avg_loss_cls = avg_loss_cls + loss_cls.item()
nb_sample_train = nb_sample_train + (m_tokens_len + 1).sum().item()
nb_iter += 1
if nb_iter % args.print_iter == 0 :
avg_loss_cls = avg_loss_cls / args.print_iter
avg_acc = right_num * 100 / nb_sample_train
writer.add_scalar('./Loss/train', avg_loss_cls, nb_iter)
writer.add_scalar('./ACC/train', avg_acc, nb_iter)
msg = f"Train. Iter {nb_iter} : Loss. {avg_loss_cls:.5f}, ACC. {avg_acc:.4f}"
logger.info(msg)
avg_loss_cls = 0.
right_num = 0
nb_sample_train = 0
if nb_iter % args.eval_iter == 0:
best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, writer, logger = eval_trans.evaluation_transformer(args.out_dir, val_loader, net, trans_encoder, logger, writer, nb_iter, best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, clip_model=clip_model, eval_wrapper=eval_wrapper)
if nb_iter == args.total_iter:
msg_final = f"Train. Iter {best_iter} : FID. {best_fid:.5f}, Diversity. {best_div:.4f}, TOP1. {best_top1:.4f}, TOP2. {best_top2:.4f}, TOP3. {best_top3:.4f}"
logger.info(msg_final)
break
+171
View File
@@ -0,0 +1,171 @@
import os
import json
import torch
import torch.optim as optim
from torch.utils.tensorboard import SummaryWriter
import models.vqvae as vqvae
import utils.losses as losses
import options.option_vq as option_vq
import utils.utils_model as utils_model
from dataset import dataset_VQ, dataset_TM_eval
import utils.eval_trans as eval_trans
from options.get_eval_option import get_opt
from models.evaluator_wrapper import EvaluatorModelWrapper
import warnings
warnings.filterwarnings('ignore')
from utils.word_vectorizer import WordVectorizer
def update_lr_warm_up(optimizer, nb_iter, warm_up_iter, lr):
current_lr = lr * (nb_iter + 1) / (warm_up_iter + 1)
for param_group in optimizer.param_groups:
param_group["lr"] = current_lr
return optimizer, current_lr
##### ---- Exp dirs ---- #####
args = option_vq.get_args_parser()
torch.manual_seed(args.seed)
args.out_dir = os.path.join(args.out_dir, f'{args.exp_name}')
os.makedirs(args.out_dir, exist_ok = True)
##### ---- Logger ---- #####
logger = utils_model.get_logger(args.out_dir)
writer = SummaryWriter(args.out_dir)
logger.info(json.dumps(vars(args), indent=4, sort_keys=True))
w_vectorizer = WordVectorizer('./glove', 'our_vab')
if args.dataname == 'kit' :
dataset_opt_path = 'checkpoints/kit/Comp_v6_KLD005/opt.txt'
args.nb_joints = 21
else :
dataset_opt_path = 'checkpoints/t2m/Comp_v6_KLD005/opt.txt'
args.nb_joints = 22
logger.info(f'Training on {args.dataname}, motions are with {args.nb_joints} joints')
wrapper_opt = get_opt(dataset_opt_path, torch.device('cuda'))
eval_wrapper = EvaluatorModelWrapper(wrapper_opt)
##### ---- Dataloader ---- #####
train_loader = dataset_VQ.DATALoader(args.dataname,
args.batch_size,
window_size=args.window_size,
unit_length=2**args.down_t)
train_loader_iter = dataset_VQ.cycle(train_loader)
val_loader = dataset_TM_eval.DATALoader(args.dataname, False,
32,
w_vectorizer,
unit_length=2**args.down_t)
##### ---- Network ---- #####
net = vqvae.HumanVQVAE(args, ## use args to define different parameters in different quantizers
args.nb_code,
args.code_dim,
args.output_emb_width,
args.down_t,
args.stride_t,
args.width,
args.depth,
args.dilation_growth_rate,
args.vq_act,
args.vq_norm)
if args.resume_pth :
logger.info('loading checkpoint from {}'.format(args.resume_pth))
ckpt = torch.load(args.resume_pth, map_location='cpu')
net.load_state_dict(ckpt['net'], strict=True)
net.train()
net.cuda()
##### ---- Optimizer & Scheduler ---- #####
optimizer = optim.AdamW(net.parameters(), lr=args.lr, betas=(0.9, 0.99), weight_decay=args.weight_decay)
scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, milestones=args.lr_scheduler, gamma=args.gamma)
Loss = losses.ReConsLoss(args.recons_loss, args.nb_joints)
##### ------ warm-up ------- #####
avg_recons, avg_perplexity, avg_commit = 0., 0., 0.
for nb_iter in range(1, args.warm_up_iter):
optimizer, current_lr = update_lr_warm_up(optimizer, nb_iter, args.warm_up_iter, args.lr)
gt_motion = next(train_loader_iter)
gt_motion = gt_motion.cuda().float() # (bs, 64, dim)
pred_motion, loss_commit, perplexity = net(gt_motion)
loss_motion = Loss(pred_motion, gt_motion)
loss_vel = Loss.forward_vel(pred_motion, gt_motion)
loss = loss_motion + args.commit * loss_commit + args.loss_vel * loss_vel
optimizer.zero_grad()
loss.backward()
optimizer.step()
avg_recons += loss_motion.item()
avg_perplexity += perplexity.item()
avg_commit += loss_commit.item()
if nb_iter % args.print_iter == 0 :
avg_recons /= args.print_iter
avg_perplexity /= args.print_iter
avg_commit /= args.print_iter
logger.info(f"Warmup. Iter {nb_iter} : lr {current_lr:.5f} \t Commit. {avg_commit:.5f} \t PPL. {avg_perplexity:.2f} \t Recons. {avg_recons:.5f}")
avg_recons, avg_perplexity, avg_commit = 0., 0., 0.
##### ---- Training ---- #####
avg_recons, avg_perplexity, avg_commit = 0., 0., 0.
best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, writer, logger = eval_trans.evaluation_vqvae(args.out_dir, val_loader, net, logger, writer, 0, best_fid=1000, best_iter=0, best_div=100, best_top1=0, best_top2=0, best_top3=0, best_matching=100, eval_wrapper=eval_wrapper)
for nb_iter in range(1, args.total_iter + 1):
gt_motion = next(train_loader_iter)
gt_motion = gt_motion.cuda().float() # bs, nb_joints, joints_dim, seq_len
pred_motion, loss_commit, perplexity = net(gt_motion)
loss_motion = Loss(pred_motion, gt_motion)
loss_vel = Loss.forward_vel(pred_motion, gt_motion)
loss = loss_motion + args.commit * loss_commit + args.loss_vel * loss_vel
optimizer.zero_grad()
loss.backward()
optimizer.step()
scheduler.step()
avg_recons += loss_motion.item()
avg_perplexity += perplexity.item()
avg_commit += loss_commit.item()
if nb_iter % args.print_iter == 0 :
avg_recons /= args.print_iter
avg_perplexity /= args.print_iter
avg_commit /= args.print_iter
writer.add_scalar('./Train/L1', avg_recons, nb_iter)
writer.add_scalar('./Train/PPL', avg_perplexity, nb_iter)
writer.add_scalar('./Train/Commit', avg_commit, nb_iter)
logger.info(f"Train. Iter {nb_iter} : \t Commit. {avg_commit:.5f} \t PPL. {avg_perplexity:.2f} \t Recons. {avg_recons:.5f}")
avg_recons, avg_perplexity, avg_commit = 0., 0., 0.,
if nb_iter % args.eval_iter==0 :
best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, writer, logger = eval_trans.evaluation_vqvae(args.out_dir, val_loader, net, logger, writer, nb_iter, best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, eval_wrapper=eval_wrapper)
+17
View File
@@ -0,0 +1,17 @@
import os
SMPL_DATA_PATH = "/group/30065/users/zhanchao/code/MMCM/mmcm/t2p/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
+580
View File
@@ -0,0 +1,580 @@
import os
import clip
import numpy as np
import torch
from scipy import linalg
import ..visualization.plot_3d_global as plot_3d
from .motion_process import recover_from_ric
def tensorborad_add_video_xyz(writer, xyz, nb_iter, tag, nb_vis=4, title_batch=None, outname=None):
xyz = xyz[:1]
bs, seq = xyz.shape[:2]
xyz = xyz.reshape(bs, seq, -1, 3)
plot_xyz = plot_3d.draw_to_batch(xyz.cpu().numpy(),title_batch, outname)
plot_xyz =np.transpose(plot_xyz, (0, 1, 4, 2, 3))
writer.add_video(tag, plot_xyz, nb_iter, fps = 20)
@torch.no_grad()
def evaluation_vqvae(out_dir, val_loader, net, logger, writer, nb_iter, best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, eval_wrapper, draw = True, save = True, savegif=False, savenpy=False) :
net.eval()
nb_sample = 0
draw_org = []
draw_pred = []
draw_text = []
motion_annotation_list = []
motion_pred_list = []
R_precision_real = 0
R_precision = 0
nb_sample = 0
matching_score_real = 0
matching_score_pred = 0
for batch in val_loader:
word_embeddings, pos_one_hots, caption, sent_len, motion, m_length, token, name = batch
motion = motion.cuda()
et, em = eval_wrapper.get_co_embeddings(word_embeddings, pos_one_hots, sent_len, motion, m_length)
bs, seq = motion.shape[0], motion.shape[1]
num_joints = 21 if motion.shape[-1] == 251 else 22
pred_pose_eval = torch.zeros((bs, seq, motion.shape[-1])).cuda()
for i in range(bs):
pose = val_loader.dataset.inv_transform(motion[i:i+1, :m_length[i], :].detach().cpu().numpy())
pose_xyz = recover_from_ric(torch.from_numpy(pose).float().cuda(), num_joints)
pred_pose, loss_commit, perplexity = net(motion[i:i+1, :m_length[i]])
pred_denorm = val_loader.dataset.inv_transform(pred_pose.detach().cpu().numpy())
pred_xyz = recover_from_ric(torch.from_numpy(pred_denorm).float().cuda(), num_joints)
if savenpy:
np.save(os.path.join(out_dir, name[i]+'_gt.npy'), pose_xyz[:, :m_length[i]].cpu().numpy())
np.save(os.path.join(out_dir, name[i]+'_pred.npy'), pred_xyz.detach().cpu().numpy())
pred_pose_eval[i:i+1,:m_length[i],:] = pred_pose
if i < min(4, bs):
draw_org.append(pose_xyz)
draw_pred.append(pred_xyz)
draw_text.append(caption[i])
et_pred, em_pred = eval_wrapper.get_co_embeddings(word_embeddings, pos_one_hots, sent_len, pred_pose_eval, m_length)
motion_pred_list.append(em_pred)
motion_annotation_list.append(em)
temp_R, temp_match = calculate_R_precision(et.cpu().numpy(), em.cpu().numpy(), top_k=3, sum_all=True)
R_precision_real += temp_R
matching_score_real += temp_match
temp_R, temp_match = calculate_R_precision(et_pred.cpu().numpy(), em_pred.cpu().numpy(), top_k=3, sum_all=True)
R_precision += temp_R
matching_score_pred += temp_match
nb_sample += bs
motion_annotation_np = torch.cat(motion_annotation_list, dim=0).cpu().numpy()
motion_pred_np = torch.cat(motion_pred_list, dim=0).cpu().numpy()
gt_mu, gt_cov = calculate_activation_statistics(motion_annotation_np)
mu, cov= calculate_activation_statistics(motion_pred_np)
diversity_real = calculate_diversity(motion_annotation_np, 300 if nb_sample > 300 else 100)
diversity = calculate_diversity(motion_pred_np, 300 if nb_sample > 300 else 100)
R_precision_real = R_precision_real / nb_sample
R_precision = R_precision / nb_sample
matching_score_real = matching_score_real / nb_sample
matching_score_pred = matching_score_pred / nb_sample
fid = calculate_frechet_distance(gt_mu, gt_cov, mu, cov)
msg = f"--> \t Eva. Iter {nb_iter} :, FID. {fid:.4f}, Diversity Real. {diversity_real:.4f}, Diversity. {diversity:.4f}, R_precision_real. {R_precision_real}, R_precision. {R_precision}, matching_score_real. {matching_score_real}, matching_score_pred. {matching_score_pred}"
logger.info(msg)
if draw:
writer.add_scalar('./Test/FID', fid, nb_iter)
writer.add_scalar('./Test/Diversity', diversity, nb_iter)
writer.add_scalar('./Test/top1', R_precision[0], nb_iter)
writer.add_scalar('./Test/top2', R_precision[1], nb_iter)
writer.add_scalar('./Test/top3', R_precision[2], nb_iter)
writer.add_scalar('./Test/matching_score', matching_score_pred, nb_iter)
if nb_iter % 5000 == 0 :
for ii in range(4):
tensorborad_add_video_xyz(writer, draw_org[ii], nb_iter, tag='./Vis/org_eval'+str(ii), nb_vis=1, title_batch=[draw_text[ii]], outname=[os.path.join(out_dir, 'gt'+str(ii)+'.gif')] if savegif else None)
if nb_iter % 5000 == 0 :
for ii in range(4):
tensorborad_add_video_xyz(writer, draw_pred[ii], nb_iter, tag='./Vis/pred_eval'+str(ii), nb_vis=1, title_batch=[draw_text[ii]], outname=[os.path.join(out_dir, 'pred'+str(ii)+'.gif')] if savegif else None)
if fid < best_fid :
msg = f"--> --> \t FID Improved from {best_fid:.5f} to {fid:.5f} !!!"
logger.info(msg)
best_fid, best_iter = fid, nb_iter
if save:
torch.save({'net' : net.state_dict()}, os.path.join(out_dir, 'net_best_fid.pth'))
if abs(diversity_real - diversity) < abs(diversity_real - best_div) :
msg = f"--> --> \t Diversity Improved from {best_div:.5f} to {diversity:.5f} !!!"
logger.info(msg)
best_div = diversity
if save:
torch.save({'net' : net.state_dict()}, os.path.join(out_dir, 'net_best_div.pth'))
if R_precision[0] > best_top1 :
msg = f"--> --> \t Top1 Improved from {best_top1:.4f} to {R_precision[0]:.4f} !!!"
logger.info(msg)
best_top1 = R_precision[0]
if save:
torch.save({'net' : net.state_dict()}, os.path.join(out_dir, 'net_best_top1.pth'))
if R_precision[1] > best_top2 :
msg = f"--> --> \t Top2 Improved from {best_top2:.4f} to {R_precision[1]:.4f} !!!"
logger.info(msg)
best_top2 = R_precision[1]
if R_precision[2] > best_top3 :
msg = f"--> --> \t Top3 Improved from {best_top3:.4f} to {R_precision[2]:.4f} !!!"
logger.info(msg)
best_top3 = R_precision[2]
if matching_score_pred < best_matching :
msg = f"--> --> \t matching_score Improved from {best_matching:.5f} to {matching_score_pred:.5f} !!!"
logger.info(msg)
best_matching = matching_score_pred
if save:
torch.save({'net' : net.state_dict()}, os.path.join(out_dir, 'net_best_matching.pth'))
if save:
torch.save({'net' : net.state_dict()}, os.path.join(out_dir, 'net_last.pth'))
net.train()
return best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, writer, logger
@torch.no_grad()
def evaluation_transformer(out_dir, val_loader, net, trans, logger, writer, nb_iter, best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, clip_model, eval_wrapper, draw = True, save = True, savegif=False) :
trans.eval()
nb_sample = 0
draw_org = []
draw_pred = []
draw_text = []
draw_text_pred = []
motion_annotation_list = []
motion_pred_list = []
R_precision_real = 0
R_precision = 0
matching_score_real = 0
matching_score_pred = 0
nb_sample = 0
for i in range(1):
for batch in val_loader:
word_embeddings, pos_one_hots, clip_text, sent_len, pose, m_length, token, name = batch
bs, seq = pose.shape[:2]
num_joints = 21 if pose.shape[-1] == 251 else 22
text = clip.tokenize(clip_text, truncate=True).cuda()
feat_clip_text = clip_model.encode_text(text).float()
pred_pose_eval = torch.zeros((bs, seq, pose.shape[-1])).cuda()
pred_len = torch.ones(bs).long()
for k in range(bs):
try:
index_motion = trans.sample(feat_clip_text[k:k+1], False)
except:
index_motion = torch.ones(1,1).cuda().long()
pred_pose = net.forward_decoder(index_motion)
cur_len = pred_pose.shape[1]
pred_len[k] = min(cur_len, seq)
pred_pose_eval[k:k+1, :cur_len] = pred_pose[:, :seq]
if draw:
pred_denorm = val_loader.dataset.inv_transform(pred_pose.detach().cpu().numpy())
pred_xyz = recover_from_ric(torch.from_numpy(pred_denorm).float().cuda(), num_joints)
if i == 0 and k < 4:
draw_pred.append(pred_xyz)
draw_text_pred.append(clip_text[k])
et_pred, em_pred = eval_wrapper.get_co_embeddings(word_embeddings, pos_one_hots, sent_len, pred_pose_eval, pred_len)
if i == 0:
pose = pose.cuda().float()
et, em = eval_wrapper.get_co_embeddings(word_embeddings, pos_one_hots, sent_len, pose, m_length)
motion_annotation_list.append(em)
motion_pred_list.append(em_pred)
if draw:
pose = val_loader.dataset.inv_transform(pose.detach().cpu().numpy())
pose_xyz = recover_from_ric(torch.from_numpy(pose).float().cuda(), num_joints)
for j in range(min(4, bs)):
draw_org.append(pose_xyz[j][:m_length[j]].unsqueeze(0))
draw_text.append(clip_text[j])
temp_R, temp_match = calculate_R_precision(et.cpu().numpy(), em.cpu().numpy(), top_k=3, sum_all=True)
R_precision_real += temp_R
matching_score_real += temp_match
temp_R, temp_match = calculate_R_precision(et_pred.cpu().numpy(), em_pred.cpu().numpy(), top_k=3, sum_all=True)
R_precision += temp_R
matching_score_pred += temp_match
nb_sample += bs
motion_annotation_np = torch.cat(motion_annotation_list, dim=0).cpu().numpy()
motion_pred_np = torch.cat(motion_pred_list, dim=0).cpu().numpy()
gt_mu, gt_cov = calculate_activation_statistics(motion_annotation_np)
mu, cov= calculate_activation_statistics(motion_pred_np)
diversity_real = calculate_diversity(motion_annotation_np, 300 if nb_sample > 300 else 100)
diversity = calculate_diversity(motion_pred_np, 300 if nb_sample > 300 else 100)
R_precision_real = R_precision_real / nb_sample
R_precision = R_precision / nb_sample
matching_score_real = matching_score_real / nb_sample
matching_score_pred = matching_score_pred / nb_sample
fid = calculate_frechet_distance(gt_mu, gt_cov, mu, cov)
msg = f"--> \t Eva. Iter {nb_iter} :, FID. {fid:.4f}, Diversity Real. {diversity_real:.4f}, Diversity. {diversity:.4f}, R_precision_real. {R_precision_real}, R_precision. {R_precision}, matching_score_real. {matching_score_real}, matching_score_pred. {matching_score_pred}"
logger.info(msg)
if draw:
writer.add_scalar('./Test/FID', fid, nb_iter)
writer.add_scalar('./Test/Diversity', diversity, nb_iter)
writer.add_scalar('./Test/top1', R_precision[0], nb_iter)
writer.add_scalar('./Test/top2', R_precision[1], nb_iter)
writer.add_scalar('./Test/top3', R_precision[2], nb_iter)
writer.add_scalar('./Test/matching_score', matching_score_pred, nb_iter)
if nb_iter % 10000 == 0 :
for ii in range(4):
tensorborad_add_video_xyz(writer, draw_org[ii], nb_iter, tag='./Vis/org_eval'+str(ii), nb_vis=1, title_batch=[draw_text[ii]], outname=[os.path.join(out_dir, 'gt'+str(ii)+'.gif')] if savegif else None)
if nb_iter % 10000 == 0 :
for ii in range(4):
tensorborad_add_video_xyz(writer, draw_pred[ii], nb_iter, tag='./Vis/pred_eval'+str(ii), nb_vis=1, title_batch=[draw_text_pred[ii]], outname=[os.path.join(out_dir, 'pred'+str(ii)+'.gif')] if savegif else None)
if fid < best_fid :
msg = f"--> --> \t FID Improved from {best_fid:.5f} to {fid:.5f} !!!"
logger.info(msg)
best_fid, best_iter = fid, nb_iter
if save:
torch.save({'trans' : trans.state_dict()}, os.path.join(out_dir, 'net_best_fid.pth'))
if matching_score_pred < best_matching :
msg = f"--> --> \t matching_score Improved from {best_matching:.5f} to {matching_score_pred:.5f} !!!"
logger.info(msg)
best_matching = matching_score_pred
if abs(diversity_real - diversity) < abs(diversity_real - best_div) :
msg = f"--> --> \t Diversity Improved from {best_div:.5f} to {diversity:.5f} !!!"
logger.info(msg)
best_div = diversity
if R_precision[0] > best_top1 :
msg = f"--> --> \t Top1 Improved from {best_top1:.4f} to {R_precision[0]:.4f} !!!"
logger.info(msg)
best_top1 = R_precision[0]
if R_precision[1] > best_top2 :
msg = f"--> --> \t Top2 Improved from {best_top2:.4f} to {R_precision[1]:.4f} !!!"
logger.info(msg)
best_top2 = R_precision[1]
if R_precision[2] > best_top3 :
msg = f"--> --> \t Top3 Improved from {best_top3:.4f} to {R_precision[2]:.4f} !!!"
logger.info(msg)
best_top3 = R_precision[2]
if save:
torch.save({'trans' : trans.state_dict()}, os.path.join(out_dir, 'net_last.pth'))
trans.train()
return best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, writer, logger
@torch.no_grad()
def evaluation_transformer_test(out_dir, val_loader, net, trans, logger, writer, nb_iter, best_fid, best_iter, best_div, best_top1, best_top2, best_top3, best_matching, best_multi, clip_model, eval_wrapper, draw = True, save = True, savegif=False, savenpy=False) :
trans.eval()
nb_sample = 0
draw_org = []
draw_pred = []
draw_text = []
draw_text_pred = []
draw_name = []
motion_annotation_list = []
motion_pred_list = []
motion_multimodality = []
R_precision_real = 0
R_precision = 0
matching_score_real = 0
matching_score_pred = 0
nb_sample = 0
for batch in val_loader:
word_embeddings, pos_one_hots, clip_text, sent_len, pose, m_length, token, name = batch
bs, seq = pose.shape[:2]
num_joints = 21 if pose.shape[-1] == 251 else 22
text = clip.tokenize(clip_text, truncate=True).cuda()
feat_clip_text = clip_model.encode_text(text).float()
motion_multimodality_batch = []
for i in range(30):
pred_pose_eval = torch.zeros((bs, seq, pose.shape[-1])).cuda()
pred_len = torch.ones(bs).long()
for k in range(bs):
try:
index_motion = trans.sample(feat_clip_text[k:k+1], True)
except:
index_motion = torch.ones(1,1).cuda().long()
pred_pose = net.forward_decoder(index_motion)
cur_len = pred_pose.shape[1]
pred_len[k] = min(cur_len, seq)
pred_pose_eval[k:k+1, :cur_len] = pred_pose[:, :seq]
if i == 0 and (draw or savenpy):
pred_denorm = val_loader.dataset.inv_transform(pred_pose.detach().cpu().numpy())
pred_xyz = recover_from_ric(torch.from_numpy(pred_denorm).float().cuda(), num_joints)
if savenpy:
np.save(os.path.join(out_dir, name[k]+'_pred.npy'), pred_xyz.detach().cpu().numpy())
if draw:
if i == 0:
draw_pred.append(pred_xyz)
draw_text_pred.append(clip_text[k])
draw_name.append(name[k])
et_pred, em_pred = eval_wrapper.get_co_embeddings(word_embeddings, pos_one_hots, sent_len, pred_pose_eval, pred_len)
motion_multimodality_batch.append(em_pred.reshape(bs, 1, -1))
if i == 0:
pose = pose.cuda().float()
et, em = eval_wrapper.get_co_embeddings(word_embeddings, pos_one_hots, sent_len, pose, m_length)
motion_annotation_list.append(em)
motion_pred_list.append(em_pred)
if draw or savenpy:
pose = val_loader.dataset.inv_transform(pose.detach().cpu().numpy())
pose_xyz = recover_from_ric(torch.from_numpy(pose).float().cuda(), num_joints)
if savenpy:
for j in range(bs):
np.save(os.path.join(out_dir, name[j]+'_gt.npy'), pose_xyz[j][:m_length[j]].unsqueeze(0).cpu().numpy())
if draw:
for j in range(bs):
draw_org.append(pose_xyz[j][:m_length[j]].unsqueeze(0))
draw_text.append(clip_text[j])
temp_R, temp_match = calculate_R_precision(et.cpu().numpy(), em.cpu().numpy(), top_k=3, sum_all=True)
R_precision_real += temp_R
matching_score_real += temp_match
temp_R, temp_match = calculate_R_precision(et_pred.cpu().numpy(), em_pred.cpu().numpy(), top_k=3, sum_all=True)
R_precision += temp_R
matching_score_pred += temp_match
nb_sample += bs
motion_multimodality.append(torch.cat(motion_multimodality_batch, dim=1))
motion_annotation_np = torch.cat(motion_annotation_list, dim=0).cpu().numpy()
motion_pred_np = torch.cat(motion_pred_list, dim=0).cpu().numpy()
gt_mu, gt_cov = calculate_activation_statistics(motion_annotation_np)
mu, cov= calculate_activation_statistics(motion_pred_np)
diversity_real = calculate_diversity(motion_annotation_np, 300 if nb_sample > 300 else 100)
diversity = calculate_diversity(motion_pred_np, 300 if nb_sample > 300 else 100)
R_precision_real = R_precision_real / nb_sample
R_precision = R_precision / nb_sample
matching_score_real = matching_score_real / nb_sample
matching_score_pred = matching_score_pred / nb_sample
multimodality = 0
motion_multimodality = torch.cat(motion_multimodality, dim=0).cpu().numpy()
multimodality = calculate_multimodality(motion_multimodality, 10)
fid = calculate_frechet_distance(gt_mu, gt_cov, mu, cov)
msg = f"--> \t Eva. Iter {nb_iter} :, FID. {fid:.4f}, Diversity Real. {diversity_real:.4f}, Diversity. {diversity:.4f}, R_precision_real. {R_precision_real}, R_precision. {R_precision}, matching_score_real. {matching_score_real}, matching_score_pred. {matching_score_pred}, multimodality. {multimodality:.4f}"
logger.info(msg)
if draw:
for ii in range(len(draw_org)):
tensorborad_add_video_xyz(writer, draw_org[ii], nb_iter, tag='./Vis/'+draw_name[ii]+'_org', nb_vis=1, title_batch=[draw_text[ii]], outname=[os.path.join(out_dir, draw_name[ii]+'_skel_gt.gif')] if savegif else None)
tensorborad_add_video_xyz(writer, draw_pred[ii], nb_iter, tag='./Vis/'+draw_name[ii]+'_pred', nb_vis=1, title_batch=[draw_text_pred[ii]], outname=[os.path.join(out_dir, draw_name[ii]+'_skel_pred.gif')] if savegif else None)
trans.train()
return fid, best_iter, diversity, R_precision[0], R_precision[1], R_precision[2], matching_score_pred, multimodality, writer, logger
# (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_R_precision(embedding1, embedding2, top_k, sum_all=False):
dist_mat = euclidean_distance_matrix(embedding1, embedding2)
matching_score = dist_mat.trace()
argmax = np.argsort(dist_mat, axis=1)
top_k_mat = calculate_top_k(argmax, top_k)
if sum_all:
return top_k_mat.sum(axis=0), matching_score
else:
return top_k_mat, matching_score
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()
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_frechet_distance(mu1, sigma1, mu2, sigma2, eps=1e-6):
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_activation_statistics(activations):
mu = np.mean(activations, axis=0)
cov = np.cov(activations, rowvar=False)
return mu, cov
def calculate_frechet_feature_distance(feature_list1, feature_list2):
feature_list1 = np.stack(feature_list1)
feature_list2 = np.stack(feature_list2)
# normalize the scale
mean = np.mean(feature_list1, axis=0)
std = np.std(feature_list1, axis=0) + 1e-10
feature_list1 = (feature_list1 - mean) / std
feature_list2 = (feature_list2 - mean) / std
dist = calculate_frechet_distance(
mu1=np.mean(feature_list1, axis=0),
sigma1=np.cov(feature_list1, rowvar=False),
mu2=np.mean(feature_list2, axis=0),
sigma2=np.cov(feature_list2, rowvar=False),
)
return dist
+30
View File
@@ -0,0 +1,30 @@
import torch
import torch.nn as nn
class ReConsLoss(nn.Module):
def __init__(self, recons_loss, nb_joints):
super(ReConsLoss, self).__init__()
if recons_loss == 'l1':
self.Loss = torch.nn.L1Loss()
elif recons_loss == 'l2' :
self.Loss = torch.nn.MSELoss()
elif recons_loss == 'l1_smooth' :
self.Loss = torch.nn.SmoothL1Loss()
# 4 global motion associated to root
# 12 local motion (3 local xyz, 3 vel xyz, 6 rot6d)
# 3 global vel xyz
# 4 foot contact
self.nb_joints = nb_joints
self.motion_dim = (nb_joints - 1) * 12 + 4 + 3 + 4
def forward(self, motion_pred, motion_gt) :
loss = self.Loss(motion_pred[..., : self.motion_dim], motion_gt[..., :self.motion_dim])
return loss
def forward_vel(self, motion_pred, motion_gt) :
loss = self.Loss(motion_pred[..., 4 : (self.nb_joints - 1) * 3 + 4], motion_gt[..., 4 : (self.nb_joints - 1) * 3 + 4])
return loss
+59
View File
@@ -0,0 +1,59 @@
import torch
from .quaternion import quaternion_to_cont6d, qrot, qinv
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_rot(data, joints_num, skeleton):
r_rot_quat, r_pos = recover_root_rot_pos(data)
r_rot_cont6d = quaternion_to_cont6d(r_rot_quat)
start_indx = 1 + 2 + 1 + (joints_num - 1) * 3
end_indx = start_indx + (joints_num - 1) * 6
cont6d_params = data[..., start_indx:end_indx]
# print(r_rot_cont6d.shape, cont6d_params.shape, r_pos.shape)
cont6d_params = torch.cat([r_rot_cont6d, cont6d_params], dim=-1)
cont6d_params = cont6d_params.view(-1, joints_num, 6)
positions = skeleton.forward_kinematics_cont6d(cont6d_params, r_pos)
return positions
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
+63
View File
@@ -0,0 +1,63 @@
import numpy as np
# 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]]
kit_tgt_skel_id = '03950'
t2m_tgt_skel_id = '000021'
+423
View File
@@ -0,0 +1,423 @@
# Copyright (c) 2018-present, Facebook, Inc.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
import torch
import numpy as np
_EPS4 = np.finfo(float).eps * 4.0
_FLOAT_EPS = np.finfo(np.float).eps
# PyTorch-backed implementations
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 qinv_np(q):
assert q.shape[-1] == 4, 'q must be a tensor of shape (*, 4)'
return qinv(torch.from_numpy(q).float()).numpy()
def qnormalize(q):
assert q.shape[-1] == 4, 'q must be a tensor of shape (*, 4)'
return q / torch.norm(q, dim=-1, keepdim=True)
def qmul(q, r):
"""
Multiply quaternion(s) q with quaternion(s) r.
Expects two equally-sized tensors of shape (*, 4), where * denotes any number of dimensions.
Returns q*r as a tensor of shape (*, 4).
"""
assert q.shape[-1] == 4
assert r.shape[-1] == 4
original_shape = q.shape
# Compute outer product
terms = torch.bmm(r.view(-1, 4, 1), q.view(-1, 1, 4))
w = terms[:, 0, 0] - terms[:, 1, 1] - terms[:, 2, 2] - terms[:, 3, 3]
x = terms[:, 0, 1] + terms[:, 1, 0] - terms[:, 2, 3] + terms[:, 3, 2]
y = terms[:, 0, 2] + terms[:, 1, 3] + terms[:, 2, 0] - terms[:, 3, 1]
z = terms[:, 0, 3] - terms[:, 1, 2] + terms[:, 2, 1] + terms[:, 3, 0]
return torch.stack((w, x, y, z), dim=1).view(original_shape)
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 qeuler(q, order, epsilon=0, deg=True):
"""
Convert quaternion(s) q to Euler angles.
Expects a tensor of shape (*, 4), where * denotes any number of dimensions.
Returns a tensor of shape (*, 3).
"""
assert q.shape[-1] == 4
original_shape = list(q.shape)
original_shape[-1] = 3
q = q.view(-1, 4)
q0 = q[:, 0]
q1 = q[:, 1]
q2 = q[:, 2]
q3 = q[:, 3]
if order == 'xyz':
x = torch.atan2(2 * (q0 * q1 - q2 * q3), 1 - 2 * (q1 * q1 + q2 * q2))
y = torch.asin(torch.clamp(2 * (q1 * q3 + q0 * q2), -1 + epsilon, 1 - epsilon))
z = torch.atan2(2 * (q0 * q3 - q1 * q2), 1 - 2 * (q2 * q2 + q3 * q3))
elif order == 'yzx':
x = torch.atan2(2 * (q0 * q1 - q2 * q3), 1 - 2 * (q1 * q1 + q3 * q3))
y = torch.atan2(2 * (q0 * q2 - q1 * q3), 1 - 2 * (q2 * q2 + q3 * q3))
z = torch.asin(torch.clamp(2 * (q1 * q2 + q0 * q3), -1 + epsilon, 1 - epsilon))
elif order == 'zxy':
x = torch.asin(torch.clamp(2 * (q0 * q1 + q2 * q3), -1 + epsilon, 1 - epsilon))
y = torch.atan2(2 * (q0 * q2 - q1 * q3), 1 - 2 * (q1 * q1 + q2 * q2))
z = torch.atan2(2 * (q0 * q3 - q1 * q2), 1 - 2 * (q1 * q1 + q3 * q3))
elif order == 'xzy':
x = torch.atan2(2 * (q0 * q1 + q2 * q3), 1 - 2 * (q1 * q1 + q3 * q3))
y = torch.atan2(2 * (q0 * q2 + q1 * q3), 1 - 2 * (q2 * q2 + q3 * q3))
z = torch.asin(torch.clamp(2 * (q0 * q3 - q1 * q2), -1 + epsilon, 1 - epsilon))
elif order == 'yxz':
x = torch.asin(torch.clamp(2 * (q0 * q1 - q2 * q3), -1 + epsilon, 1 - epsilon))
y = torch.atan2(2 * (q1 * q3 + q0 * q2), 1 - 2 * (q1 * q1 + q2 * q2))
z = torch.atan2(2 * (q1 * q2 + q0 * q3), 1 - 2 * (q1 * q1 + q3 * q3))
elif order == 'zyx':
x = torch.atan2(2 * (q0 * q1 + q2 * q3), 1 - 2 * (q1 * q1 + q2 * q2))
y = torch.asin(torch.clamp(2 * (q0 * q2 - q1 * q3), -1 + epsilon, 1 - epsilon))
z = torch.atan2(2 * (q0 * q3 + q1 * q2), 1 - 2 * (q2 * q2 + q3 * q3))
else:
raise
if deg:
return torch.stack((x, y, z), dim=1).view(original_shape) * 180 / np.pi
else:
return torch.stack((x, y, z), dim=1).view(original_shape)
# Numpy-backed implementations
def qmul_np(q, r):
q = torch.from_numpy(q).contiguous().float()
r = torch.from_numpy(r).contiguous().float()
return qmul(q, r).numpy()
def qrot_np(q, v):
q = torch.from_numpy(q).contiguous().float()
v = torch.from_numpy(v).contiguous().float()
return qrot(q, v).numpy()
def qeuler_np(q, order, epsilon=0, use_gpu=False):
if use_gpu:
q = torch.from_numpy(q).cuda().float()
return qeuler(q, order, epsilon).cpu().numpy()
else:
q = torch.from_numpy(q).contiguous().float()
return qeuler(q, order, epsilon).numpy()
def qfix(q):
"""
Enforce quaternion continuity across the time dimension by selecting
the representation (q or -q) with minimal distance (or, equivalently, maximal dot product)
between two consecutive frames.
Expects a tensor of shape (L, J, 4), where L is the sequence length and J is the number of joints.
Returns a tensor of the same shape.
"""
assert len(q.shape) == 3
assert q.shape[-1] == 4
result = q.copy()
dot_products = np.sum(q[1:] * q[:-1], axis=2)
mask = dot_products < 0
mask = (np.cumsum(mask, axis=0) % 2).astype(bool)
result[1:][mask] *= -1
return result
def euler2quat(e, order, deg=True):
"""
Convert Euler angles to quaternions.
"""
assert e.shape[-1] == 3
original_shape = list(e.shape)
original_shape[-1] = 4
e = e.view(-1, 3)
## if euler angles in degrees
if deg:
e = e * np.pi / 180.
x = e[:, 0]
y = e[:, 1]
z = e[:, 2]
rx = torch.stack((torch.cos(x / 2), torch.sin(x / 2), torch.zeros_like(x), torch.zeros_like(x)), dim=1)
ry = torch.stack((torch.cos(y / 2), torch.zeros_like(y), torch.sin(y / 2), torch.zeros_like(y)), dim=1)
rz = torch.stack((torch.cos(z / 2), torch.zeros_like(z), torch.zeros_like(z), torch.sin(z / 2)), dim=1)
result = None
for coord in order:
if coord == 'x':
r = rx
elif coord == 'y':
r = ry
elif coord == 'z':
r = rz
else:
raise
if result is None:
result = r
else:
result = qmul(result, r)
# Reverse antipodal representation to have a non-negative "w"
if order in ['xyz', 'yzx', 'zxy']:
result *= -1
return result.view(original_shape)
def expmap_to_quaternion(e):
"""
Convert axis-angle rotations (aka exponential maps) to quaternions.
Stable formula from "Practical Parameterization of Rotations Using the Exponential Map".
Expects a tensor of shape (*, 3), where * denotes any number of dimensions.
Returns a tensor of shape (*, 4).
"""
assert e.shape[-1] == 3
original_shape = list(e.shape)
original_shape[-1] = 4
e = e.reshape(-1, 3)
theta = np.linalg.norm(e, axis=1).reshape(-1, 1)
w = np.cos(0.5 * theta).reshape(-1, 1)
xyz = 0.5 * np.sinc(0.5 * theta / np.pi) * e
return np.concatenate((w, xyz), axis=1).reshape(original_shape)
def euler_to_quaternion(e, order):
"""
Convert Euler angles to quaternions.
"""
assert e.shape[-1] == 3
original_shape = list(e.shape)
original_shape[-1] = 4
e = e.reshape(-1, 3)
x = e[:, 0]
y = e[:, 1]
z = e[:, 2]
rx = np.stack((np.cos(x / 2), np.sin(x / 2), np.zeros_like(x), np.zeros_like(x)), axis=1)
ry = np.stack((np.cos(y / 2), np.zeros_like(y), np.sin(y / 2), np.zeros_like(y)), axis=1)
rz = np.stack((np.cos(z / 2), np.zeros_like(z), np.zeros_like(z), np.sin(z / 2)), axis=1)
result = None
for coord in order:
if coord == 'x':
r = rx
elif coord == 'y':
r = ry
elif coord == 'z':
r = rz
else:
raise
if result is None:
result = r
else:
result = qmul_np(result, r)
# Reverse antipodal representation to have a non-negative "w"
if order in ['xyz', 'yzx', 'zxy']:
result *= -1
return result.reshape(original_shape)
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 quaternion_to_matrix_np(quaternions):
q = torch.from_numpy(quaternions).contiguous().float()
return quaternion_to_matrix(q).numpy()
def quaternion_to_cont6d_np(quaternions):
rotation_mat = quaternion_to_matrix_np(quaternions)
cont_6d = np.concatenate([rotation_mat[..., 0], rotation_mat[..., 1]], axis=-1)
return cont_6d
def quaternion_to_cont6d(quaternions):
rotation_mat = quaternion_to_matrix(quaternions)
cont_6d = torch.cat([rotation_mat[..., 0], rotation_mat[..., 1]], dim=-1)
return cont_6d
def cont6d_to_matrix(cont6d):
assert cont6d.shape[-1] == 6, "The last dimension must be 6"
x_raw = cont6d[..., 0:3]
y_raw = cont6d[..., 3:6]
x = x_raw / torch.norm(x_raw, dim=-1, keepdim=True)
z = torch.cross(x, y_raw, dim=-1)
z = z / torch.norm(z, dim=-1, keepdim=True)
y = torch.cross(z, x, dim=-1)
x = x[..., None]
y = y[..., None]
z = z[..., None]
mat = torch.cat([x, y, z], dim=-1)
return mat
def cont6d_to_matrix_np(cont6d):
q = torch.from_numpy(cont6d).contiguous().float()
return cont6d_to_matrix(q).numpy()
def qpow(q0, t, dtype=torch.float):
''' q0 : tensor of quaternions
t: tensor of powers
'''
q0 = qnormalize(q0)
theta0 = torch.acos(q0[..., 0])
## if theta0 is close to zero, add epsilon to avoid NaNs
mask = (theta0 <= 10e-10) * (theta0 >= -10e-10)
theta0 = (1 - mask) * theta0 + mask * 10e-10
v0 = q0[..., 1:] / torch.sin(theta0).view(-1, 1)
if isinstance(t, torch.Tensor):
q = torch.zeros(t.shape + q0.shape)
theta = t.view(-1, 1) * theta0.view(1, -1)
else: ## if t is a number
q = torch.zeros(q0.shape)
theta = t * theta0
q[..., 0] = torch.cos(theta)
q[..., 1:] = v0 * torch.sin(theta).unsqueeze(-1)
return q.to(dtype)
def qslerp(q0, q1, t):
'''
q0: starting quaternion
q1: ending quaternion
t: array of points along the way
Returns:
Tensor of Slerps: t.shape + q0.shape
'''
q0 = qnormalize(q0)
q1 = qnormalize(q1)
q_ = qpow(qmul(q1, qinv(q0)), t)
return qmul(q_,
q0.contiguous().view(torch.Size([1] * len(t.shape)) + q0.shape).expand(t.shape + q0.shape).contiguous())
def qbetween(v0, v1):
'''
find the quaternion used to rotate v0 to v1
'''
assert v0.shape[-1] == 3, 'v0 must be of the shape (*, 3)'
assert v1.shape[-1] == 3, 'v1 must be of the shape (*, 3)'
v = torch.cross(v0, v1)
w = torch.sqrt((v0 ** 2).sum(dim=-1, keepdim=True) * (v1 ** 2).sum(dim=-1, keepdim=True)) + (v0 * v1).sum(dim=-1,
keepdim=True)
return qnormalize(torch.cat([w, v], dim=-1))
def qbetween_np(v0, v1):
'''
find the quaternion used to rotate v0 to v1
'''
assert v0.shape[-1] == 3, 'v0 must be of the shape (*, 3)'
assert v1.shape[-1] == 3, 'v1 must be of the shape (*, 3)'
v0 = torch.from_numpy(v0).float()
v1 = torch.from_numpy(v1).float()
return qbetween(v0, v1).numpy()
def lerp(p0, p1, t):
if not isinstance(t, torch.Tensor):
t = torch.Tensor([t])
new_shape = t.shape + p0.shape
new_view_t = t.shape + torch.Size([1] * len(p0.shape))
new_view_p = torch.Size([1] * len(t.shape)) + p0.shape
p0 = p0.view(new_view_p).expand(new_shape)
p1 = p1.view(new_view_p).expand(new_shape)
t = t.view(new_view_t).expand(new_shape)
return p0 + t * (p1 - p0)
+532
View File
@@ -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
+199
View File
@@ -0,0 +1,199 @@
from .quaternion import *
import scipy.ndimage.filters as filters
class Skeleton(object):
def __init__(self, offset, kinematic_tree, device):
self.device = device
self._raw_offset_np = offset.numpy()
self._raw_offset = offset.clone().detach().to(device).float()
self._kinematic_tree = kinematic_tree
self._offset = None
self._parents = [0] * len(self._raw_offset)
self._parents[0] = -1
for chain in self._kinematic_tree:
for j in range(1, len(chain)):
self._parents[chain[j]] = chain[j-1]
def njoints(self):
return len(self._raw_offset)
def offset(self):
return self._offset
def set_offset(self, offsets):
self._offset = offsets.clone().detach().to(self.device).float()
def kinematic_tree(self):
return self._kinematic_tree
def parents(self):
return self._parents
# joints (batch_size, joints_num, 3)
def get_offsets_joints_batch(self, joints):
assert len(joints.shape) == 3
_offsets = self._raw_offset.expand(joints.shape[0], -1, -1).clone()
for i in range(1, self._raw_offset.shape[0]):
_offsets[:, i] = torch.norm(joints[:, i] - joints[:, self._parents[i]], p=2, dim=1)[:, None] * _offsets[:, i]
self._offset = _offsets.detach()
return _offsets
# joints (joints_num, 3)
def get_offsets_joints(self, joints):
assert len(joints.shape) == 2
_offsets = self._raw_offset.clone()
for i in range(1, self._raw_offset.shape[0]):
# print(joints.shape)
_offsets[i] = torch.norm(joints[i] - joints[self._parents[i]], p=2, dim=0) * _offsets[i]
self._offset = _offsets.detach()
return _offsets
# face_joint_idx should follow the order of right hip, left hip, right shoulder, left shoulder
# joints (batch_size, joints_num, 3)
def inverse_kinematics_np(self, joints, face_joint_idx, smooth_forward=False):
assert len(face_joint_idx) == 4
'''Get Forward Direction'''
l_hip, r_hip, sdr_r, sdr_l = face_joint_idx
across1 = joints[:, r_hip] - joints[:, l_hip]
across2 = joints[:, sdr_r] - joints[:, sdr_l]
across = across1 + across2
across = across / np.sqrt((across**2).sum(axis=-1))[:, np.newaxis]
# print(across1.shape, across2.shape)
# forward (batch_size, 3)
forward = np.cross(np.array([[0, 1, 0]]), across, axis=-1)
if smooth_forward:
forward = filters.gaussian_filter1d(forward, 20, axis=0, mode='nearest')
# forward (batch_size, 3)
forward = forward / np.sqrt((forward**2).sum(axis=-1))[..., np.newaxis]
'''Get Root Rotation'''
target = np.array([[0,0,1]]).repeat(len(forward), axis=0)
root_quat = qbetween_np(forward, target)
'''Inverse Kinematics'''
# quat_params (batch_size, joints_num, 4)
# print(joints.shape[:-1])
quat_params = np.zeros(joints.shape[:-1] + (4,))
# print(quat_params.shape)
root_quat[0] = np.array([[1.0, 0.0, 0.0, 0.0]])
quat_params[:, 0] = root_quat
# quat_params[0, 0] = np.array([[1.0, 0.0, 0.0, 0.0]])
for chain in self._kinematic_tree:
R = root_quat
for j in range(len(chain) - 1):
# (batch, 3)
u = self._raw_offset_np[chain[j+1]][np.newaxis,...].repeat(len(joints), axis=0)
# print(u.shape)
# (batch, 3)
v = joints[:, chain[j+1]] - joints[:, chain[j]]
v = v / np.sqrt((v**2).sum(axis=-1))[:, np.newaxis]
# print(u.shape, v.shape)
rot_u_v = qbetween_np(u, v)
R_loc = qmul_np(qinv_np(R), rot_u_v)
quat_params[:,chain[j + 1], :] = R_loc
R = qmul_np(R, R_loc)
return quat_params
# Be sure root joint is at the beginning of kinematic chains
def forward_kinematics(self, quat_params, root_pos, skel_joints=None, do_root_R=True):
# quat_params (batch_size, joints_num, 4)
# joints (batch_size, joints_num, 3)
# root_pos (batch_size, 3)
if skel_joints is not None:
offsets = self.get_offsets_joints_batch(skel_joints)
if len(self._offset.shape) == 2:
offsets = self._offset.expand(quat_params.shape[0], -1, -1)
joints = torch.zeros(quat_params.shape[:-1] + (3,)).to(self.device)
joints[:, 0] = root_pos
for chain in self._kinematic_tree:
if do_root_R:
R = quat_params[:, 0]
else:
R = torch.tensor([[1.0, 0.0, 0.0, 0.0]]).expand(len(quat_params), -1).detach().to(self.device)
for i in range(1, len(chain)):
R = qmul(R, quat_params[:, chain[i]])
offset_vec = offsets[:, chain[i]]
joints[:, chain[i]] = qrot(R, offset_vec) + joints[:, chain[i-1]]
return joints
# Be sure root joint is at the beginning of kinematic chains
def forward_kinematics_np(self, quat_params, root_pos, skel_joints=None, do_root_R=True):
# quat_params (batch_size, joints_num, 4)
# joints (batch_size, joints_num, 3)
# root_pos (batch_size, 3)
if skel_joints is not None:
skel_joints = torch.from_numpy(skel_joints)
offsets = self.get_offsets_joints_batch(skel_joints)
if len(self._offset.shape) == 2:
offsets = self._offset.expand(quat_params.shape[0], -1, -1)
offsets = offsets.numpy()
joints = np.zeros(quat_params.shape[:-1] + (3,))
joints[:, 0] = root_pos
for chain in self._kinematic_tree:
if do_root_R:
R = quat_params[:, 0]
else:
R = np.array([[1.0, 0.0, 0.0, 0.0]]).repeat(len(quat_params), axis=0)
for i in range(1, len(chain)):
R = qmul_np(R, quat_params[:, chain[i]])
offset_vec = offsets[:, chain[i]]
joints[:, chain[i]] = qrot_np(R, offset_vec) + joints[:, chain[i - 1]]
return joints
def forward_kinematics_cont6d_np(self, cont6d_params, root_pos, skel_joints=None, do_root_R=True):
# cont6d_params (batch_size, joints_num, 6)
# joints (batch_size, joints_num, 3)
# root_pos (batch_size, 3)
if skel_joints is not None:
skel_joints = torch.from_numpy(skel_joints)
offsets = self.get_offsets_joints_batch(skel_joints)
if len(self._offset.shape) == 2:
offsets = self._offset.expand(cont6d_params.shape[0], -1, -1)
offsets = offsets.numpy()
joints = np.zeros(cont6d_params.shape[:-1] + (3,))
joints[:, 0] = root_pos
for chain in self._kinematic_tree:
if do_root_R:
matR = cont6d_to_matrix_np(cont6d_params[:, 0])
else:
matR = np.eye(3)[np.newaxis, :].repeat(len(cont6d_params), axis=0)
for i in range(1, len(chain)):
matR = np.matmul(matR, cont6d_to_matrix_np(cont6d_params[:, chain[i]]))
offset_vec = offsets[:, chain[i]][..., np.newaxis]
# print(matR.shape, offset_vec.shape)
joints[:, chain[i]] = np.matmul(matR, offset_vec).squeeze(-1) + joints[:, chain[i-1]]
return joints
def forward_kinematics_cont6d(self, cont6d_params, root_pos, skel_joints=None, do_root_R=True):
# cont6d_params (batch_size, joints_num, 6)
# joints (batch_size, joints_num, 3)
# root_pos (batch_size, 3)
if skel_joints is not None:
# skel_joints = torch.from_numpy(skel_joints)
offsets = self.get_offsets_joints_batch(skel_joints)
if len(self._offset.shape) == 2:
offsets = self._offset.expand(cont6d_params.shape[0], -1, -1)
joints = torch.zeros(cont6d_params.shape[:-1] + (3,)).to(cont6d_params.device)
joints[..., 0, :] = root_pos
for chain in self._kinematic_tree:
if do_root_R:
matR = cont6d_to_matrix(cont6d_params[:, 0])
else:
matR = torch.eye(3).expand((len(cont6d_params), -1, -1)).detach().to(cont6d_params.device)
for i in range(1, len(chain)):
matR = torch.matmul(matR, cont6d_to_matrix(cont6d_params[:, chain[i]]))
offset_vec = offsets[:, chain[i]].unsqueeze(-1)
# print(matR.shape, offset_vec.shape)
joints[:, chain[i]] = torch.matmul(matR, offset_vec).squeeze(-1) + joints[:, chain[i-1]]
return joints
+66
View File
@@ -0,0 +1,66 @@
import numpy as np
import torch
import torch.optim as optim
import logging
import os
import sys
def getCi(accLog):
mean = np.mean(accLog)
std = np.std(accLog)
ci95 = 1.96*std/np.sqrt(len(accLog))
return mean, ci95
def get_logger(out_dir):
logger = logging.getLogger('Exp')
logger.setLevel(logging.INFO)
formatter = logging.Formatter("%(asctime)s %(levelname)s %(message)s")
file_path = os.path.join(out_dir, "run.log")
file_hdlr = logging.FileHandler(file_path)
file_hdlr.setFormatter(formatter)
strm_hdlr = logging.StreamHandler(sys.stdout)
strm_hdlr.setFormatter(formatter)
logger.addHandler(file_hdlr)
logger.addHandler(strm_hdlr)
return logger
## Optimizer
def initial_optim(decay_option, lr, weight_decay, net, optimizer) :
if optimizer == 'adamw' :
optimizer_adam_family = optim.AdamW
elif optimizer == 'adam' :
optimizer_adam_family = optim.Adam
if decay_option == 'all':
#optimizer = optimizer_adam_family(net.parameters(), lr=lr, betas=(0.9, 0.999), weight_decay=weight_decay)
optimizer = optimizer_adam_family(net.parameters(), lr=lr, betas=(0.5, 0.9), weight_decay=weight_decay)
elif decay_option == 'noVQ':
all_params = set(net.parameters())
no_decay = set([net.vq_layer])
decay = all_params - no_decay
optimizer = optimizer_adam_family([
{'params': list(no_decay), 'weight_decay': 0},
{'params': list(decay), 'weight_decay' : weight_decay}], lr=lr)
return optimizer
def get_motion_with_trans(motion, velocity) :
'''
motion : torch.tensor, shape (batch_size, T, 72), with the global translation = 0
velocity : torch.tensor, shape (batch_size, T, 3), contain the information of velocity = 0
'''
trans = torch.cumsum(velocity, dim=1)
trans = trans - trans[:, :1] ## the first root is initialized at 0 (just for visualization)
trans = trans.repeat((1, 1, 21))
motion_with_trans = motion + trans
return motion_with_trans
+99
View File
@@ -0,0 +1,99 @@
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'))
self.word2idx = pickle.load(open(pjoin(meta_root, '%s_idx.pkl'%prefix), 'rb'))
self.word2vec = {w: vectors[self.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
class WordVectorizerV2(WordVectorizer):
def __init__(self, meta_root, prefix):
super(WordVectorizerV2, self).__init__(meta_root, prefix)
self.idx2word = {self.word2idx[w]: w for w in self.word2idx}
def __getitem__(self, item):
word_vec, pose_vec = super(WordVectorizerV2, self).__getitem__(item)
word, pos = item.split('/')
if word in self.word2vec:
return word_vec, pose_vec, self.word2idx[word]
else:
return word_vec, pose_vec, self.word2idx['unk']
def itos(self, idx):
if idx == len(self.idx2word):
return "pad"
return self.idx2word[idx]
@@ -0,0 +1,131 @@
import torch
import matplotlib.pyplot as plt
import numpy as np
import io
import matplotlib
from mpl_toolkits.mplot3d.art3d import Poly3DCollection
import mpl_toolkits.mplot3d.axes3d as p3
from textwrap import wrap
import imageio
def plot_3d_motion(args, figsize=(10, 10), fps=120, radius=4):
matplotlib.use('Agg')
plt.style.use('dark_background')
joints, out_name, title = args #kit(192,22,3)
data = joints.copy().reshape(len(joints), -1, 3)
nb_joints = joints.shape[1]# kit:22 openpose 25
smpl_kinetic_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]] if nb_joints == 21 else [[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]]
# 22关键点 [0, 2, 5, 8, 11]表示连接了五个关键点,分别是左脚踝(0号关键点)、左髋部(2号关键点)、左肩部(5号关键点)、左手腕(8号关键点)和左肘部(11号关键点),这五个关键点按照顺序连接起来。
# [0, 11, 12, 13, 14, 15]表示连接了六个关键点,分别是骨盆(0号关键点)、左大腿(11号关键点)、左小腿(12号关键点)、左脚踝(13号关键点)、左脚尖(14号关键点)和左脚掌(15号关键点
limits = 1000 if nb_joints == 21 else 2
MINS = data.min(axis=0).min(axis=0)
MAXS = data.max(axis=0).max(axis=0)
colors = ['red', 'blue', 'black', 'red', 'blue',
'darkblue', 'darkblue', 'darkblue', 'darkblue', 'darkblue',
'darkred', 'darkred', 'darkred', 'darkred', 'darkred']
frame_number = data.shape[0]
# print(data.shape)
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):
def init():
ax.set_xlim(-limits, limits)
ax.set_ylim(-limits, limits)
ax.set_zlim(0, limits)
ax.grid(b=False)
def plot_xzPlane(minx, maxx, miny, minz, maxz):
## Plot a plane XZ
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))
#xz_plane.set_facecolor(())
#ax.add_collection3d(xz_plane)#绘制行走平面
fig = plt.figure(figsize=(480/96., 320/96.), dpi=96) if nb_joints == 21 else plt.figure(figsize=(10, 10), dpi=96)
if title is not None :
wraped_title = '\n'.join(wrap(title, 40))
fig.suptitle(wraped_title, fontsize=16)
ax = p3.Axes3D(fig)
init()
#ax.lines = []
#ax.collections = []
ax.view_init(elev=110, azim=-90)
ax.dist = 7.5
# ax =
plot_xzPlane(MINS[0] - trajec[index, 0], MAXS[0] - trajec[index, 0], 0, MINS[2] - trajec[index, 1],
MAXS[2] - trajec[index, 1])
# ax.scatter(data[index, :22, 0], data[index, :22, 1], data[index, :22, 2], color='black', s=3)
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')
# ax = plot_xzPlane(ax, MINS[0], MAXS[0], 0, MINS[2], MAXS[2])
for i, (chain, color) in enumerate(zip(smpl_kinetic_chain, colors)):
if i < 5:
linewidth = 4.0
else:
linewidth = 2.0
ax.plot3D(data[index, chain, 0], data[index, chain, 1], data[index, chain, 2], linewidth=linewidth,color=color)#xyz width color
#ax.text(data[index, chain, 0], data[index, chain, 1], data[index, chain, 2], str(chain), fontsize = 15)
plt.axis('off')
ax.set_xticklabels([])
ax.set_yticklabels([])
ax.set_zticklabels([])
#plt.savefig(f'./smpl_{index}.jpg', dpi=96)
if out_name is not None :
plt.savefig(out_name, dpi=96)
plt.close()
else :
io_buf = io.BytesIO()
fig.savefig(io_buf, format='raw', dpi=96)
io_buf.seek(0)
# print(fig.bbox.bounds)
arr = np.reshape(np.frombuffer(io_buf.getvalue(), dtype=np.uint8),
newshape=(int(fig.bbox.bounds[3]), int(fig.bbox.bounds[2]), -1))
io_buf.close()
plt.close()
return arr
out = []
for i in range(frame_number) :
out.append(update(i))
out = np.stack(out, axis=0)
return torch.from_numpy(out)
def draw_to_batch(smpl_joints_batch, title_batch=None, outname=None) :
batch_size = len(smpl_joints_batch)
out = []
for i in range(batch_size) :
out.append(plot_3d_motion([smpl_joints_batch[i], None, title_batch[i] if title_batch is not None else None]))
if outname is not None:
imageio.mimsave(outname[i], np.array(out[-1]), duration=1000/20)
out = torch.stack(out, axis=0)
return out
@@ -0,0 +1,40 @@
import numpy as np
# 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 = "/group/30065/users/zhanchao/code/MMCM/mmcm/t2p/body_models/"
GMM_MODEL_DIR = "/group/30065/users/zhanchao/code/MMCM/mmcm/t2p/visualize/joints2smpl/smpl_models"
SMPL_MEAN_FILE = "/group/30065/users/zhanchao/code/MMCM/mmcm/t2p/visualize/joints2smpl/smpl_models/neutral_smpl_mean_params.h5"
# for collsion
Part_Seg_DIR = "/group/30065/users/zhanchao/code/MMCM/mmcm/t2p/visualize/joints2smpl/smpl_models/smplx_parts_segm.pkl"
@@ -0,0 +1,224 @@
from __future__ import absolute_import
import torch
import torch.nn.functional as F
from mmcm.t2p.visualize.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,230 @@
# -*- 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))
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)
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,281 @@
from __future__ import absolute_import
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 mmcm.t2p.visualize.joints2smpl.src import config
@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
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')
for i in range(10):
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()
return loss
camera_optimizer.step(closure)
else:
camera_optimizer = torch.optim.Adam(camera_opt_params, lr=self.step_size, betas=(0.9, 0.999))
for i in range(20):
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()
# Fix camera translation after optimizing camera
# --------Step 2: Optimize body joints --------------------------
# Optimize only the body pose and global orientation of the body
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')
for i in range(self.num_iters):
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()
return loss
body_optimizer.step(closure)
else:
body_optimizer = torch.optim.Adam(body_opt_params, lr=self.step_size, betas=(0.9, 0.999))
for i in range(self.num_iters):
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()
# 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
+33
View File
@@ -0,0 +1,33 @@
import argparse
import os
from .visualize import vis_utils
import shutil
from tqdm import tqdm
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument("--input_path", type=str, required=True, help='stick figure mp4 file to be rendered.')
parser.add_argument("--cuda", type=bool, default=True, help='')
parser.add_argument("--device", type=int, default=0, help='')
params = parser.parse_args()
assert params.input_path.endswith('.mp4')
parsed_name = os.path.basename(params.input_path).replace('.mp4', '').replace('sample', '').replace('rep', '')
sample_i, rep_i = [int(e) for e in parsed_name.split('_')]
npy_path = os.path.join(os.path.dirname(params.input_path), 'results.npy')
out_npy_path = params.input_path.replace('.mp4', '_smpl_params.npy')
assert os.path.exists(npy_path)
results_dir = params.input_path.replace('.mp4', '_obj')
if os.path.exists(results_dir):
shutil.rmtree(results_dir)
os.makedirs(results_dir)
npy2obj = vis_utils.npy2obj(npy_path, sample_i, rep_i,
device=params.device, cuda=params.cuda)
print('Saving obj files to [{}]'.format(os.path.abspath(results_dir)))
for frame_i in tqdm(range(npy2obj.real_num_frames)):
npy2obj.save_obj(os.path.join(results_dir, 'frame{:03d}.obj'.format(frame_i)), frame_i)
print('Saving SMPL params to [{}]'.format(os.path.abspath(out_npy_path)))
npy2obj.save_npy(out_npy_path)

Some files were not shown because too many files have changed in this diff Show More