Add 4DHuman

This commit is contained in:
Hacker 17082006
2024-03-31 15:46:19 +07:00
parent e4f50bbc63
commit cd340c210e
249 changed files with 29981 additions and 12 deletions
View File
+107
View File
@@ -0,0 +1,107 @@
# 4DHumans: Reconstructing and Tracking Humans with Transformers
Code repository for the paper:
**Humans in 4D: Reconstructing and Tracking Humans with Transformers**
[Shubham Goel](https://people.eecs.berkeley.edu/~shubham-goel/), [Georgios Pavlakos](https://geopavlakos.github.io/), [Jathushan Rajasegaran](http://people.eecs.berkeley.edu/~jathushan/), [Angjoo Kanazawa](https://people.eecs.berkeley.edu/~kanazawa/)<sup>\*</sup>, [Jitendra Malik](http://people.eecs.berkeley.edu/~malik/)<sup>\*</sup>
[![arXiv](https://img.shields.io/badge/arXiv-2305.20091-00ff00.svg)](https://arxiv.org/pdf/2305.20091.pdf) [![Website shields.io](https://img.shields.io/website-up-down-green-red/http/shields.io.svg)](https://shubham-goel.github.io/4dhumans/) [![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/drive/1Ex4gE5v1bPR3evfhtG7sDHxQGsWwNwby?usp=sharing) [![Hugging Face Spaces](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Spaces-blue)](https://huggingface.co/spaces/brjathu/HMR2.0)
![teaser](assets/teaser.png)
## Installation and Setup
First, clone the repo. Then, we recommend creating a clean [conda](https://docs.conda.io/) environment, installing all dependencies, and finally activating the environment, as follows:
```bash
git clone https://github.com/shubham-goel/4D-Humans.git
cd 4D-Humans
conda env create -f environment.yml
conda activate 4D-humans
```
If conda is too slow, you can use pip:
```bash
conda create --name 4D-humans python=3.10
conda activate 4D-humans
pip install torch
pip install -e .[all]
```
All checkpoints and data will automatically be downloaded to `$HOME/.cache/4DHumans` the first time you run the demo code.
Besides these files, you also need to download the *SMPL* model. You will need the [neutral model](http://smplify.is.tue.mpg.de) for training and running the demo code. Please go to the corresponding website and register to get access to the downloads section. Download the model and place `basicModel_neutral_lbs_10_207_0_v1.0.0.pkl` in `./data/`.
## Run demo on images
The following command will run ViTDet and HMR2.0 on all images in the specified `--img_folder`, and save renderings of the reconstructions in `--out_folder`. `--batch_size` batches the images together for faster processing. The `--side_view` flags additionally renders the side view of the reconstructed mesh, `--full_frame` renders all people together in front view, `--save_mesh` saves meshes as `.obj`s.
```bash
python demo.py \
--img_folder example_data/images \
--out_folder demo_out \
--batch_size=48 --side_view --save_mesh --full_frame
```
## Run tracking demo on videos
Our tracker builds on PHALP, please install that first:
```bash
pip install git+https://github.com/brjathu/PHALP.git
```
Now, run `track.py` to reconstruct and track humans in any video. Input video source may be a video file, a folder of frames, or a youtube link:
```bash
# Run on video file
python track.py video.source="example_data/videos/gymnasts.mp4"
# Run on extracted frames
python track.py video.source="/path/to/frames_folder/"
# Run on a youtube link (depends on pytube working properly)
python track.py video.source=\'"https://www.youtube.com/watch?v=xEH_5T9jMVU"\'
```
The output directory (`./outputs` by default) will contain a video rendering of the tracklets and a `.pkl` file containing the tracklets with 3D pose and shape. Please see the [PHALP](https://github.com/brjathu/PHALP) repository for details.
## Training
Download the [training data](https://www.dropbox.com/sh/mjdwu59fxuhls5h/AACQ6FCGSrggUXmRzuubRHXIa) to `./hmr2_training_data/`, then start training using the following command:
```
bash fetch_training_data.sh
python train.py exp_name=hmr2 data=mix_all experiment=hmr_vit_transformer trainer=gpu launcher=local
```
Checkpoints and logs will be saved to `./logs/`. We trained on 8 A100 GPUs for 7 days using PyTorch 1.13.1 and PyTorch-Lightning 1.8.1 with CUDA 11.6 on a Linux system. You may adjust batch size and number of GPUs per your convenience.
## Evaluation
Download the [evaluation metadata](https://www.dropbox.com/scl/fi/kl79djemdgqcl6d691er7/hmr2_evaluation_data.tar.gz?rlkey=ttmbdu3x5etxwqqyzwk581zjl) to `./hmr2_evaluation_data/`. Additionally, download the Human3.6M, 3DPW, LSP-Extended, COCO, and PoseTrack dataset images and update the corresponding paths in `hmr2/configs/datasets_eval.yaml`.
Run evaluation on multiple datasets as follows, results are stored in `results/eval_regression.csv`.
```bash
python eval.py --dataset 'H36M-VAL-P2,3DPW-TEST,LSP-EXTENDED,POSETRACK-VAL,COCO-VAL'
```
By default, our code uses the released checkpoint (mentioned as HMR2.0b in the paper). To use the HMR2.0a checkpoint, you may download and untar from [here](https://people.eecs.berkeley.edu/~jathushan/projects/4dhumans/hmr2a_model.tar.gz)
## Preprocess code
To preprocess LSP Extended and Posetrack into metadata zip files for evaluation, see `hmr2/datasets/preprocess`.
Training data preprocessing coming soon.
## Open Source Contributions
[carlosedubarreto](https://github.com/carlosedubarreto/) has created a tutorial to import 4D Humans in Blender: https://www.patreon.com/posts/86992009
## Acknowledgements
Parts of the code are taken or adapted from the following repos:
- [ProHMR](https://github.com/nkolot/ProHMR)
- [SPIN](https://github.com/nkolot/SPIN)
- [SMPLify-X](https://github.com/vchoutas/smplify-x)
- [HMR](https://github.com/akanazawa/hmr)
- [ViTPose](https://github.com/ViTAE-Transformer/ViTPose)
- [Detectron2](https://github.com/facebookresearch/detectron2)
Additionally, we thank [StabilityAI](https://stability.ai/) for a generous compute grant that enabled this work.
## Citing
If you find this code useful for your research, please consider citing the following paper:
```bibtex
@inproceedings{goel2023humans,
title={Humans in 4{D}: Reconstructing and Tracking Humans with Transformers},
author={Goel, Shubham and Pavlakos, Georgios and Rajasegaran, Jathushan and Kanazawa, Angjoo and Malik, Jitendra},
booktitle={ICCV},
year={2023}
}
```
View File
+116
View File
@@ -0,0 +1,116 @@
import os
from typing import Dict
from yacs.config import CfgNode as CN
from pathlib import Path
CACHE_DIR_4DHUMANS = os.environ.get("4DHUMAN_CACHE", str(Path(__file__).parent.parent.parent.parent / "ckpts"))
def to_lower(x: Dict) -> Dict:
"""
Convert all dictionary keys to lowercase
Args:
x (dict): Input dictionary
Returns:
dict: Output dictionary with all keys converted to lowercase
"""
return {k.lower(): v for k, v in x.items()}
_C = CN(new_allowed=True)
_C.GENERAL = CN(new_allowed=True)
_C.GENERAL.RESUME = True
_C.GENERAL.TIME_TO_RUN = 3300
_C.GENERAL.VAL_STEPS = 100
_C.GENERAL.LOG_STEPS = 100
_C.GENERAL.CHECKPOINT_STEPS = 20000
_C.GENERAL.CHECKPOINT_DIR = "checkpoints"
_C.GENERAL.SUMMARY_DIR = "tensorboard"
_C.GENERAL.NUM_GPUS = 1
_C.GENERAL.NUM_WORKERS = 4
_C.GENERAL.MIXED_PRECISION = True
_C.GENERAL.ALLOW_CUDA = True
_C.GENERAL.PIN_MEMORY = False
_C.GENERAL.DISTRIBUTED = False
_C.GENERAL.LOCAL_RANK = 0
_C.GENERAL.USE_SYNCBN = False
_C.GENERAL.WORLD_SIZE = 1
_C.TRAIN = CN(new_allowed=True)
_C.TRAIN.NUM_EPOCHS = 100
_C.TRAIN.BATCH_SIZE = 32
_C.TRAIN.SHUFFLE = True
_C.TRAIN.WARMUP = False
_C.TRAIN.NORMALIZE_PER_IMAGE = False
_C.TRAIN.CLIP_GRAD = False
_C.TRAIN.CLIP_GRAD_VALUE = 1.0
_C.LOSS_WEIGHTS = CN(new_allowed=True)
_C.DATASETS = CN(new_allowed=True)
_C.MODEL = CN(new_allowed=True)
_C.MODEL.IMAGE_SIZE = 224
_C.EXTRA = CN(new_allowed=True)
_C.EXTRA.FOCAL_LENGTH = 5000
_C.DATASETS.CONFIG = CN(new_allowed=True)
_C.DATASETS.CONFIG.SCALE_FACTOR = 0.3
_C.DATASETS.CONFIG.ROT_FACTOR = 30
_C.DATASETS.CONFIG.TRANS_FACTOR = 0.02
_C.DATASETS.CONFIG.COLOR_SCALE = 0.2
_C.DATASETS.CONFIG.ROT_AUG_RATE = 0.6
_C.DATASETS.CONFIG.TRANS_AUG_RATE = 0.5
_C.DATASETS.CONFIG.DO_FLIP = True
_C.DATASETS.CONFIG.FLIP_AUG_RATE = 0.5
_C.DATASETS.CONFIG.EXTREME_CROP_AUG_RATE = 0.10
def default_config() -> CN:
"""
Get a yacs CfgNode object with the default config values.
"""
# Return a clone so that the defaults will not be altered
# This is for the "local variable" use pattern
return _C.clone()
def dataset_config(name='datasets_tar.yaml') -> CN:
"""
Get dataset config file
Returns:
CfgNode: Dataset config as a yacs CfgNode object.
"""
cfg = CN(new_allowed=True)
config_file = os.path.join(os.path.dirname(os.path.realpath(__file__)), name)
cfg.merge_from_file(config_file)
cfg.freeze()
return cfg
def dataset_eval_config() -> CN:
return dataset_config('datasets_eval.yaml')
def get_config(config_file: str, merge: bool = True, update_cachedir: bool = False) -> CN:
"""
Read a config file and optionally merge it with the default config file.
Args:
config_file (str): Path to config file.
merge (bool): Whether to merge with the default config or not.
Returns:
CfgNode: Config as a yacs CfgNode object.
"""
if merge:
cfg = default_config()
else:
cfg = CN(new_allowed=True)
cfg.merge_from_file(config_file)
if update_cachedir:
def update_path(path: str) -> str:
if os.path.isabs(path):
return path
return os.path.join(CACHE_DIR_4DHUMANS, path)
cfg.SMPL.MODEL_PATH = update_path(cfg.SMPL.MODEL_PATH)
cfg.SMPL.JOINT_REGRESSOR_EXTRA = update_path(cfg.SMPL.JOINT_REGRESSOR_EXTRA)
cfg.SMPL.MEAN_PARAMS = update_path(cfg.SMPL.MEAN_PARAMS)
cfg.freeze()
return cfg
@@ -0,0 +1,129 @@
## coco_loader_lsj.py
import detectron2.data.transforms as T
from detectron2 import model_zoo
from detectron2.config import LazyCall as L
# Data using LSJ
image_size = 1024
dataloader = model_zoo.get_config("common/data/coco.py").dataloader
dataloader.train.mapper.augmentations = [
L(T.RandomFlip)(horizontal=True), # flip first
L(T.ResizeScale)(
min_scale=0.1, max_scale=2.0, target_height=image_size, target_width=image_size
),
L(T.FixedSizeCrop)(crop_size=(image_size, image_size), pad=False),
]
dataloader.train.mapper.image_format = "RGB"
dataloader.train.total_batch_size = 64
# recompute boxes due to cropping
dataloader.train.mapper.recompute_boxes = True
dataloader.test.mapper.augmentations = [
L(T.ResizeShortestEdge)(short_edge_length=image_size, max_size=image_size),
]
from functools import partial
from fvcore.common.param_scheduler import MultiStepParamScheduler
from detectron2 import model_zoo
from detectron2.config import LazyCall as L
from detectron2.solver import WarmupParamScheduler
from detectron2.modeling.backbone.vit import get_vit_lr_decay_rate
# mask_rcnn_vitdet_b_100ep.py
model = model_zoo.get_config("common/models/mask_rcnn_vitdet.py").model
# Initialization and trainer settings
train = model_zoo.get_config("common/train.py").train
train.amp.enabled = True
train.ddp.fp16_compression = True
train.init_checkpoint = "detectron2://ImageNetPretrained/MAE/mae_pretrain_vit_base.pth"
# Schedule
# 100 ep = 184375 iters * 64 images/iter / 118000 images/ep
train.max_iter = 184375
lr_multiplier = L(WarmupParamScheduler)(
scheduler=L(MultiStepParamScheduler)(
values=[1.0, 0.1, 0.01],
milestones=[163889, 177546],
num_updates=train.max_iter,
),
warmup_length=250 / train.max_iter,
warmup_factor=0.001,
)
# Optimizer
optimizer = model_zoo.get_config("common/optim.py").AdamW
optimizer.params.lr_factor_func = partial(get_vit_lr_decay_rate, num_layers=12, lr_decay_rate=0.7)
optimizer.params.overrides = {"pos_embed": {"weight_decay": 0.0}}
# cascade_mask_rcnn_vitdet_b_100ep.py
from detectron2.config import LazyCall as L
from detectron2.layers import ShapeSpec
from detectron2.modeling.box_regression import Box2BoxTransform
from detectron2.modeling.matcher import Matcher
from detectron2.modeling.roi_heads import (
FastRCNNOutputLayers,
FastRCNNConvFCHead,
CascadeROIHeads,
)
# arguments that don't exist for Cascade R-CNN
[model.roi_heads.pop(k) for k in ["box_head", "box_predictor", "proposal_matcher"]]
model.roi_heads.update(
_target_=CascadeROIHeads,
box_heads=[
L(FastRCNNConvFCHead)(
input_shape=ShapeSpec(channels=256, height=7, width=7),
conv_dims=[256, 256, 256, 256],
fc_dims=[1024],
conv_norm="LN",
)
for _ in range(3)
],
box_predictors=[
L(FastRCNNOutputLayers)(
input_shape=ShapeSpec(channels=1024),
test_score_thresh=0.05,
box2box_transform=L(Box2BoxTransform)(weights=(w1, w1, w2, w2)),
cls_agnostic_bbox_reg=True,
num_classes="${...num_classes}",
)
for (w1, w2) in [(10, 5), (20, 10), (30, 15)]
],
proposal_matchers=[
L(Matcher)(thresholds=[th], labels=[0, 1], allow_low_quality_matches=False)
for th in [0.5, 0.6, 0.7]
],
)
# cascade_mask_rcnn_vitdet_h_75ep.py
from functools import partial
train.init_checkpoint = "detectron2://ImageNetPretrained/MAE/mae_pretrain_vit_huge_p14to16.pth"
model.backbone.net.embed_dim = 1280
model.backbone.net.depth = 32
model.backbone.net.num_heads = 16
model.backbone.net.drop_path_rate = 0.5
# 7, 15, 23, 31 for global attention
model.backbone.net.window_block_indexes = (
list(range(0, 7)) + list(range(8, 15)) + list(range(16, 23)) + list(range(24, 31))
)
optimizer.params.lr_factor_func = partial(get_vit_lr_decay_rate, lr_decay_rate=0.9, num_layers=32)
optimizer.params.overrides = {}
optimizer.params.weight_decay_norm = None
train.max_iter = train.max_iter * 3 // 4 # 100ep -> 75ep
lr_multiplier.scheduler.milestones = [
milestone * 3 // 4 for milestone in lr_multiplier.scheduler.milestones
]
lr_multiplier.scheduler.num_updates = train.max_iter
@@ -0,0 +1,31 @@
H36M-VAL-P2:
TYPE: ImageDataset
DATASET_FILE: hmr2_evaluation_data/h36m_val_p2.npz
IMG_DIR: /shared/pavlakos/datasets/h36m/images/
KEYPOINT_LIST: [25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35, 36, 37, 43]
USE_HIPS: True
3DPW-TEST:
TYPE: ImageDataset
DATASET_FILE: hmr2_evaluation_data/3dpw_test.npz
IMG_DIR: /shared/pavlakos/datasets/3DPW/
KEYPOINT_LIST: [25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35, 36, 37, 43]
USE_HIPS: False
POSETRACK-VAL:
TYPE: ImageDataset
DATASET_FILE: hmr2_evaluation_data/posetrack_2018_val.npz
IMG_DIR: /shared/pavlakos/datasets/posetrack/posetrack2018/posetrack_data/
KEYPOINT_LIST: [0] # Dummy
LSP-EXTENDED:
TYPE: ImageDataset
DATASET_FILE: hmr2_evaluation_data/hr-lspet_train.npz
IMG_DIR: /shared/pavlakos/datasets/hr-lspet/
KEYPOINT_LIST: [0] # Dummy
COCO-VAL:
TYPE: ImageDataset
DATASET_FILE: hmr2_evaluation_data/coco_val.npz
IMG_DIR: /shared/pavlakos/datasets/coco/
KEYPOINT_LIST: [0] # Dummy
@@ -0,0 +1,38 @@
MPI-INF-TRAIN-PRUNED:
TYPE: ImageDataset
URLS: hmr2_training_data/dataset_tars/mpi-inf-train-pruned/{000000..00006}.tar
epoch_size: 12_000
H36M-TRAIN-WMASK:
TYPE: ImageDataset
URLS: hmr2_training_data/dataset_tars/h36m-train/{000000..000312}.tar
epoch_size: 314_000 # Can be changed arbitrarily. Enables resampling.
MPII-TRAIN-WMASK:
TYPE: ImageDataset
URLS: hmr2_training_data/dataset_tars/mpii-train/{000000..000009}.tar
epoch_size: 100_000 # Increase to ensure dataset doesn't end before epoch
COCO-TRAIN-2014-WMASK-PRUNED:
TYPE: ImageDataset
URLS: hmr2_training_data/dataset_tars/coco-train-2014-pruned/{000000..000017}.tar
epoch_size: 18_000
AVA-TRAIN-MIDFRAMES-1FPS-WMASK:
TYPE: ImageDataset
URLS: hmr2_training_data/dataset_tars/ava-train-midframes-1fps-vitpose/{000000..000092}.tar
epoch_size: 200_000
AIC-TRAIN-WMASK:
TYPE: ImageDataset
URLS: hmr2_training_data/dataset_tars/aic-train-vitpose/{000000..000104}.tar
epoch_size: 200_000
INSTA-TRAIN-WMASK:
TYPE: ImageDataset
URLS: hmr2_training_data/dataset_tars/insta-train-vitpose-replicate/{000000..003657}.tar
epoch_size: 4_000_000
COCO-TRAIN-2014-VITPOSE-REPLICATE-PRUNED12:
TYPE: ImageDataset
URLS: hmr2_training_data/dataset_tars/coco-train-2014-vitpose-pruned/{000000..000044}.tar
epoch_size: 45_000
COCO-VAL:
TYPE: ImageDataset
URLS: hmr2_training_data/dataset_tars/coco-val/{000000..000000}.tar
KEYPOINT_LIST: [0] # Dummy
CMU-MOCAP:
DATASET_FILE: hmr2_training_data/cmu_mocap.npz
@@ -0,0 +1,27 @@
# @package _global_
defaults:
- /data_filtering: low1
DATASETS:
TRAIN:
H36M-TRAIN-WMASK:
WEIGHT: 0.1
MPII-TRAIN-WMASK:
WEIGHT: 0.1
COCO-TRAIN-2014-WMASK-PRUNED:
WEIGHT: 0.1
COCO-TRAIN-2014-VITPOSE-REPLICATE-PRUNED12:
WEIGHT: 0.1
MPI-INF-TRAIN-PRUNED:
WEIGHT: 0.02
AVA-TRAIN-MIDFRAMES-1FPS-WMASK:
WEIGHT: 0.19
AIC-TRAIN-WMASK:
WEIGHT: 0.19
INSTA-TRAIN-WMASK:
WEIGHT: 0.2
VAL:
COCO-VAL:
WEIGHT: 1.0
MOCAP: CMU-MOCAP
@@ -0,0 +1,13 @@
# @package _global_
DATASETS:
# Data filtering during training
SUPPRESS_KP_CONF_THRESH: 0.3
FILTER_NUM_KP: 4
FILTER_NUM_KP_THRESH: 0.0
FILTER_REPROJ_THRESH: 31000
SUPPRESS_BETAS_THRESH: 3.0
SUPPRESS_BAD_POSES: True
POSES_BETAS_SIMULTANEOUS: True
FILTER_NO_POSES: False # If True, filters images that don't have poses
@@ -0,0 +1,29 @@
# @package _global_
SMPL:
DATA_DIR: ${oc.env:HOME}/.cache/4DHumans/data/
MODEL_PATH: ${SMPL.DATA_DIR}/smpl
GENDER: neutral
NUM_BODY_JOINTS: 23
JOINT_REGRESSOR_EXTRA: ${SMPL.DATA_DIR}/SMPL_to_J19.pkl
MEAN_PARAMS: ${SMPL.DATA_DIR}/smpl_mean_params.npz
EXTRA:
FOCAL_LENGTH: 5000
NUM_LOG_IMAGES: 4
NUM_LOG_SAMPLES_PER_IMAGE: 8
PELVIS_IND: 39
DATASETS:
BETAS_REG: True
CONFIG:
SCALE_FACTOR: 0.3
ROT_FACTOR: 30
TRANS_FACTOR: 0.02
COLOR_SCALE: 0.2
ROT_AUG_RATE: 0.6
TRANS_AUG_RATE: 0.5
DO_FLIP: True
FLIP_AUG_RATE: 0.5
EXTREME_CROP_AUG_RATE: 0.10
EXTREME_CROP_AUG_LEVEL: 1
@@ -0,0 +1,51 @@
# @package _global_
defaults:
- default.yaml
GENERAL:
TOTAL_STEPS: 1_000_000
LOG_STEPS: 1000
VAL_STEPS: 1000
CHECKPOINT_STEPS: 10000
CHECKPOINT_SAVE_TOP_K: 1
NUM_WORKERS: 6
PREFETCH_FACTOR: 2
TRAIN:
LR: 1e-5
WEIGHT_DECAY: 1e-4
BATCH_SIZE: 48
LOSS_REDUCTION: mean
NUM_TRAIN_SAMPLES: 2
NUM_TEST_SAMPLES: 64
POSE_2D_NOISE_RATIO: 0.01
SMPL_PARAM_NOISE_RATIO: 0.005
MODEL:
IMAGE_SIZE: 256
IMAGE_MEAN: [0.485, 0.456, 0.406]
IMAGE_STD: [0.229, 0.224, 0.225]
BACKBONE:
TYPE: vit
PRETRAINED_WEIGHTS: hmr2_training_data/vitpose_backbone.pth
SMPL_HEAD:
TYPE: transformer_decoder
IN_CHANNELS: 2048
TRANSFORMER_DECODER:
depth: 6
heads: 8
mlp_dim: 1024
dim_head: 64
dropout: 0.0
emb_dropout: 0.0
norm: layer
context_dim: 1280 # from vitpose-H
LOSS_WEIGHTS:
KEYPOINTS_3D: 0.05
KEYPOINTS_2D: 0.01
GLOBAL_ORIENT: 0.001
BODY_POSE: 0.001
BETAS: 0.0005
ADVERSARIAL: 0.0005
@@ -0,0 +1,8 @@
# disable python warnings if they annoy you
ignore_warnings: False
# ask user for tags if none are provided in the config
enforce_tags: True
# pretty print config tree at the start of the run using Rich library
print_config: True
@@ -0,0 +1,26 @@
# @package _global_
# https://hydra.cc/docs/configure_hydra/intro/
# enable color logging
defaults:
- override /hydra/hydra_logging: colorlog
- override /hydra/job_logging: colorlog
# exp_name: ovrd_${hydra:job.override_dirname}
exp_name: ${now:%Y-%m-%d}_${now:%H-%M-%S}
hydra:
run:
dir: ${paths.log_dir}/${task_name}/runs/${exp_name}
sweep:
dir: ${paths.log_dir}/${task_name}/multiruns/${exp_name}
subdir: ${hydra.job.num}
job:
config:
override_dirname:
exclude_keys:
- trainer
- trainer.devices
- trainer.num_nodes
- callbacks
- debug
@@ -0,0 +1,13 @@
# @package _global_
defaults:
- override /hydra/launcher: submitit_local
hydra:
launcher:
timeout_min: 10_080 # 7 days
nodes: 1
tasks_per_node: ${trainer.devices}
cpus_per_task: 6
gpus_per_node: ${trainer.devices}
name: hmr2
@@ -0,0 +1,22 @@
# @package _global_
defaults:
- override /hydra/launcher: submitit_slurm
hydra:
launcher:
timeout_min: 10_080 # 7 days
max_num_timeout: 3
partition: g40
qos: idle
nodes: 1
tasks_per_node: ${trainer.devices}
gpus_per_task: null
cpus_per_task: 12
gpus_per_node: ${trainer.devices}
cpus_per_gpu: null
comment: laion
name: hmr2
setup:
- module load cuda openmpi libfabric-aws
- export CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7
@@ -0,0 +1,18 @@
# path to root directory
# this requires PROJECT_ROOT environment variable to exist
# PROJECT_ROOT is inferred and set by pyrootutils package in `train.py` and `eval.py`
root_dir: ${oc.env:PROJECT_ROOT}
# path to data directory
data_dir: ${paths.root_dir}/data/
# path to logging directory
log_dir: logs/
# path to output directory, created dynamically by hydra
# path generation pattern is specified in `configs/hydra/default.yaml`
# use it to store all files generated during the run, like ckpts and metrics
output_dir: ${hydra:runtime.output_dir}
# path to working directory
work_dir: ${hydra:runtime.cwd}
@@ -0,0 +1,47 @@
# @package _global_
# specify here default configuration
# order of defaults determines the order in which configs override each other
defaults:
- _self_
- data: mix_all.yaml
- trainer: ddp.yaml
- paths: default.yaml
- extras: default.yaml
- hydra: default.yaml
# experiment configs allow for version control of specific hyperparameters
# e.g. best hyperparameters for given model and datamodule
- experiment: null
- texture_exp: null
# optional local config for machine/user specific settings
# it's optional since it doesn't need to exist and is excluded from version control
- optional launcher: local.yaml
# - optional launcher: slurm.yaml
# debugging config (enable through command line, e.g. `python train.py debug=default)
- debug: null
# task name, determines output directory path
task_name: "train"
# tags to help you identify your experiments
# you can overwrite this in experiment configs
# overwrite from command line with `python train.py tags="[first_tag, second_tag]"`
# appending lists from command line is currently not supported :(
# https://github.com/facebookresearch/hydra/issues/1547
tags: ["dev"]
# set False to skip model training
train: True
# evaluate on test set, using best model weights achieved during training
# lightning chooses best weights based on the metric specified in checkpoint callback
test: False
# simply provide checkpoint path to resume training
ckpt_path: null
# seed for random number generators in pytorch, numpy and python.random
seed: null
@@ -0,0 +1,6 @@
defaults:
- default.yaml
- default_hmr.yaml
accelerator: cpu
devices: 1
@@ -0,0 +1,14 @@
defaults:
- default.yaml
- default_hmr.yaml
# use "ddp_spawn" instead of "ddp",
# it's slower but normal "ddp" currently doesn't work ideally with hydra
# https://github.com/facebookresearch/hydra/issues/2070
# https://pytorch-lightning.readthedocs.io/en/latest/accelerators/gpu_intermediate.html#distributed-data-parallel-spawn
strategy: ddp
accelerator: gpu
devices: 8
num_nodes: 1
sync_batchnorm: True
@@ -0,0 +1,10 @@
_target_: pytorch_lightning.Trainer
default_root_dir: ${paths.output_dir}
accelerator: cpu
devices: 1
# set True to to ensure deterministic results
# makes training slower but gives more reproducibility than just setting seeds
deterministic: False
@@ -0,0 +1,8 @@
num_sanity_val_steps: 0
log_every_n_steps: ${GENERAL.LOG_STEPS}
val_check_interval: ${GENERAL.VAL_STEPS}
precision: 16
max_steps: ${GENERAL.TOTAL_STEPS}
# move_metrics_to_cpu: True
limit_val_batches: 1
# track_grad_norm: -1
@@ -0,0 +1,6 @@
defaults:
- default.yaml
- default_hmr.yaml
accelerator: gpu
devices: 1
@@ -0,0 +1,6 @@
defaults:
- default.yaml
- default_hmr.yaml
accelerator: mps
devices: 1
@@ -0,0 +1,88 @@
from typing import Dict, Optional
import torch
import numpy as np
import pytorch_lightning as pl
from yacs.config import CfgNode
import webdataset as wds
from ..configs import to_lower
from .dataset import Dataset
from .image_dataset import ImageDataset
from .mocap_dataset import MoCapDataset
def create_dataset(cfg: CfgNode, dataset_cfg: CfgNode, train: bool = True, **kwargs) -> Dataset:
"""
Instantiate a dataset from a config file.
Args:
cfg (CfgNode): Model configuration file.
dataset_cfg (CfgNode): Dataset configuration info.
train (bool): Variable to select between train and val datasets.
"""
dataset_type = Dataset.registry[dataset_cfg.TYPE]
return dataset_type(cfg, **to_lower(dataset_cfg), train=train, **kwargs)
def create_webdataset(cfg: CfgNode, dataset_cfg: CfgNode, train: bool = True) -> Dataset:
"""
Like `create_dataset` but load data from tars.
"""
dataset_type = Dataset.registry[dataset_cfg.TYPE]
return dataset_type.load_tars_as_webdataset(cfg, **to_lower(dataset_cfg), train=train)
class MixedWebDataset(wds.WebDataset):
def __init__(self, cfg: CfgNode, dataset_cfg: CfgNode, train: bool = True) -> None:
super(wds.WebDataset, self).__init__()
dataset_list = cfg.DATASETS.TRAIN if train else cfg.DATASETS.VAL
datasets = [create_webdataset(cfg, dataset_cfg[dataset], train=train) for dataset, v in dataset_list.items()]
weights = np.array([v.WEIGHT for dataset, v in dataset_list.items()])
weights = weights / weights.sum() # normalize
self.append(wds.RandomMix(datasets, weights))
class HMR2DataModule(pl.LightningDataModule):
def __init__(self, cfg: CfgNode, dataset_cfg: CfgNode) -> None:
"""
Initialize LightningDataModule for HMR2 training
Args:
cfg (CfgNode): Config file as a yacs CfgNode containing necessary dataset info.
dataset_cfg (CfgNode): Dataset configuration file
"""
super().__init__()
self.cfg = cfg
self.dataset_cfg = dataset_cfg
self.train_dataset = None
self.val_dataset = None
self.test_dataset = None
self.mocap_dataset = None
def setup(self, stage: Optional[str] = None) -> None:
"""
Load datasets necessary for training
Args:
cfg (CfgNode): Config file as a yacs CfgNode containing necessary dataset info.
"""
if self.train_dataset == None:
self.train_dataset = MixedWebDataset(self.cfg, self.dataset_cfg, train=True).with_epoch(100_000).shuffle(4000)
self.val_dataset = MixedWebDataset(self.cfg, self.dataset_cfg, train=False).shuffle(4000)
self.mocap_dataset = MoCapDataset(**to_lower(self.dataset_cfg[self.cfg.DATASETS.MOCAP]))
def train_dataloader(self) -> Dict:
"""
Setup training data loader.
Returns:
Dict: Dictionary containing image and mocap data dataloaders
"""
train_dataloader = torch.utils.data.DataLoader(self.train_dataset, self.cfg.TRAIN.BATCH_SIZE, drop_last=True, num_workers=self.cfg.GENERAL.NUM_WORKERS, prefetch_factor=self.cfg.GENERAL.PREFETCH_FACTOR)
mocap_dataloader = torch.utils.data.DataLoader(self.mocap_dataset, self.cfg.TRAIN.NUM_TRAIN_SAMPLES * self.cfg.TRAIN.BATCH_SIZE, shuffle=True, drop_last=True, num_workers=1)
return {'img': train_dataloader, 'mocap': mocap_dataloader}
def val_dataloader(self) -> torch.utils.data.DataLoader:
"""
Setup val data loader.
Returns:
torch.utils.data.DataLoader: Validation dataloader
"""
val_dataloader = torch.utils.data.DataLoader(self.val_dataset, self.cfg.TRAIN.BATCH_SIZE, drop_last=True, num_workers=self.cfg.GENERAL.NUM_WORKERS)
return val_dataloader
@@ -0,0 +1,27 @@
"""
This file contains the defition of the base Dataset class.
"""
class DatasetRegistration(type):
"""
Metaclass for registering different datasets
"""
def __init__(cls, name, bases, nmspc):
super().__init__(name, bases, nmspc)
if not hasattr(cls, 'registry'):
cls.registry = dict()
cls.registry[name] = cls
# Metamethods, called on class objects:
def __iter__(cls):
return iter(cls.registry)
def __str__(cls):
return str(cls.registry)
class Dataset(metaclass=DatasetRegistration):
"""
Base Dataset class
"""
def __init__(self, *args, **kwargs):
pass
@@ -0,0 +1,454 @@
import copy
import os
import numpy as np
import torch
from typing import Any, Dict, List
from yacs.config import CfgNode
import braceexpand
import cv2
from .dataset import Dataset
from .utils import get_example, expand_to_aspect_ratio
from .smplh_prob_filter import poses_check_probable, load_amass_hist_smooth
def expand(s):
return os.path.expanduser(os.path.expandvars(s))
def expand_urls(urls: str|List[str]):
if isinstance(urls, str):
urls = [urls]
urls = [u for url in urls for u in braceexpand.braceexpand(expand(url))]
return urls
AIC_TRAIN_CORRUPT_KEYS = {
'0a047f0124ae48f8eee15a9506ce1449ee1ba669',
'1a703aa174450c02fbc9cfbf578a5435ef403689',
'0394e6dc4df78042929b891dbc24f0fd7ffb6b6d',
'5c032b9626e410441544c7669123ecc4ae077058',
'ca018a7b4c5f53494006ebeeff9b4c0917a55f07',
'4a77adb695bef75a5d34c04d589baf646fe2ba35',
'a0689017b1065c664daef4ae2d14ea03d543217e',
'39596a45cbd21bed4a5f9c2342505532f8ec5cbb',
'3d33283b40610d87db660b62982f797d50a7366b',
}
CORRUPT_KEYS = {
*{f'aic-train/{k}' for k in AIC_TRAIN_CORRUPT_KEYS},
*{f'aic-train-vitpose/{k}' for k in AIC_TRAIN_CORRUPT_KEYS},
}
body_permutation = [0, 1, 5, 6, 7, 2, 3, 4, 8, 12, 13, 14, 9, 10, 11, 16, 15, 18, 17, 22, 23, 24, 19, 20, 21]
extra_permutation = [5, 4, 3, 2, 1, 0, 11, 10, 9, 8, 7, 6, 12, 13, 14, 15, 16, 17, 18]
FLIP_KEYPOINT_PERMUTATION = body_permutation + [25 + i for i in extra_permutation]
DEFAULT_MEAN = 255. * np.array([0.485, 0.456, 0.406])
DEFAULT_STD = 255. * np.array([0.229, 0.224, 0.225])
DEFAULT_IMG_SIZE = 256
class ImageDataset(Dataset):
def __init__(self,
cfg: CfgNode,
dataset_file: str,
img_dir: str,
train: bool = True,
prune: Dict[str, Any] = {},
**kwargs):
"""
Dataset class used for loading images and corresponding annotations.
Args:
cfg (CfgNode): Model config file.
dataset_file (str): Path to npz file containing dataset info.
img_dir (str): Path to image folder.
train (bool): Whether it is for training or not (enables data augmentation).
"""
super(ImageDataset, self).__init__()
self.train = train
self.cfg = cfg
self.img_size = cfg.MODEL.IMAGE_SIZE
self.mean = 255. * np.array(self.cfg.MODEL.IMAGE_MEAN)
self.std = 255. * np.array(self.cfg.MODEL.IMAGE_STD)
self.img_dir = img_dir
self.data = np.load(dataset_file, allow_pickle=True)
self.imgname = self.data['imgname']
self.personid = np.zeros(len(self.imgname), dtype=np.int32)
self.extra_info = self.data.get('extra_info', [{} for _ in range(len(self.imgname))])
self.flip_keypoint_permutation = copy.copy(FLIP_KEYPOINT_PERMUTATION)
num_pose = 3 * (self.cfg.SMPL.NUM_BODY_JOINTS + 1)
# Bounding boxes are assumed to be in the center and scale format
self.center = self.data['center']
self.scale = self.data['scale'].reshape(len(self.center), -1) / 200.0
if self.scale.shape[1] == 1:
self.scale = np.tile(self.scale, (1, 2))
assert self.scale.shape == (len(self.center), 2)
# Get gt SMPLX parameters, if available
try:
self.body_pose = self.data['body_pose'].astype(np.float32)
self.has_body_pose = self.data['has_body_pose'].astype(np.float32)
except KeyError:
self.body_pose = np.zeros((len(self.imgname), num_pose), dtype=np.float32)
self.has_body_pose = np.zeros(len(self.imgname), dtype=np.float32)
try:
self.betas = self.data['betas'].astype(np.float32)
self.has_betas = self.data['has_betas'].astype(np.float32)
except KeyError:
self.betas = np.zeros((len(self.imgname), 10), dtype=np.float32)
self.has_betas = np.zeros(len(self.imgname), dtype=np.float32)
# Try to get 2d keypoints, if available
try:
body_keypoints_2d = self.data['body_keypoints_2d']
except KeyError:
body_keypoints_2d = np.zeros((len(self.center), 25, 3))
# Try to get extra 2d keypoints, if available
try:
extra_keypoints_2d = self.data['extra_keypoints_2d']
except KeyError:
extra_keypoints_2d = np.zeros((len(self.center), 19, 3))
self.keypoints_2d = np.concatenate((body_keypoints_2d, extra_keypoints_2d), axis=1).astype(np.float32)
# Try to get 3d keypoints, if available
try:
body_keypoints_3d = self.data['body_keypoints_3d'].astype(np.float32)
except KeyError:
body_keypoints_3d = np.zeros((len(self.center), 25, 4), dtype=np.float32)
# Try to get extra 3d keypoints, if available
try:
extra_keypoints_3d = self.data['extra_keypoints_3d'].astype(np.float32)
except KeyError:
extra_keypoints_3d = np.zeros((len(self.center), 19, 4), dtype=np.float32)
body_keypoints_3d[:, [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14], -1] = 0
self.keypoints_3d = np.concatenate((body_keypoints_3d, extra_keypoints_3d), axis=1).astype(np.float32)
def __len__(self) -> int:
return len(self.scale)
def __getitem__(self, idx: int) -> Dict:
"""
Returns an example from the dataset.
"""
try:
image_file_rel = self.imgname[idx].decode('utf-8')
except AttributeError:
image_file_rel = self.imgname[idx]
image_file = os.path.join(self.img_dir, image_file_rel)
keypoints_2d = self.keypoints_2d[idx].copy()
keypoints_3d = self.keypoints_3d[idx].copy()
center = self.center[idx].copy()
center_x = center[0]
center_y = center[1]
scale = self.scale[idx]
BBOX_SHAPE = self.cfg.MODEL.get('BBOX_SHAPE', None)
bbox_size = expand_to_aspect_ratio(scale*200, target_aspect_ratio=BBOX_SHAPE).max()
bbox_expand_factor = bbox_size / ((scale*200).max())
body_pose = self.body_pose[idx].copy().astype(np.float32)
betas = self.betas[idx].copy().astype(np.float32)
has_body_pose = self.has_body_pose[idx].copy()
has_betas = self.has_betas[idx].copy()
smpl_params = {'global_orient': body_pose[:3],
'body_pose': body_pose[3:],
'betas': betas
}
has_smpl_params = {'global_orient': has_body_pose,
'body_pose': has_body_pose,
'betas': has_betas
}
smpl_params_is_axis_angle = {'global_orient': True,
'body_pose': True,
'betas': False
}
augm_config = self.cfg.DATASETS.CONFIG
# Crop image and (possibly) perform data augmentation
img_patch, keypoints_2d, keypoints_3d, smpl_params, has_smpl_params, img_size = get_example(image_file,
center_x, center_y,
bbox_size, bbox_size,
keypoints_2d, keypoints_3d,
smpl_params, has_smpl_params,
self.flip_keypoint_permutation,
self.img_size, self.img_size,
self.mean, self.std, self.train, augm_config)
item = {}
# These are the keypoints in the original image coordinates (before cropping)
orig_keypoints_2d = self.keypoints_2d[idx].copy()
item['img'] = img_patch
item['keypoints_2d'] = keypoints_2d.astype(np.float32)
item['keypoints_3d'] = keypoints_3d.astype(np.float32)
item['orig_keypoints_2d'] = orig_keypoints_2d
item['box_center'] = self.center[idx].copy()
item['box_size'] = bbox_size
item['bbox_expand_factor'] = bbox_expand_factor
item['img_size'] = 1.0 * img_size[::-1].copy()
item['smpl_params'] = smpl_params
item['has_smpl_params'] = has_smpl_params
item['smpl_params_is_axis_angle'] = smpl_params_is_axis_angle
item['imgname'] = image_file
item['imgname_rel'] = image_file_rel
item['personid'] = int(self.personid[idx])
item['extra_info'] = copy.deepcopy(self.extra_info[idx])
item['idx'] = idx
item['_scale'] = scale
return item
@staticmethod
def load_tars_as_webdataset(cfg: CfgNode, urls: str|List[str], train: bool,
resampled=False,
epoch_size=None,
cache_dir=None,
**kwargs) -> Dataset:
"""
Loads the dataset from a webdataset tar file.
"""
IMG_SIZE = cfg.MODEL.IMAGE_SIZE
BBOX_SHAPE = cfg.MODEL.get('BBOX_SHAPE', None)
MEAN = 255. * np.array(cfg.MODEL.IMAGE_MEAN)
STD = 255. * np.array(cfg.MODEL.IMAGE_STD)
def split_data(source):
for item in source:
datas = item['data.pyd']
for data in datas:
if 'detection.npz' in item:
det_idx = data['extra_info']['detection_npz_idx']
mask = item['detection.npz']['masks'][det_idx]
else:
mask = np.ones_like(item['jpg'][:,:,0], dtype=bool)
yield {
'__key__': item['__key__'],
'jpg': item['jpg'],
'data.pyd': data,
'mask': mask,
}
def suppress_bad_kps(item, thresh=0.0):
if thresh > 0:
kp2d = item['data.pyd']['keypoints_2d']
kp2d_conf = np.where(kp2d[:, 2] < thresh, 0.0, kp2d[:, 2])
item['data.pyd']['keypoints_2d'] = np.concatenate([kp2d[:,:2], kp2d_conf[:,None]], axis=1)
return item
def filter_numkp(item, numkp=4, thresh=0.0):
kp_conf = item['data.pyd']['keypoints_2d'][:, 2]
return (kp_conf > thresh).sum() > numkp
def filter_reproj_error(item, thresh=10**4.5):
losses = item['data.pyd'].get('extra_info', {}).get('fitting_loss', np.array({})).item()
reproj_loss = losses.get('reprojection_loss', None)
return reproj_loss is None or reproj_loss < thresh
def filter_bbox_size(item, thresh=1):
bbox_size_min = item['data.pyd']['scale'].min().item() * 200.
return bbox_size_min > thresh
def filter_no_poses(item):
return (item['data.pyd']['has_body_pose'] > 0)
def supress_bad_betas(item, thresh=3):
has_betas = item['data.pyd']['has_betas']
if thresh > 0 and has_betas:
betas_abs = np.abs(item['data.pyd']['betas'])
if (betas_abs > thresh).any():
item['data.pyd']['has_betas'] = False
return item
amass_poses_hist100_smooth = load_amass_hist_smooth()
def supress_bad_poses(item):
has_body_pose = item['data.pyd']['has_body_pose']
if has_body_pose:
body_pose = item['data.pyd']['body_pose']
pose_is_probable = poses_check_probable(torch.from_numpy(body_pose)[None, 3:], amass_poses_hist100_smooth).item()
if not pose_is_probable:
item['data.pyd']['has_body_pose'] = False
return item
def poses_betas_simultaneous(item):
# We either have both body_pose and betas, or neither
has_betas = item['data.pyd']['has_betas']
has_body_pose = item['data.pyd']['has_body_pose']
item['data.pyd']['has_betas'] = item['data.pyd']['has_body_pose'] = np.array(float((has_body_pose>0) and (has_betas>0)))
return item
def set_betas_for_reg(item):
# Always have betas set to true
has_betas = item['data.pyd']['has_betas']
betas = item['data.pyd']['betas']
if not (has_betas>0):
item['data.pyd']['has_betas'] = np.array(float((True)))
item['data.pyd']['betas'] = betas * 0
return item
# Load the dataset
if epoch_size is not None:
resampled = True
corrupt_filter = lambda sample: (sample['__key__'] not in CORRUPT_KEYS)
import webdataset as wds
dataset = wds.WebDataset(expand_urls(urls),
nodesplitter=wds.split_by_node,
shardshuffle=True,
resampled=resampled,
cache_dir=cache_dir,
).select(corrupt_filter)
if train:
dataset = dataset.shuffle(100)
dataset = dataset.decode('rgb8').rename(jpg='jpg;jpeg;png')
# Process the dataset
dataset = dataset.compose(split_data)
# Filter/clean the dataset
SUPPRESS_KP_CONF_THRESH = cfg.DATASETS.get('SUPPRESS_KP_CONF_THRESH', 0.0)
SUPPRESS_BETAS_THRESH = cfg.DATASETS.get('SUPPRESS_BETAS_THRESH', 0.0)
SUPPRESS_BAD_POSES = cfg.DATASETS.get('SUPPRESS_BAD_POSES', False)
POSES_BETAS_SIMULTANEOUS = cfg.DATASETS.get('POSES_BETAS_SIMULTANEOUS', False)
BETAS_REG = cfg.DATASETS.get('BETAS_REG', False)
FILTER_NO_POSES = cfg.DATASETS.get('FILTER_NO_POSES', False)
FILTER_NUM_KP = cfg.DATASETS.get('FILTER_NUM_KP', 4)
FILTER_NUM_KP_THRESH = cfg.DATASETS.get('FILTER_NUM_KP_THRESH', 0.0)
FILTER_REPROJ_THRESH = cfg.DATASETS.get('FILTER_REPROJ_THRESH', 0.0)
FILTER_MIN_BBOX_SIZE = cfg.DATASETS.get('FILTER_MIN_BBOX_SIZE', 0.0)
if SUPPRESS_KP_CONF_THRESH > 0:
dataset = dataset.map(lambda x: suppress_bad_kps(x, thresh=SUPPRESS_KP_CONF_THRESH))
if SUPPRESS_BETAS_THRESH > 0:
dataset = dataset.map(lambda x: supress_bad_betas(x, thresh=SUPPRESS_BETAS_THRESH))
if SUPPRESS_BAD_POSES:
dataset = dataset.map(lambda x: supress_bad_poses(x))
if POSES_BETAS_SIMULTANEOUS:
dataset = dataset.map(lambda x: poses_betas_simultaneous(x))
if FILTER_NO_POSES:
dataset = dataset.select(lambda x: filter_no_poses(x))
if FILTER_NUM_KP > 0:
dataset = dataset.select(lambda x: filter_numkp(x, numkp=FILTER_NUM_KP, thresh=FILTER_NUM_KP_THRESH))
if FILTER_REPROJ_THRESH > 0:
dataset = dataset.select(lambda x: filter_reproj_error(x, thresh=FILTER_REPROJ_THRESH))
if FILTER_MIN_BBOX_SIZE > 0:
dataset = dataset.select(lambda x: filter_bbox_size(x, thresh=FILTER_MIN_BBOX_SIZE))
if BETAS_REG:
dataset = dataset.map(lambda x: set_betas_for_reg(x)) # NOTE: Must be at the end
use_skimage_antialias = cfg.DATASETS.get('USE_SKIMAGE_ANTIALIAS', False)
border_mode = {
'constant': cv2.BORDER_CONSTANT,
'replicate': cv2.BORDER_REPLICATE,
}[cfg.DATASETS.get('BORDER_MODE', 'constant')]
# Process the dataset further
dataset = dataset.map(lambda x: ImageDataset.process_webdataset_tar_item(x, train,
augm_config=cfg.DATASETS.CONFIG,
MEAN=MEAN, STD=STD, IMG_SIZE=IMG_SIZE,
BBOX_SHAPE=BBOX_SHAPE,
use_skimage_antialias=use_skimage_antialias,
border_mode=border_mode,
))
if epoch_size is not None:
dataset = dataset.with_epoch(epoch_size)
return dataset
@staticmethod
def process_webdataset_tar_item(item, train,
augm_config=None,
MEAN=DEFAULT_MEAN,
STD=DEFAULT_STD,
IMG_SIZE=DEFAULT_IMG_SIZE,
BBOX_SHAPE=None,
use_skimage_antialias=False,
border_mode=cv2.BORDER_CONSTANT,
):
# Read data from item
key = item['__key__']
image = item['jpg']
data = item['data.pyd']
mask = item['mask']
keypoints_2d = data['keypoints_2d']
keypoints_3d = data['keypoints_3d']
center = data['center']
scale = data['scale']
body_pose = data['body_pose']
betas = data['betas']
has_body_pose = data['has_body_pose']
has_betas = data['has_betas']
# image_file = data['image_file']
# Process data
orig_keypoints_2d = keypoints_2d.copy()
center_x = center[0]
center_y = center[1]
bbox_size = expand_to_aspect_ratio(scale*200, target_aspect_ratio=BBOX_SHAPE).max()
if bbox_size < 1:
breakpoint()
smpl_params = {'global_orient': body_pose[:3],
'body_pose': body_pose[3:],
'betas': betas
}
has_smpl_params = {'global_orient': has_body_pose,
'body_pose': has_body_pose,
'betas': has_betas
}
smpl_params_is_axis_angle = {'global_orient': True,
'body_pose': True,
'betas': False
}
augm_config = copy.deepcopy(augm_config)
# Crop image and (possibly) perform data augmentation
img_rgba = np.concatenate([image, mask.astype(np.uint8)[:,:,None]*255], axis=2)
img_patch_rgba, keypoints_2d, keypoints_3d, smpl_params, has_smpl_params, img_size, trans = get_example(img_rgba,
center_x, center_y,
bbox_size, bbox_size,
keypoints_2d, keypoints_3d,
smpl_params, has_smpl_params,
FLIP_KEYPOINT_PERMUTATION,
IMG_SIZE, IMG_SIZE,
MEAN, STD, train, augm_config,
is_bgr=False, return_trans=True,
use_skimage_antialias=use_skimage_antialias,
border_mode=border_mode,
)
img_patch = img_patch_rgba[:3,:,:]
mask_patch = (img_patch_rgba[3,:,:] / 255.0).clip(0,1)
if (mask_patch < 0.5).all():
mask_patch = np.ones_like(mask_patch)
item = {}
item['img'] = img_patch
item['mask'] = mask_patch
# item['img_og'] = image
# item['mask_og'] = mask
item['keypoints_2d'] = keypoints_2d.astype(np.float32)
item['keypoints_3d'] = keypoints_3d.astype(np.float32)
item['orig_keypoints_2d'] = orig_keypoints_2d
item['box_center'] = center.copy()
item['box_size'] = bbox_size
item['img_size'] = 1.0 * img_size[::-1].copy()
item['smpl_params'] = smpl_params
item['has_smpl_params'] = has_smpl_params
item['smpl_params_is_axis_angle'] = smpl_params_is_axis_angle
item['_scale'] = scale
item['_trans'] = trans
item['imgname'] = key
# item['idx'] = idx
return item
@@ -0,0 +1,25 @@
import numpy as np
from typing import Dict
class MoCapDataset:
def __init__(self, dataset_file: str):
"""
Dataset class used for loading a dataset of unpaired SMPL parameter annotations
Args:
cfg (CfgNode): Model config file.
dataset_file (str): Path to npz file containing dataset info.
"""
data = np.load(dataset_file)
self.pose = data['body_pose'].astype(np.float32)[:, 3:]
self.betas = data['betas'].astype(np.float32)
self.length = len(self.pose)
def __getitem__(self, idx: int) -> Dict:
pose = self.pose[idx].copy()
betas = self.betas[idx].copy()
item = {'body_pose': pose, 'betas': betas}
return item
def __len__(self) -> int:
return self.length
@@ -0,0 +1,74 @@
# Adapted from https://raw.githubusercontent.com/nkolot/SPIN/master/datasets/preprocess/hr_lspet.py
import os
import glob
import numpy as np
import scipy.io as sio
# from .read_openpose import read_openpose
def hr_lspet_extract(dataset_path, out_path):
# training mode
png_path = os.path.join(dataset_path, '*.png')
imgs = glob.glob(png_path)
imgs.sort()
# structs we use
imgnames_, scales_, centers_, parts_, openposes_= [], [], [], [], []
# scale factor
scaleFactor = 1.2
# annotation files
annot_file = os.path.join(dataset_path, 'joints.mat')
joints = sio.loadmat(annot_file)['joints']
# main loop
for i, imgname in enumerate(imgs):
# image name
imgname = imgname.split('/')[-1]
# read keypoints
part14 = joints[:,:2,i]
# scale and center
bbox = [min(part14[:,0]), min(part14[:,1]),
max(part14[:,0]), max(part14[:,1])]
center = [(bbox[2]+bbox[0])/2, (bbox[3]+bbox[1])/2]
# scale = scaleFactor*max(bbox[2]-bbox[0], bbox[3]-bbox[1]) # Don't /200
scale = scaleFactor*np.array([bbox[2]-bbox[0], bbox[3]-bbox[1]]) # Don't /200
# update keypoints
part = np.zeros([24,3])
part[:14] = np.hstack([part14, np.ones([14,1])])
# # read openpose detections
# json_file = os.path.join(openpose_path, 'hrlspet',
# imgname.replace('.png', '_keypoints.json'))
# openpose = read_openpose(json_file, part, 'hrlspet')
# store the data
imgnames_.append(imgname)
centers_.append(center)
scales_.append(scale)
parts_.append(part)
# openposes_.append(openpose)
# Populate extra_keypoints_2d: N,19,3
# extra_keypoints_2d[:14] = parts[:14]
extra_keypoints_2d = np.zeros((len(parts_), 19, 3))
extra_keypoints_2d[:,:14,:] = np.stack(parts_)[:,:14,:3]
print(f'{extra_keypoints_2d.shape=}')
# store the data struct
if not os.path.isdir(out_path):
os.makedirs(out_path)
out_file = os.path.join(out_path, 'hr-lspet_train.npz')
np.savez(out_file, imgname=imgnames_,
center=centers_,
scale=scales_,
part=parts_,
extra_keypoints_2d=extra_keypoints_2d,
# openpose=openposes_
)
if __name__ == '__main__':
hr_lspet_extract('/shared/pavlakos/datasets/hr-lspet/', 'hmr2_evaluation_data/')
@@ -0,0 +1,92 @@
# Adapted from https://raw.githubusercontent.com/nkolot/SPIN/master/datasets/preprocess/coco.py
import os
from os.path import join
import sys
import json
import numpy as np
from pathlib import Path
# from .read_openpose import read_openpose
def coco_extract(dataset_path, out_path):
# # convert joints to global order
# joints_idx = [19, 20, 21, 22, 23, 9, 8, 10, 7, 11, 6, 3, 2, 4, 1, 5, 0]
# bbox expansion factor
scaleFactor = 1.2
# structs we need
imgnames_, scales_, centers_, parts_, openposes_ = [], [], [], [], []
# json annotation file
SPLIT='val'
json_paths = (Path(dataset_path)/'posetrack_data/annotations'/SPLIT).glob('*.json')
for json_path in json_paths:
json_data = json.load(open(json_path, 'r'))
imgs = {}
for img in json_data['images']:
imgs[img['id']] = img
for annot in json_data['annotations']:
# keypoints processing
keypoints = annot['keypoints']
keypoints = np.reshape(keypoints, (17,3))
keypoints[keypoints[:,2]>0,2] = 1
# check if all major body joints are annotated
if sum(keypoints[5:,2]>0) < 12:
continue
# image name
image_id = annot['image_id']
img_name = str(imgs[image_id]['file_name'])
# img_name_full = f'images/{SPLIT}/{json_path.stem}/{img_name}'
img_name_full = img_name
# keypoints
part = np.zeros([17,3])
# part[joints_idx] = keypoints
part = keypoints
# scale and center
bbox = annot['bbox']
center = [bbox[0] + bbox[2]/2, bbox[1] + bbox[3]/2]
# scale = scaleFactor*max(bbox[2], bbox[3]) # Don't do /200
scale = scaleFactor*np.array([bbox[2], bbox[3]]) # Don't /200
# # read openpose detections
# json_file = os.path.join(openpose_path, 'coco',
# img_name.replace('.jpg', '_keypoints.json'))
# openpose = read_openpose(json_file, part, 'coco')
# store data
imgnames_.append(img_name_full)
centers_.append(center)
scales_.append(scale)
parts_.append(part)
# openposes_.append(openpose)
# NOTE: Posetrack val doesn't annotate ears (17,18)
# But Posetrack does annotate head, neck so that wil have to live in extra_kps.
posetrack_to_op_extra = [0, 37, 38, 18, 17, 5, 2, 6, 3, 7, 4, 12, 9, 13, 10, 14, 11] # Will contain 15 keypoints.
all_keypoints_2d = np.zeros((len(parts_), 44, 3))
all_keypoints_2d[:,posetrack_to_op_extra] = np.stack(parts_)[:,:len(posetrack_to_op_extra),:3]
body_keypoints_2d = all_keypoints_2d[:,:25,:]
extra_keypoints_2d = all_keypoints_2d[:,25:,:]
print(f'{extra_keypoints_2d.shape=}')
# store the data struct
if not os.path.isdir(out_path):
os.makedirs(out_path)
out_file = os.path.join(out_path, f'posetrack_2018_{SPLIT}.npz')
np.savez(out_file, imgname=imgnames_,
center=centers_,
scale=scales_,
part=parts_,
body_keypoints_2d=body_keypoints_2d,
extra_keypoints_2d=extra_keypoints_2d,
# openpose=openposes_
)
if __name__ == '__main__':
coco_extract('/shared/pavlakos/datasets/posetrack/posetrack2018/', 'hmr2_evaluation_data/')
@@ -0,0 +1,154 @@
import os
import numpy as np
import torch
import torch.nn.functional as F
JOINT_NAMES = [
'left_hip',
'right_hip',
'spine1',
'left_knee',
'right_knee',
'spine2',
'left_ankle',
'right_ankle',
'spine3',
'left_foot',
'right_foot',
'neck',
'left_collar',
'right_collar',
'head',
'left_shoulder',
'right_shoulder',
'left_elbow',
'right_elbow',
'left_wrist',
'right_wrist'
]
# Manually chosen probability density thresholds for each joint
# Probablities computed using SIGMA=2 gaussian blur on AMASS pose 3D histogram for range (-pi,pi) with 100x100x100 bins
JOINT_NAME_PROB_THRESHOLDS = {
'left_hip': 5e-5,
'right_hip': 5e-5,
'spine1': 2e-3,
'left_knee': 5e-6,
'right_knee': 5e-6,
'spine2': 0.01,
'left_ankle': 5e-6,
'right_ankle': 5e-6,
'spine3': 0.025,
'left_foot': 0,
'right_foot': 0,
'neck': 2e-4,
'left_collar': 4.5e-4 ,
'right_collar': 4.5e-4,
'head': 5e-4,
'left_shoulder': 2e-4,
'right_shoulder': 2e-4,
'left_elbow': 4e-5,
'right_elbow': 4e-5,
'left_wrist': 1e-3,
'right_wrist': 1e-3,
}
JOINT_IDX_PROB_THRESHOLDS = torch.tensor([JOINT_NAME_PROB_THRESHOLDS[joint_name] for joint_name in JOINT_NAMES])
###################################################################
POSE_RANGE_MIN = -np.pi
POSE_RANGE_MAX = np.pi
# Create 21x100x100x100 histogram of all 21 AMASS body pose joints using `create_pose_hist(amass_poses, nbins=100)`
AMASS_HIST100_PATH = 'hmr2_training_data/amass_poses_hist100_SMPL+H_G.npy'
if not os.path.exists(AMASS_HIST100_PATH):
AMASS_HIST100_PATH = '/shared/shubham/code/hmr2023/amass_poses_hist100_SMPL+H_G.npy'
if not os.path.exists(AMASS_HIST100_PATH):
AMASS_HIST100_PATH = '/fsx/shubham/code/stable-humans/notebooks/amass_poses_hist100_SMPL+H_G.npy'
def create_pose_hist(poses: np.ndarray, nbins: int = 100) -> np.ndarray:
N,K,C = poses.shape
assert C==3, poses.shape
poses_21x3 = normalize_axis_angle(torch.fromnumpy(poses).view(N*K,3)).numpy().reshape(N, K, 3)
assert (poses_21x3 > -np.pi).all() and (poses_21x3 < np.pi).all()
Hs, Es = [], []
for i in range(K):
H, edges = np.histogramdd(poses_21x3[:, i, :], bins=nbins, range=[(-np.pi, np.pi)]*3)
Hs.append(H)
Es.append(edges)
Hs = np.stack(Hs, axis=0)
return Hs
def load_amass_hist_smooth(sigma=2) -> torch.Tensor:
amass_poses_hist100 = np.load(AMASS_HIST100_PATH)
amass_poses_hist100 = torch.from_numpy(amass_poses_hist100)
assert amass_poses_hist100.shape == (21,100,100,100)
nbins = amass_poses_hist100.shape[1]
amass_poses_hist100 = amass_poses_hist100/amass_poses_hist100.sum() / (2*np.pi/nbins)**3
# Gaussian filter on amass_poses_hist100
from scipy.ndimage import gaussian_filter
amass_poses_hist100_smooth = gaussian_filter(amass_poses_hist100.numpy(), sigma=sigma, mode='constant')
amass_poses_hist100_smooth = torch.from_numpy(amass_poses_hist100_smooth)
return amass_poses_hist100_smooth
# Normalize axis angle representation s.t. angle is in [-pi, pi]
def normalize_axis_angle(poses: torch.Tensor) -> torch.Tensor:
# poses: N, 3
# print(f'normalize_axis_angle ...')
assert poses.shape[1] == 3, poses.shape
angle = poses.norm(dim=1)
axis = F.normalize(poses, p=2, dim=1, eps=1e-8)
angle_fixed = angle.clone()
axis_fixed = axis.clone()
eps = 1e-6
ii = 0
while True:
# print(f'normalize_axis_angle iter {ii}')
ii += 1
angle_too_big = (angle_fixed > np.pi + eps)
if not angle_too_big.any():
break
angle_fixed[angle_too_big] -= 2 * np.pi
angle_too_small = (angle_fixed < -eps)
axis_fixed[angle_too_small] *= -1
angle_fixed[angle_too_small] *= -1
return axis_fixed * angle_fixed[:,None]
def poses_to_joint_probs(poses: torch.Tensor, amass_poses_100_smooth: torch.Tensor) -> torch.Tensor:
# poses: Nx69
# amass_poses_100_smooth: 21xBINSxBINSxBINS
# returns: poses_prob: Nx21
N=poses.shape[0]
assert poses.shape == (N,69)
poses = poses[:,:63].reshape(N*21,3)
nbins = amass_poses_100_smooth.shape[1]
assert amass_poses_100_smooth.shape == (21,nbins,nbins,nbins)
poses_bin = (poses - POSE_RANGE_MIN) / (POSE_RANGE_MAX - POSE_RANGE_MIN) * (nbins - 1e-6)
poses_bin = poses_bin.long().clip(0, nbins-1)
joint_id = torch.arange(21, device=poses.device).view(1,21).expand(N,21).reshape(N*21)
poses_prob = amass_poses_100_smooth[joint_id, poses_bin[:,0], poses_bin[:,1], poses_bin[:,2]]
poses_bad = ((poses < POSE_RANGE_MIN) | (poses >= POSE_RANGE_MAX)).any(dim=1)
poses_prob[poses_bad] = 0
return poses_prob.view(N,21)
def poses_check_probable(
poses: torch.Tensor,
amass_poses_100_smooth: torch.Tensor,
prob_thresholds: torch.Tensor = JOINT_IDX_PROB_THRESHOLDS
) -> torch.Tensor:
N,C=poses.shape
poses_norm = normalize_axis_angle(poses.reshape(N*(C//3),3)).reshape(N,C)
poses_prob = poses_to_joint_probs(poses_norm, amass_poses_100_smooth)
return (poses_prob > prob_thresholds).all(dim=1)
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,88 @@
from typing import Dict
import cv2
import numpy as np
from skimage.filters import gaussian
from yacs.config import CfgNode
import torch
from .utils import (convert_cvimg_to_tensor,
expand_to_aspect_ratio,
generate_image_patch_cv2)
DEFAULT_MEAN = 255. * np.array([0.485, 0.456, 0.406])
DEFAULT_STD = 255. * np.array([0.229, 0.224, 0.225])
class ViTDetDataset(torch.utils.data.Dataset):
def __init__(self,
cfg: CfgNode,
img_cv2: np.array,
boxes: np.array,
train: bool = False,
**kwargs):
super().__init__()
self.cfg = cfg
self.img_cv2 = img_cv2
# self.boxes = boxes
assert train == False, "ViTDetDataset is only for inference"
self.train = train
self.img_size = cfg.MODEL.IMAGE_SIZE
self.mean = 255. * np.array(self.cfg.MODEL.IMAGE_MEAN)
self.std = 255. * np.array(self.cfg.MODEL.IMAGE_STD)
# Preprocess annotations
boxes = boxes.astype(np.float32)
self.center = (boxes[:, 2:4] + boxes[:, 0:2]) / 2.0
self.scale = (boxes[:, 2:4] - boxes[:, 0:2]) / 200.0
self.personid = np.arange(len(boxes), dtype=np.int32)
def __len__(self) -> int:
return len(self.personid)
def __getitem__(self, idx: int) -> Dict[str, np.array]:
center = self.center[idx].copy()
center_x = center[0]
center_y = center[1]
scale = self.scale[idx]
BBOX_SHAPE = self.cfg.MODEL.get('BBOX_SHAPE', None)
bbox_size = expand_to_aspect_ratio(scale*200, target_aspect_ratio=BBOX_SHAPE).max()
patch_width = patch_height = self.img_size
# 3. generate image patch
# if use_skimage_antialias:
cvimg = self.img_cv2.copy()
if True:
# Blur image to avoid aliasing artifacts
downsampling_factor = ((bbox_size*1.0) / patch_width)
# print(f'{downsampling_factor=}')
downsampling_factor = downsampling_factor / 2.0
if downsampling_factor > 1.1:
cvimg = gaussian(cvimg, sigma=(downsampling_factor-1)/2, channel_axis=2, preserve_range=True)
img_patch_cv, trans = generate_image_patch_cv2(cvimg,
center_x, center_y,
bbox_size, bbox_size,
patch_width, patch_height,
False, 1.0, 0,
border_mode=cv2.BORDER_CONSTANT)
img_patch_cv = img_patch_cv[:, :, ::-1]
img_patch = convert_cvimg_to_tensor(img_patch_cv)
# apply normalization
for n_c in range(min(self.img_cv2.shape[2], 3)):
img_patch[n_c, :, :] = (img_patch[n_c, :, :] - self.mean[n_c]) / self.std[n_c]
item = {
'img': img_patch,
'personid': int(self.personid[idx]),
}
item['box_center'] = self.center[idx].copy()
item['box_size'] = bbox_size
item['img_size'] = 1.0 * np.array([cvimg.shape[1], cvimg.shape[0]])
return item
@@ -0,0 +1,90 @@
from .smpl_wrapper import SMPL
from .hmr2_arch import HMR2
from .discriminator import Discriminator
from ..utils.download import cache_url
from ..configs import CACHE_DIR_4DHUMANS
def download_models(folder=CACHE_DIR_4DHUMANS, extra_filename_links={}):
"""Download checkpoints and files for running inference.
"""
import os
os.makedirs(folder, exist_ok=True)
download_files = {
"model_config.yaml": ["https://huggingface.co/spaces/brjathu/HMR2.0/raw/main/logs/train/multiruns/hmr2/0/model_config.yaml", folder],
"epoch=35-step=1000000.ckpt": ["https://huggingface.co/spaces/brjathu/HMR2.0/resolve/main/logs/train/multiruns/hmr2/0/checkpoints/epoch%3D35-step%3D1000000.ckpt", folder],
"SMPL_NEUTRAL.pkl": ["https://huggingface.co/spaces/brjathu/HMR2.0/resolve/main/data/smpl/SMPL_NEUTRAL.pkl", os.path.join(folder, 'data', 'smpl')],
"SMPL_to_J19.pkl": ["https://huggingface.co/spaces/brjathu/HMR2.0/resolve/main/data/SMPL_to_J19.pkl", os.path.join(folder, 'data')],
"smpl_mean_params.npz": ["https://huggingface.co/spaces/brjathu/HMR2.0/resolve/main/data/smpl_mean_params.npz", os.path.join(folder, 'data')],
**{filename: [link, folder] for filename, link in extra_filename_links.items()}
}
for file_name, url in download_files.items():
output_path = os.path.join(url[1], file_name)
if not os.path.exists(output_path):
log = "smpl" not in file_name.lower()
if log:
print("Downloading file: " + file_name)
# output = gdown.cached_download(url[0], output_path, fuzzy=True)
output = cache_url(url[0], output_path, log=log)
assert os.path.exists(output_path), f"{output} does not exist"
# if ends with tar.gz, tar -xzf
if file_name.endswith(".tar.gz"):
print("Extracting file: " + file_name)
os.system("tar -xvf " + output_path + " -C " + url[1])
def check_smpl_exists():
import os
candidates = [
f'{CACHE_DIR_4DHUMANS}/data/smpl/SMPL_NEUTRAL.pkl'
]
candidates_exist = [os.path.exists(c) for c in candidates]
if not any(candidates_exist):
raise FileNotFoundError(f"SMPL model not found. Please download it from https://smplify.is.tue.mpg.de/ and place it at {candidates[0]}")
# Code edxpects SMPL model at CACHE_DIR_4DHUMANS/data/smpl/SMPL_NEUTRAL.pkl. Copy there if needed
if (not candidates_exist[0]) and candidates_exist[1]:
convert_pkl(candidates[1], candidates[0])
return True
# Convert SMPL pkl file to be compatible with Python 3
# Script is from https://rebeccabilbro.github.io/convert-py2-pickles-to-py3/
def convert_pkl(old_pkl, new_pkl):
"""
Convert a Python 2 pickle to Python 3
"""
import dill
import pickle
# Convert Python 2 "ObjectType" to Python 3 object
dill._dill._reverse_typemap["ObjectType"] = object
# Open the pickle using latin1 encoding
with open(old_pkl, "rb") as f:
loaded = pickle.load(f, encoding="latin1")
# Re-save as Python 3 pickle
with open(new_pkl, "wb") as outfile:
pickle.dump(loaded, outfile)
DEFAULT_CHECKPOINT=f'{CACHE_DIR_4DHUMANS}/epoch=35-step=1000000.ckpt'
def load_hmr2(checkpoint_path=DEFAULT_CHECKPOINT):
from pathlib import Path
from ..configs import get_config
model_cfg = str(Path(checkpoint_path).parent / 'model_config.yaml')
model_cfg = get_config(model_cfg, update_cachedir=True)
# Override some config values, to crop bbox correctly
if (model_cfg.MODEL.BACKBONE.TYPE == 'vit') and ('BBOX_SHAPE' not in model_cfg.MODEL):
model_cfg.defrost()
assert model_cfg.MODEL.IMAGE_SIZE == 256, f"MODEL.IMAGE_SIZE ({model_cfg.MODEL.IMAGE_SIZE}) should be 256 for ViT backbone"
model_cfg.MODEL.BBOX_SHAPE = [192,256]
model_cfg.freeze()
# Ensure SMPL model exists
check_smpl_exists()
model = HMR2.load_from_checkpoint(checkpoint_path, strict=False, cfg=model_cfg)
return model, model_cfg
@@ -0,0 +1,7 @@
from .vit import vit
def create_backbone(cfg):
if cfg.MODEL.BACKBONE.TYPE == 'vit':
return vit(cfg)
else:
raise NotImplementedError('Backbone type is not implemented')
@@ -0,0 +1,348 @@
# Copyright (c) OpenMMLab. All rights reserved.
import math
import torch
from functools import partial
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.checkpoint as checkpoint
from timm.models.layers import drop_path, to_2tuple, trunc_normal_
def vit(cfg):
return ViT(
img_size=(256, 192),
patch_size=16,
embed_dim=1280,
depth=32,
num_heads=16,
ratio=1,
use_checkpoint=False,
mlp_ratio=4,
qkv_bias=True,
drop_path_rate=0.55,
)
def get_abs_pos(abs_pos, h, w, ori_h, ori_w, has_cls_token=True):
"""
Calculate absolute positional embeddings. If needed, resize embeddings and remove cls_token
dimension for the original embeddings.
Args:
abs_pos (Tensor): absolute positional embeddings with (1, num_position, C).
has_cls_token (bool): If true, has 1 embedding in abs_pos for cls token.
hw (Tuple): size of input image tokens.
Returns:
Absolute positional embeddings after processing with shape (1, H, W, C)
"""
cls_token = None
B, L, C = abs_pos.shape
if has_cls_token:
cls_token = abs_pos[:, 0:1]
abs_pos = abs_pos[:, 1:]
if ori_h != h or ori_w != w:
new_abs_pos = F.interpolate(
abs_pos.reshape(1, ori_h, ori_w, -1).permute(0, 3, 1, 2),
size=(h, w),
mode="bicubic",
align_corners=False,
).permute(0, 2, 3, 1).reshape(B, -1, C)
else:
new_abs_pos = abs_pos
if cls_token is not None:
new_abs_pos = torch.cat([cls_token, new_abs_pos], dim=1)
return new_abs_pos
class DropPath(nn.Module):
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
"""
def __init__(self, drop_prob=None):
super(DropPath, self).__init__()
self.drop_prob = drop_prob
def forward(self, x):
return drop_path(x, self.drop_prob, self.training)
def extra_repr(self):
return 'p={}'.format(self.drop_prob)
class Mlp(nn.Module):
def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
super().__init__()
out_features = out_features or in_features
hidden_features = hidden_features or in_features
self.fc1 = nn.Linear(in_features, hidden_features)
self.act = act_layer()
self.fc2 = nn.Linear(hidden_features, out_features)
self.drop = nn.Dropout(drop)
def forward(self, x):
x = self.fc1(x)
x = self.act(x)
x = self.fc2(x)
x = self.drop(x)
return x
class Attention(nn.Module):
def __init__(
self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0.,
proj_drop=0., attn_head_dim=None,):
super().__init__()
self.num_heads = num_heads
head_dim = dim // num_heads
self.dim = dim
if attn_head_dim is not None:
head_dim = attn_head_dim
all_head_dim = head_dim * self.num_heads
self.scale = qk_scale or head_dim ** -0.5
self.qkv = nn.Linear(dim, all_head_dim * 3, bias=qkv_bias)
self.attn_drop = nn.Dropout(attn_drop)
self.proj = nn.Linear(all_head_dim, dim)
self.proj_drop = nn.Dropout(proj_drop)
def forward(self, x):
B, N, C = x.shape
qkv = self.qkv(x)
qkv = qkv.reshape(B, N, 3, self.num_heads, -1).permute(2, 0, 3, 1, 4)
q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)
q = q * self.scale
attn = (q @ k.transpose(-2, -1))
attn = attn.softmax(dim=-1)
attn = self.attn_drop(attn)
x = (attn @ v).transpose(1, 2).reshape(B, N, -1)
x = self.proj(x)
x = self.proj_drop(x)
return x
class Block(nn.Module):
def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None,
drop=0., attn_drop=0., drop_path=0., act_layer=nn.GELU,
norm_layer=nn.LayerNorm, attn_head_dim=None
):
super().__init__()
self.norm1 = norm_layer(dim)
self.attn = Attention(
dim, num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale,
attn_drop=attn_drop, proj_drop=drop, attn_head_dim=attn_head_dim
)
# NOTE: drop path for stochastic depth, we shall see if this is better than dropout here
self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
self.norm2 = norm_layer(dim)
mlp_hidden_dim = int(dim * mlp_ratio)
self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
def forward(self, x):
x = x + self.drop_path(self.attn(self.norm1(x)))
x = x + self.drop_path(self.mlp(self.norm2(x)))
return x
class PatchEmbed(nn.Module):
""" Image to Patch Embedding
"""
def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768, ratio=1):
super().__init__()
img_size = to_2tuple(img_size)
patch_size = to_2tuple(patch_size)
num_patches = (img_size[1] // patch_size[1]) * (img_size[0] // patch_size[0]) * (ratio ** 2)
self.patch_shape = (int(img_size[0] // patch_size[0] * ratio), int(img_size[1] // patch_size[1] * ratio))
self.origin_patch_shape = (int(img_size[0] // patch_size[0]), int(img_size[1] // patch_size[1]))
self.img_size = img_size
self.patch_size = patch_size
self.num_patches = num_patches
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=(patch_size[0] // ratio), padding=4 + 2 * (ratio//2-1))
def forward(self, x, **kwargs):
B, C, H, W = x.shape
x = self.proj(x)
Hp, Wp = x.shape[2], x.shape[3]
x = x.flatten(2).transpose(1, 2)
return x, (Hp, Wp)
class HybridEmbed(nn.Module):
""" CNN Feature Map Embedding
Extract feature map from CNN, flatten, project to embedding dim.
"""
def __init__(self, backbone, img_size=224, feature_size=None, in_chans=3, embed_dim=768):
super().__init__()
assert isinstance(backbone, nn.Module)
img_size = to_2tuple(img_size)
self.img_size = img_size
self.backbone = backbone
if feature_size is None:
with torch.no_grad():
training = backbone.training
if training:
backbone.eval()
o = self.backbone(torch.zeros(1, in_chans, img_size[0], img_size[1]))[-1]
feature_size = o.shape[-2:]
feature_dim = o.shape[1]
backbone.train(training)
else:
feature_size = to_2tuple(feature_size)
feature_dim = self.backbone.feature_info.channels()[-1]
self.num_patches = feature_size[0] * feature_size[1]
self.proj = nn.Linear(feature_dim, embed_dim)
def forward(self, x):
x = self.backbone(x)[-1]
x = x.flatten(2).transpose(1, 2)
x = self.proj(x)
return x
class ViT(nn.Module):
def __init__(self,
img_size=224, patch_size=16, in_chans=3, num_classes=80, embed_dim=768, depth=12,
num_heads=12, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop_rate=0., attn_drop_rate=0.,
drop_path_rate=0., hybrid_backbone=None, norm_layer=None, use_checkpoint=False,
frozen_stages=-1, ratio=1, last_norm=True,
patch_padding='pad', freeze_attn=False, freeze_ffn=False,
):
# Protect mutable default arguments
super(ViT, self).__init__()
norm_layer = norm_layer or partial(nn.LayerNorm, eps=1e-6)
self.num_classes = num_classes
self.num_features = self.embed_dim = embed_dim # num_features for consistency with other models
self.frozen_stages = frozen_stages
self.use_checkpoint = use_checkpoint
self.patch_padding = patch_padding
self.freeze_attn = freeze_attn
self.freeze_ffn = freeze_ffn
self.depth = depth
if hybrid_backbone is not None:
self.patch_embed = HybridEmbed(
hybrid_backbone, img_size=img_size, in_chans=in_chans, embed_dim=embed_dim)
else:
self.patch_embed = PatchEmbed(
img_size=img_size, patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim, ratio=ratio)
num_patches = self.patch_embed.num_patches
# since the pretraining model has class token
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)] # stochastic depth decay rule
self.blocks = nn.ModuleList([
Block(
dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, qk_scale=qk_scale,
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[i], norm_layer=norm_layer,
)
for i in range(depth)])
self.last_norm = norm_layer(embed_dim) if last_norm else nn.Identity()
if self.pos_embed is not None:
trunc_normal_(self.pos_embed, std=.02)
self._freeze_stages()
def _freeze_stages(self):
"""Freeze parameters."""
if self.frozen_stages >= 0:
self.patch_embed.eval()
for param in self.patch_embed.parameters():
param.requires_grad = False
for i in range(1, self.frozen_stages + 1):
m = self.blocks[i]
m.eval()
for param in m.parameters():
param.requires_grad = False
if self.freeze_attn:
for i in range(0, self.depth):
m = self.blocks[i]
m.attn.eval()
m.norm1.eval()
for param in m.attn.parameters():
param.requires_grad = False
for param in m.norm1.parameters():
param.requires_grad = False
if self.freeze_ffn:
self.pos_embed.requires_grad = False
self.patch_embed.eval()
for param in self.patch_embed.parameters():
param.requires_grad = False
for i in range(0, self.depth):
m = self.blocks[i]
m.mlp.eval()
m.norm2.eval()
for param in m.mlp.parameters():
param.requires_grad = False
for param in m.norm2.parameters():
param.requires_grad = False
def init_weights(self):
"""Initialize the weights in backbone.
Args:
pretrained (str, optional): Path to pre-trained weights.
Defaults to None.
"""
def _init_weights(m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=.02)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
self.apply(_init_weights)
def get_num_layers(self):
return len(self.blocks)
@torch.jit.ignore
def no_weight_decay(self):
return {'pos_embed', 'cls_token'}
def forward_features(self, x):
B, C, H, W = x.shape
x, (Hp, Wp) = self.patch_embed(x)
if self.pos_embed is not None:
# fit for multiple GPU training
# since the first element for pos embed (sin-cos manner) is zero, it will cause no difference
x = x + self.pos_embed[:, 1:] + self.pos_embed[:, :1]
for blk in self.blocks:
if self.use_checkpoint:
x = checkpoint.checkpoint(blk, x)
else:
x = blk(x)
x = self.last_norm(x)
xp = x.permute(0, 2, 1).reshape(B, -1, Hp, Wp).contiguous()
return xp
def forward(self, x):
x = self.forward_features(x)
return x
def train(self, mode=True):
"""Convert the model into training mode."""
super().train(mode)
self._freeze_stages()
@@ -0,0 +1,17 @@
# import mmcv
# import mmpose
# from mmpose.models import build_posenet
# from mmcv.runner import load_checkpoint
# from pathlib import Path
# def vit(cfg):
# vitpose_dir = Path(mmpose.__file__).parent.parent
# config = f'{vitpose_dir}/configs/body/2d_kpt_sview_rgb_img/topdown_heatmap/coco/ViTPose_huge_coco_256x192.py'
# # checkpoint = f'{vitpose_dir}/models/vitpose-h-multi-coco.pth'
# config = mmcv.Config.fromfile(config)
# config.model.pretrained = None
# model = build_posenet(config.model)
# # load_checkpoint(model, checkpoint, map_location='cpu')
# return model.backbone
@@ -0,0 +1,358 @@
from inspect import isfunction
from typing import Callable, Optional
import torch
from einops import rearrange
from einops.layers.torch import Rearrange
from torch import nn
from .t_cond_mlp import (
AdaptiveLayerNorm1D,
FrequencyEmbedder,
normalization_layer,
)
# from .vit import Attention, FeedForward
def exists(val):
return val is not None
def default(val, d):
if exists(val):
return val
return d() if isfunction(d) else d
class PreNorm(nn.Module):
def __init__(self, dim: int, fn: Callable, norm: str = "layer", norm_cond_dim: int = -1):
super().__init__()
self.norm = normalization_layer(norm, dim, norm_cond_dim)
self.fn = fn
def forward(self, x: torch.Tensor, *args, **kwargs):
if isinstance(self.norm, AdaptiveLayerNorm1D):
return self.fn(self.norm(x, *args), **kwargs)
else:
return self.fn(self.norm(x), **kwargs)
class FeedForward(nn.Module):
def __init__(self, dim, hidden_dim, dropout=0.0):
super().__init__()
self.net = nn.Sequential(
nn.Linear(dim, hidden_dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(hidden_dim, dim),
nn.Dropout(dropout),
)
def forward(self, x):
return self.net(x)
class Attention(nn.Module):
def __init__(self, dim, heads=8, dim_head=64, dropout=0.0):
super().__init__()
inner_dim = dim_head * heads
project_out = not (heads == 1 and dim_head == dim)
self.heads = heads
self.scale = dim_head**-0.5
self.attend = nn.Softmax(dim=-1)
self.dropout = nn.Dropout(dropout)
self.to_qkv = nn.Linear(dim, inner_dim * 3, bias=False)
self.to_out = (
nn.Sequential(nn.Linear(inner_dim, dim), nn.Dropout(dropout))
if project_out
else nn.Identity()
)
def forward(self, x):
qkv = self.to_qkv(x).chunk(3, dim=-1)
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=self.heads), qkv)
dots = torch.matmul(q, k.transpose(-1, -2)) * self.scale
attn = self.attend(dots)
attn = self.dropout(attn)
out = torch.matmul(attn, v)
out = rearrange(out, "b h n d -> b n (h d)")
return self.to_out(out)
class CrossAttention(nn.Module):
def __init__(self, dim, context_dim=None, heads=8, dim_head=64, dropout=0.0):
super().__init__()
inner_dim = dim_head * heads
project_out = not (heads == 1 and dim_head == dim)
self.heads = heads
self.scale = dim_head**-0.5
self.attend = nn.Softmax(dim=-1)
self.dropout = nn.Dropout(dropout)
context_dim = default(context_dim, dim)
self.to_kv = nn.Linear(context_dim, inner_dim * 2, bias=False)
self.to_q = nn.Linear(dim, inner_dim, bias=False)
self.to_out = (
nn.Sequential(nn.Linear(inner_dim, dim), nn.Dropout(dropout))
if project_out
else nn.Identity()
)
def forward(self, x, context=None):
context = default(context, x)
k, v = self.to_kv(context).chunk(2, dim=-1)
q = self.to_q(x)
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=self.heads), [q, k, v])
dots = torch.matmul(q, k.transpose(-1, -2)) * self.scale
attn = self.attend(dots)
attn = self.dropout(attn)
out = torch.matmul(attn, v)
out = rearrange(out, "b h n d -> b n (h d)")
return self.to_out(out)
class Transformer(nn.Module):
def __init__(
self,
dim: int,
depth: int,
heads: int,
dim_head: int,
mlp_dim: int,
dropout: float = 0.0,
norm: str = "layer",
norm_cond_dim: int = -1,
):
super().__init__()
self.layers = nn.ModuleList([])
for _ in range(depth):
sa = Attention(dim, heads=heads, dim_head=dim_head, dropout=dropout)
ff = FeedForward(dim, mlp_dim, dropout=dropout)
self.layers.append(
nn.ModuleList(
[
PreNorm(dim, sa, norm=norm, norm_cond_dim=norm_cond_dim),
PreNorm(dim, ff, norm=norm, norm_cond_dim=norm_cond_dim),
]
)
)
def forward(self, x: torch.Tensor, *args):
for attn, ff in self.layers:
x = attn(x, *args) + x
x = ff(x, *args) + x
return x
class TransformerCrossAttn(nn.Module):
def __init__(
self,
dim: int,
depth: int,
heads: int,
dim_head: int,
mlp_dim: int,
dropout: float = 0.0,
norm: str = "layer",
norm_cond_dim: int = -1,
context_dim: Optional[int] = None,
):
super().__init__()
self.layers = nn.ModuleList([])
for _ in range(depth):
sa = Attention(dim, heads=heads, dim_head=dim_head, dropout=dropout)
ca = CrossAttention(
dim, context_dim=context_dim, heads=heads, dim_head=dim_head, dropout=dropout
)
ff = FeedForward(dim, mlp_dim, dropout=dropout)
self.layers.append(
nn.ModuleList(
[
PreNorm(dim, sa, norm=norm, norm_cond_dim=norm_cond_dim),
PreNorm(dim, ca, norm=norm, norm_cond_dim=norm_cond_dim),
PreNorm(dim, ff, norm=norm, norm_cond_dim=norm_cond_dim),
]
)
)
def forward(self, x: torch.Tensor, *args, context=None, context_list=None):
if context_list is None:
context_list = [context] * len(self.layers)
if len(context_list) != len(self.layers):
raise ValueError(f"len(context_list) != len(self.layers) ({len(context_list)} != {len(self.layers)})")
for i, (self_attn, cross_attn, ff) in enumerate(self.layers):
x = self_attn(x, *args) + x
x = cross_attn(x, *args, context=context_list[i]) + x
x = ff(x, *args) + x
return x
class DropTokenDropout(nn.Module):
def __init__(self, p: float = 0.1):
super().__init__()
if p < 0 or p > 1:
raise ValueError(
"dropout probability has to be between 0 and 1, " "but got {}".format(p)
)
self.p = p
def forward(self, x: torch.Tensor):
# x: (batch_size, seq_len, dim)
if self.training and self.p > 0:
zero_mask = torch.full_like(x[0, :, 0], self.p).bernoulli().bool()
# TODO: permutation idx for each batch using torch.argsort
if zero_mask.any():
x = x[:, ~zero_mask, :]
return x
class ZeroTokenDropout(nn.Module):
def __init__(self, p: float = 0.1):
super().__init__()
if p < 0 or p > 1:
raise ValueError(
"dropout probability has to be between 0 and 1, " "but got {}".format(p)
)
self.p = p
def forward(self, x: torch.Tensor):
# x: (batch_size, seq_len, dim)
if self.training and self.p > 0:
zero_mask = torch.full_like(x[:, :, 0], self.p).bernoulli().bool()
# Zero-out the masked tokens
x[zero_mask, :] = 0
return x
class TransformerEncoder(nn.Module):
def __init__(
self,
num_tokens: int,
token_dim: int,
dim: int,
depth: int,
heads: int,
mlp_dim: int,
dim_head: int = 64,
dropout: float = 0.0,
emb_dropout: float = 0.0,
emb_dropout_type: str = "drop",
emb_dropout_loc: str = "token",
norm: str = "layer",
norm_cond_dim: int = -1,
token_pe_numfreq: int = -1,
):
super().__init__()
if token_pe_numfreq > 0:
token_dim_new = token_dim * (2 * token_pe_numfreq + 1)
self.to_token_embedding = nn.Sequential(
Rearrange("b n d -> (b n) d", n=num_tokens, d=token_dim),
FrequencyEmbedder(token_pe_numfreq, token_pe_numfreq - 1),
Rearrange("(b n) d -> b n d", n=num_tokens, d=token_dim_new),
nn.Linear(token_dim_new, dim),
)
else:
self.to_token_embedding = nn.Linear(token_dim, dim)
self.pos_embedding = nn.Parameter(torch.randn(1, num_tokens, dim))
if emb_dropout_type == "drop":
self.dropout = DropTokenDropout(emb_dropout)
elif emb_dropout_type == "zero":
self.dropout = ZeroTokenDropout(emb_dropout)
else:
raise ValueError(f"Unknown emb_dropout_type: {emb_dropout_type}")
self.emb_dropout_loc = emb_dropout_loc
self.transformer = Transformer(
dim, depth, heads, dim_head, mlp_dim, dropout, norm=norm, norm_cond_dim=norm_cond_dim
)
def forward(self, inp: torch.Tensor, *args, **kwargs):
x = inp
if self.emb_dropout_loc == "input":
x = self.dropout(x)
x = self.to_token_embedding(x)
if self.emb_dropout_loc == "token":
x = self.dropout(x)
b, n, _ = x.shape
x += self.pos_embedding[:, :n]
if self.emb_dropout_loc == "token_afterpos":
x = self.dropout(x)
x = self.transformer(x, *args)
return x
class TransformerDecoder(nn.Module):
def __init__(
self,
num_tokens: int,
token_dim: int,
dim: int,
depth: int,
heads: int,
mlp_dim: int,
dim_head: int = 64,
dropout: float = 0.0,
emb_dropout: float = 0.0,
emb_dropout_type: str = 'drop',
norm: str = "layer",
norm_cond_dim: int = -1,
context_dim: Optional[int] = None,
skip_token_embedding: bool = False,
):
super().__init__()
if not skip_token_embedding:
self.to_token_embedding = nn.Linear(token_dim, dim)
else:
self.to_token_embedding = nn.Identity()
if token_dim != dim:
raise ValueError(
f"token_dim ({token_dim}) != dim ({dim}) when skip_token_embedding is True"
)
self.pos_embedding = nn.Parameter(torch.randn(1, num_tokens, dim))
if emb_dropout_type == "drop":
self.dropout = DropTokenDropout(emb_dropout)
elif emb_dropout_type == "zero":
self.dropout = ZeroTokenDropout(emb_dropout)
elif emb_dropout_type == "normal":
self.dropout = nn.Dropout(emb_dropout)
self.transformer = TransformerCrossAttn(
dim,
depth,
heads,
dim_head,
mlp_dim,
dropout,
norm=norm,
norm_cond_dim=norm_cond_dim,
context_dim=context_dim,
)
def forward(self, inp: torch.Tensor, *args, context=None, context_list=None):
x = self.to_token_embedding(inp)
b, n, _ = x.shape
x = self.dropout(x)
x += self.pos_embedding[:, :n]
x = self.transformer(x, *args, context=context, context_list=context_list)
return x
@@ -0,0 +1,199 @@
import copy
from typing import List, Optional
import torch
class AdaptiveLayerNorm1D(torch.nn.Module):
def __init__(self, data_dim: int, norm_cond_dim: int):
super().__init__()
if data_dim <= 0:
raise ValueError(f"data_dim must be positive, but got {data_dim}")
if norm_cond_dim <= 0:
raise ValueError(f"norm_cond_dim must be positive, but got {norm_cond_dim}")
self.norm = torch.nn.LayerNorm(
data_dim
) # TODO: Check if elementwise_affine=True is correct
self.linear = torch.nn.Linear(norm_cond_dim, 2 * data_dim)
torch.nn.init.zeros_(self.linear.weight)
torch.nn.init.zeros_(self.linear.bias)
def forward(self, x: torch.Tensor, t: torch.Tensor) -> torch.Tensor:
# x: (batch, ..., data_dim)
# t: (batch, norm_cond_dim)
# return: (batch, data_dim)
x = self.norm(x)
alpha, beta = self.linear(t).chunk(2, dim=-1)
# Add singleton dimensions to alpha and beta
if x.dim() > 2:
alpha = alpha.view(alpha.shape[0], *([1] * (x.dim() - 2)), alpha.shape[1])
beta = beta.view(beta.shape[0], *([1] * (x.dim() - 2)), beta.shape[1])
return x * (1 + alpha) + beta
class SequentialCond(torch.nn.Sequential):
def forward(self, input, *args, **kwargs):
for module in self:
if isinstance(module, (AdaptiveLayerNorm1D, SequentialCond, ResidualMLPBlock)):
# print(f'Passing on args to {module}', [a.shape for a in args])
input = module(input, *args, **kwargs)
else:
# print(f'Skipping passing args to {module}', [a.shape for a in args])
input = module(input)
return input
def normalization_layer(norm: Optional[str], dim: int, norm_cond_dim: int = -1):
if norm == "batch":
return torch.nn.BatchNorm1d(dim)
elif norm == "layer":
return torch.nn.LayerNorm(dim)
elif norm == "ada":
assert norm_cond_dim > 0, f"norm_cond_dim must be positive, got {norm_cond_dim}"
return AdaptiveLayerNorm1D(dim, norm_cond_dim)
elif norm is None:
return torch.nn.Identity()
else:
raise ValueError(f"Unknown norm: {norm}")
def linear_norm_activ_dropout(
input_dim: int,
output_dim: int,
activation: torch.nn.Module = torch.nn.ReLU(),
bias: bool = True,
norm: Optional[str] = "layer", # Options: ada/batch/layer
dropout: float = 0.0,
norm_cond_dim: int = -1,
) -> SequentialCond:
layers = []
layers.append(torch.nn.Linear(input_dim, output_dim, bias=bias))
if norm is not None:
layers.append(normalization_layer(norm, output_dim, norm_cond_dim))
layers.append(copy.deepcopy(activation))
if dropout > 0.0:
layers.append(torch.nn.Dropout(dropout))
return SequentialCond(*layers)
def create_simple_mlp(
input_dim: int,
hidden_dims: List[int],
output_dim: int,
activation: torch.nn.Module = torch.nn.ReLU(),
bias: bool = True,
norm: Optional[str] = "layer", # Options: ada/batch/layer
dropout: float = 0.0,
norm_cond_dim: int = -1,
) -> SequentialCond:
layers = []
prev_dim = input_dim
for hidden_dim in hidden_dims:
layers.extend(
linear_norm_activ_dropout(
prev_dim, hidden_dim, activation, bias, norm, dropout, norm_cond_dim
)
)
prev_dim = hidden_dim
layers.append(torch.nn.Linear(prev_dim, output_dim, bias=bias))
return SequentialCond(*layers)
class ResidualMLPBlock(torch.nn.Module):
def __init__(
self,
input_dim: int,
hidden_dim: int,
num_hidden_layers: int,
output_dim: int,
activation: torch.nn.Module = torch.nn.ReLU(),
bias: bool = True,
norm: Optional[str] = "layer", # Options: ada/batch/layer
dropout: float = 0.0,
norm_cond_dim: int = -1,
):
super().__init__()
if not (input_dim == output_dim == hidden_dim):
raise NotImplementedError(
f"input_dim {input_dim} != output_dim {output_dim} is not implemented"
)
layers = []
prev_dim = input_dim
for i in range(num_hidden_layers):
layers.append(
linear_norm_activ_dropout(
prev_dim, hidden_dim, activation, bias, norm, dropout, norm_cond_dim
)
)
prev_dim = hidden_dim
self.model = SequentialCond(*layers)
self.skip = torch.nn.Identity()
def forward(self, x: torch.Tensor, *args, **kwargs) -> torch.Tensor:
return x + self.model(x, *args, **kwargs)
class ResidualMLP(torch.nn.Module):
def __init__(
self,
input_dim: int,
hidden_dim: int,
num_hidden_layers: int,
output_dim: int,
activation: torch.nn.Module = torch.nn.ReLU(),
bias: bool = True,
norm: Optional[str] = "layer", # Options: ada/batch/layer
dropout: float = 0.0,
num_blocks: int = 1,
norm_cond_dim: int = -1,
):
super().__init__()
self.input_dim = input_dim
self.model = SequentialCond(
linear_norm_activ_dropout(
input_dim, hidden_dim, activation, bias, norm, dropout, norm_cond_dim
),
*[
ResidualMLPBlock(
hidden_dim,
hidden_dim,
num_hidden_layers,
hidden_dim,
activation,
bias,
norm,
dropout,
norm_cond_dim,
)
for _ in range(num_blocks)
],
torch.nn.Linear(hidden_dim, output_dim, bias=bias),
)
def forward(self, x: torch.Tensor, *args, **kwargs) -> torch.Tensor:
return self.model(x, *args, **kwargs)
class FrequencyEmbedder(torch.nn.Module):
def __init__(self, num_frequencies, max_freq_log2):
super().__init__()
frequencies = 2 ** torch.linspace(0, max_freq_log2, steps=num_frequencies)
self.register_buffer("frequencies", frequencies)
def forward(self, x):
# x should be of size (N,) or (N, D)
N = x.size(0)
if x.dim() == 1: # (N,)
x = x.unsqueeze(1) # (N, D) where D=1
x_unsqueezed = x.unsqueeze(-1) # (N, D, 1)
scaled = self.frequencies.view(1, 1, -1) * x_unsqueezed # (N, D, num_frequencies)
s = torch.sin(scaled)
c = torch.cos(scaled)
embedded = torch.cat([s, c, x_unsqueezed], dim=-1).view(
N, -1
) # (N, D * 2 * num_frequencies + D)
return embedded
@@ -0,0 +1,99 @@
import torch
import torch.nn as nn
class Discriminator(nn.Module):
def __init__(self):
"""
Pose + Shape discriminator proposed in HMR
"""
super(Discriminator, self).__init__()
self.num_joints = 23
# poses_alone
self.D_conv1 = nn.Conv2d(9, 32, kernel_size=1)
nn.init.xavier_uniform_(self.D_conv1.weight)
nn.init.zeros_(self.D_conv1.bias)
self.relu = nn.ReLU(inplace=True)
self.D_conv2 = nn.Conv2d(32, 32, kernel_size=1)
nn.init.xavier_uniform_(self.D_conv2.weight)
nn.init.zeros_(self.D_conv2.bias)
pose_out = []
for i in range(self.num_joints):
pose_out_temp = nn.Linear(32, 1)
nn.init.xavier_uniform_(pose_out_temp.weight)
nn.init.zeros_(pose_out_temp.bias)
pose_out.append(pose_out_temp)
self.pose_out = nn.ModuleList(pose_out)
# betas
self.betas_fc1 = nn.Linear(10, 10)
nn.init.xavier_uniform_(self.betas_fc1.weight)
nn.init.zeros_(self.betas_fc1.bias)
self.betas_fc2 = nn.Linear(10, 5)
nn.init.xavier_uniform_(self.betas_fc2.weight)
nn.init.zeros_(self.betas_fc2.bias)
self.betas_out = nn.Linear(5, 1)
nn.init.xavier_uniform_(self.betas_out.weight)
nn.init.zeros_(self.betas_out.bias)
# poses_joint
self.D_alljoints_fc1 = nn.Linear(32*self.num_joints, 1024)
nn.init.xavier_uniform_(self.D_alljoints_fc1.weight)
nn.init.zeros_(self.D_alljoints_fc1.bias)
self.D_alljoints_fc2 = nn.Linear(1024, 1024)
nn.init.xavier_uniform_(self.D_alljoints_fc2.weight)
nn.init.zeros_(self.D_alljoints_fc2.bias)
self.D_alljoints_out = nn.Linear(1024, 1)
nn.init.xavier_uniform_(self.D_alljoints_out.weight)
nn.init.zeros_(self.D_alljoints_out.bias)
def forward(self, poses: torch.Tensor, betas: torch.Tensor) -> torch.Tensor:
"""
Forward pass of the discriminator.
Args:
poses (torch.Tensor): Tensor of shape (B, 23, 3, 3) containing a batch of SMPL body poses (excluding the global orientation).
betas (torch.Tensor): Tensor of shape (B, 10) containign a batch of SMPL beta coefficients.
Returns:
torch.Tensor: Discriminator output with shape (B, 25)
"""
#import ipdb; ipdb.set_trace()
#bn = poses.shape[0]
# poses B x 207
#poses = poses.reshape(bn, -1)
# poses B x num_joints x 1 x 9
poses = poses.reshape(-1, self.num_joints, 1, 9)
bn = poses.shape[0]
# poses B x 9 x num_joints x 1
poses = poses.permute(0, 3, 1, 2).contiguous()
# poses_alone
poses = self.D_conv1(poses)
poses = self.relu(poses)
poses = self.D_conv2(poses)
poses = self.relu(poses)
poses_out = []
for i in range(self.num_joints):
poses_out_ = self.pose_out[i](poses[:, :, i, 0])
poses_out.append(poses_out_)
poses_out = torch.cat(poses_out, dim=1)
# betas
betas = self.betas_fc1(betas)
betas = self.relu(betas)
betas = self.betas_fc2(betas)
betas = self.relu(betas)
betas_out = self.betas_out(betas)
# poses_joint
poses = poses.reshape(bn,-1)
poses_all = self.D_alljoints_fc1(poses)
poses_all = self.relu(poses_all)
poses_all = self.D_alljoints_fc2(poses_all)
poses_all = self.relu(poses_all)
poses_all_out = self.D_alljoints_out(poses_all)
disc_out = torch.cat((poses_out, betas_out, poses_all_out), 1)
return disc_out
@@ -0,0 +1 @@
from .smpl_head import build_smpl_head
@@ -0,0 +1,111 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import einops
from ...utils.geometry import rot6d_to_rotmat, aa_to_rotmat
from ..components.pose_transformer import TransformerDecoder
def build_smpl_head(cfg):
smpl_head_type = cfg.MODEL.SMPL_HEAD.get('TYPE', 'hmr')
if smpl_head_type == 'transformer_decoder':
return SMPLTransformerDecoderHead(cfg)
else:
raise ValueError('Unknown SMPL head type: {}'.format(smpl_head_type))
class SMPLTransformerDecoderHead(nn.Module):
""" Cross-attention based SMPL Transformer decoder
"""
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
self.joint_rep_type = cfg.MODEL.SMPL_HEAD.get('JOINT_REP', '6d')
self.joint_rep_dim = {'6d': 6, 'aa': 3}[self.joint_rep_type]
npose = self.joint_rep_dim * (cfg.SMPL.NUM_BODY_JOINTS + 1)
self.npose = npose
self.input_is_mean_shape = cfg.MODEL.SMPL_HEAD.get('TRANSFORMER_INPUT', 'zero') == 'mean_shape'
transformer_args = dict(
num_tokens=1,
token_dim=(npose + 10 + 3) if self.input_is_mean_shape else 1,
dim=1024,
)
transformer_args = (transformer_args | dict(cfg.MODEL.SMPL_HEAD.TRANSFORMER_DECODER))
self.transformer = TransformerDecoder(
**transformer_args
)
dim=transformer_args['dim']
self.decpose = nn.Linear(dim, npose)
self.decshape = nn.Linear(dim, 10)
self.deccam = nn.Linear(dim, 3)
if cfg.MODEL.SMPL_HEAD.get('INIT_DECODER_XAVIER', False):
# True by default in MLP. False by default in Transformer
nn.init.xavier_uniform_(self.decpose.weight, gain=0.01)
nn.init.xavier_uniform_(self.decshape.weight, gain=0.01)
nn.init.xavier_uniform_(self.deccam.weight, gain=0.01)
mean_params = np.load(cfg.SMPL.MEAN_PARAMS)
init_body_pose = torch.from_numpy(mean_params['pose'].astype(np.float32)).unsqueeze(0)
init_betas = torch.from_numpy(mean_params['shape'].astype('float32')).unsqueeze(0)
init_cam = torch.from_numpy(mean_params['cam'].astype(np.float32)).unsqueeze(0)
self.register_buffer('init_body_pose', init_body_pose)
self.register_buffer('init_betas', init_betas)
self.register_buffer('init_cam', init_cam)
def forward(self, x, **kwargs):
batch_size = x.shape[0]
# vit pretrained backbone is channel-first. Change to token-first
x = einops.rearrange(x, 'b c h w -> b (h w) c')
init_body_pose = self.init_body_pose.expand(batch_size, -1)
init_betas = self.init_betas.expand(batch_size, -1)
init_cam = self.init_cam.expand(batch_size, -1)
# TODO: Convert init_body_pose to aa rep if needed
if self.joint_rep_type == 'aa':
raise NotImplementedError
pred_body_pose = init_body_pose
pred_betas = init_betas
pred_cam = init_cam
pred_body_pose_list = []
pred_betas_list = []
pred_cam_list = []
for i in range(self.cfg.MODEL.SMPL_HEAD.get('IEF_ITERS', 1)):
# Input token to transformer is zero token
if self.input_is_mean_shape:
token = torch.cat([pred_body_pose, pred_betas, pred_cam], dim=1)[:,None,:]
else:
token = torch.zeros(batch_size, 1, 1).to(x.device)
# Pass through transformer
token_out = self.transformer(token, context=x)
token_out = token_out.squeeze(1) # (B, C)
# Readout from token_out
pred_body_pose = self.decpose(token_out) + pred_body_pose
pred_betas = self.decshape(token_out) + pred_betas
pred_cam = self.deccam(token_out) + pred_cam
pred_body_pose_list.append(pred_body_pose)
pred_betas_list.append(pred_betas)
pred_cam_list.append(pred_cam)
# Convert self.joint_rep_type -> rotmat
joint_conversion_fn = {
'6d': rot6d_to_rotmat,
'aa': lambda x: aa_to_rotmat(x.view(-1, 3).contiguous())
}[self.joint_rep_type]
pred_smpl_params_list = {}
pred_smpl_params_list['body_pose'] = torch.cat([joint_conversion_fn(pbp).view(batch_size, -1, 3, 3)[:, 1:, :, :] for pbp in pred_body_pose_list], dim=0)
pred_smpl_params_list['betas'] = torch.cat(pred_betas_list, dim=0)
pred_smpl_params_list['cam'] = torch.cat(pred_cam_list, dim=0)
pred_body_pose = joint_conversion_fn(pred_body_pose).view(batch_size, self.cfg.SMPL.NUM_BODY_JOINTS+1, 3, 3)
pred_smpl_params = {'global_orient': pred_body_pose[:, [0]],
'body_pose': pred_body_pose[:, 1:],
'betas': pred_betas}
return pred_smpl_params, pred_cam, pred_smpl_params_list
+355
View File
@@ -0,0 +1,355 @@
import torch
import pytorch_lightning as pl
from typing import Any, Dict, Mapping, Tuple
from yacs.config import CfgNode
from ..utils import SkeletonRenderer, MeshRenderer
from ..utils.geometry import aa_to_rotmat, perspective_projection
from ..utils.pylogger import get_pylogger
from .backbones import create_backbone
from .heads import build_smpl_head
from .discriminator import Discriminator
from .losses import Keypoint3DLoss, Keypoint2DLoss, ParameterLoss
from . import SMPL
log = get_pylogger(__name__)
class HMR2(pl.LightningModule):
def __init__(self, cfg: CfgNode, init_renderer: bool = False):
"""
Setup HMR2 model
Args:
cfg (CfgNode): Config file as a yacs CfgNode
"""
super().__init__()
# Save hyperparameters
self.save_hyperparameters(logger=False, ignore=['init_renderer'])
self.cfg = cfg
# Create backbone feature extractor
self.backbone = create_backbone(cfg)
if cfg.MODEL.BACKBONE.get('PRETRAINED_WEIGHTS', None):
log.info(f'Loading backbone weights from {cfg.MODEL.BACKBONE.PRETRAINED_WEIGHTS}')
self.backbone.load_state_dict(torch.load(cfg.MODEL.BACKBONE.PRETRAINED_WEIGHTS, map_location='cpu')['state_dict'])
# Create SMPL head
self.smpl_head = build_smpl_head(cfg)
# Create discriminator
if self.cfg.LOSS_WEIGHTS.ADVERSARIAL > 0:
self.discriminator = Discriminator()
# Define loss functions
self.keypoint_3d_loss = Keypoint3DLoss(loss_type='l1')
self.keypoint_2d_loss = Keypoint2DLoss(loss_type='l1')
self.smpl_parameter_loss = ParameterLoss()
# Instantiate SMPL model
smpl_cfg = {k.lower(): v for k,v in dict(cfg.SMPL).items()}
self.smpl = SMPL(**smpl_cfg)
# Buffer that shows whetheer we need to initialize ActNorm layers
self.register_buffer('initialized', torch.tensor(False))
# Setup renderer for visualization
if init_renderer:
self.renderer = SkeletonRenderer(self.cfg)
self.mesh_renderer = MeshRenderer(self.cfg, faces=self.smpl.faces)
else:
self.renderer = None
self.mesh_renderer = None
# Disable automatic optimization since we use adversarial training
self.automatic_optimization = False
def get_parameters(self):
all_params = list(self.smpl_head.parameters())
all_params += list(self.backbone.parameters())
return all_params
def configure_optimizers(self) -> Tuple[torch.optim.Optimizer, torch.optim.Optimizer]:
"""
Setup model and distriminator Optimizers
Returns:
Tuple[torch.optim.Optimizer, torch.optim.Optimizer]: Model and discriminator optimizers
"""
param_groups = [{'params': filter(lambda p: p.requires_grad, self.get_parameters()), 'lr': self.cfg.TRAIN.LR}]
optimizer = torch.optim.AdamW(params=param_groups,
# lr=self.cfg.TRAIN.LR,
weight_decay=self.cfg.TRAIN.WEIGHT_DECAY)
optimizer_disc = torch.optim.AdamW(params=self.discriminator.parameters(),
lr=self.cfg.TRAIN.LR,
weight_decay=self.cfg.TRAIN.WEIGHT_DECAY)
return optimizer, optimizer_disc
def forward_step(self, batch: Dict, train: bool = False) -> Dict:
"""
Run a forward step of the network
Args:
batch (Dict): Dictionary containing batch data
train (bool): Flag indicating whether it is training or validation mode
Returns:
Dict: Dictionary containing the regression output
"""
# Use RGB image as input
x = batch['img']
batch_size = x.shape[0]
# Compute conditioning features using the backbone
# if using ViT backbone, we need to use a different aspect ratio
conditioning_feats = self.backbone(x[:,:,:,32:-32])
pred_smpl_params, pred_cam, _ = self.smpl_head(conditioning_feats)
# Store useful regression outputs to the output dict
output = {}
output['pred_cam'] = pred_cam
output['pred_smpl_params'] = {k: v.clone() for k,v in pred_smpl_params.items()}
# Compute camera translation
device = pred_smpl_params['body_pose'].device
dtype = pred_smpl_params['body_pose'].dtype
focal_length = self.cfg.EXTRA.FOCAL_LENGTH * torch.ones(batch_size, 2, device=device, dtype=dtype)
pred_cam_t = torch.stack([pred_cam[:, 1],
pred_cam[:, 2],
2*focal_length[:, 0]/(self.cfg.MODEL.IMAGE_SIZE * pred_cam[:, 0] +1e-9)],dim=-1)
output['pred_cam_t'] = pred_cam_t
output['focal_length'] = focal_length
# Compute model vertices, joints and the projected joints
pred_smpl_params['global_orient'] = pred_smpl_params['global_orient'].reshape(batch_size, -1, 3, 3)
pred_smpl_params['body_pose'] = pred_smpl_params['body_pose'].reshape(batch_size, -1, 3, 3)
pred_smpl_params['betas'] = pred_smpl_params['betas'].reshape(batch_size, -1)
smpl_output = self.smpl(**{k: v.float() for k,v in pred_smpl_params.items()}, pose2rot=False)
pred_keypoints_3d = smpl_output.joints
pred_vertices = smpl_output.vertices
output['pred_keypoints_3d'] = pred_keypoints_3d.reshape(batch_size, -1, 3)
output['pred_vertices'] = pred_vertices.reshape(batch_size, -1, 3)
pred_cam_t = pred_cam_t.reshape(-1, 3)
focal_length = focal_length.reshape(-1, 2)
pred_keypoints_2d = perspective_projection(pred_keypoints_3d,
translation=pred_cam_t,
focal_length=focal_length / self.cfg.MODEL.IMAGE_SIZE)
output['pred_keypoints_2d'] = pred_keypoints_2d.reshape(batch_size, -1, 2)
return output
def compute_loss(self, batch: Dict, output: Dict, train: bool = True) -> torch.Tensor:
"""
Compute losses given the input batch and the regression output
Args:
batch (Dict): Dictionary containing batch data
output (Dict): Dictionary containing the regression output
train (bool): Flag indicating whether it is training or validation mode
Returns:
torch.Tensor : Total loss for current batch
"""
pred_smpl_params = output['pred_smpl_params']
pred_keypoints_2d = output['pred_keypoints_2d']
pred_keypoints_3d = output['pred_keypoints_3d']
batch_size = pred_smpl_params['body_pose'].shape[0]
device = pred_smpl_params['body_pose'].device
dtype = pred_smpl_params['body_pose'].dtype
# Get annotations
gt_keypoints_2d = batch['keypoints_2d']
gt_keypoints_3d = batch['keypoints_3d']
gt_smpl_params = batch['smpl_params']
has_smpl_params = batch['has_smpl_params']
is_axis_angle = batch['smpl_params_is_axis_angle']
# Compute 3D keypoint loss
loss_keypoints_2d = self.keypoint_2d_loss(pred_keypoints_2d, gt_keypoints_2d)
loss_keypoints_3d = self.keypoint_3d_loss(pred_keypoints_3d, gt_keypoints_3d, pelvis_id=25+14)
# Compute loss on SMPL parameters
loss_smpl_params = {}
for k, pred in pred_smpl_params.items():
gt = gt_smpl_params[k].view(batch_size, -1)
if is_axis_angle[k].all():
gt = aa_to_rotmat(gt.reshape(-1, 3)).view(batch_size, -1, 3, 3)
has_gt = has_smpl_params[k]
loss_smpl_params[k] = self.smpl_parameter_loss(pred.reshape(batch_size, -1), gt.reshape(batch_size, -1), has_gt)
loss = self.cfg.LOSS_WEIGHTS['KEYPOINTS_3D'] * loss_keypoints_3d+\
self.cfg.LOSS_WEIGHTS['KEYPOINTS_2D'] * loss_keypoints_2d+\
sum([loss_smpl_params[k] * self.cfg.LOSS_WEIGHTS[k.upper()] for k in loss_smpl_params])
losses = dict(loss=loss.detach(),
loss_keypoints_2d=loss_keypoints_2d.detach(),
loss_keypoints_3d=loss_keypoints_3d.detach())
for k, v in loss_smpl_params.items():
losses['loss_' + k] = v.detach()
output['losses'] = losses
return loss
# Tensoroboard logging should run from first rank only
@pl.utilities.rank_zero.rank_zero_only
def tensorboard_logging(self, batch: Dict, output: Dict, step_count: int, train: bool = True, write_to_summary_writer: bool = True) -> None:
"""
Log results to Tensorboard
Args:
batch (Dict): Dictionary containing batch data
output (Dict): Dictionary containing the regression output
step_count (int): Global training step count
train (bool): Flag indicating whether it is training or validation mode
"""
mode = 'train' if train else 'val'
batch_size = batch['keypoints_2d'].shape[0]
images = batch['img']
images = images * torch.tensor([0.229, 0.224, 0.225], device=images.device).reshape(1,3,1,1)
images = images + torch.tensor([0.485, 0.456, 0.406], device=images.device).reshape(1,3,1,1)
#images = 255*images.permute(0, 2, 3, 1).cpu().numpy()
pred_keypoints_3d = output['pred_keypoints_3d'].detach().reshape(batch_size, -1, 3)
pred_vertices = output['pred_vertices'].detach().reshape(batch_size, -1, 3)
focal_length = output['focal_length'].detach().reshape(batch_size, 2)
gt_keypoints_3d = batch['keypoints_3d']
gt_keypoints_2d = batch['keypoints_2d']
losses = output['losses']
pred_cam_t = output['pred_cam_t'].detach().reshape(batch_size, 3)
pred_keypoints_2d = output['pred_keypoints_2d'].detach().reshape(batch_size, -1, 2)
if write_to_summary_writer:
summary_writer = self.logger.experiment
for loss_name, val in losses.items():
summary_writer.add_scalar(mode +'/' + loss_name, val.detach().item(), step_count)
num_images = min(batch_size, self.cfg.EXTRA.NUM_LOG_IMAGES)
gt_keypoints_3d = batch['keypoints_3d']
pred_keypoints_3d = output['pred_keypoints_3d'].detach().reshape(batch_size, -1, 3)
# We render the skeletons instead of the full mesh because rendering a lot of meshes will make the training slow.
#predictions = self.renderer(pred_keypoints_3d[:num_images],
# gt_keypoints_3d[:num_images],
# 2 * gt_keypoints_2d[:num_images],
# images=images[:num_images],
# camera_translation=pred_cam_t[:num_images])
predictions = self.mesh_renderer.visualize_tensorboard(pred_vertices[:num_images].cpu().numpy(),
pred_cam_t[:num_images].cpu().numpy(),
images[:num_images].cpu().numpy(),
pred_keypoints_2d[:num_images].cpu().numpy(),
gt_keypoints_2d[:num_images].cpu().numpy(),
focal_length=focal_length[:num_images].cpu().numpy())
if write_to_summary_writer:
summary_writer.add_image('%s/predictions' % mode, predictions, step_count)
return predictions
def forward(self, batch: Dict) -> Dict:
"""
Run a forward step of the network in val mode
Args:
batch (Dict): Dictionary containing batch data
Returns:
Dict: Dictionary containing the regression output
"""
return self.forward_step(batch, train=False)
def training_step_discriminator(self, batch: Dict,
body_pose: torch.Tensor,
betas: torch.Tensor,
optimizer: torch.optim.Optimizer) -> torch.Tensor:
"""
Run a discriminator training step
Args:
batch (Dict): Dictionary containing mocap batch data
body_pose (torch.Tensor): Regressed body pose from current step
betas (torch.Tensor): Regressed betas from current step
optimizer (torch.optim.Optimizer): Discriminator optimizer
Returns:
torch.Tensor: Discriminator loss
"""
batch_size = body_pose.shape[0]
gt_body_pose = batch['body_pose']
gt_betas = batch['betas']
gt_rotmat = aa_to_rotmat(gt_body_pose.view(-1,3)).view(batch_size, -1, 3, 3)
disc_fake_out = self.discriminator(body_pose.detach(), betas.detach())
loss_fake = ((disc_fake_out - 0.0) ** 2).sum() / batch_size
disc_real_out = self.discriminator(gt_rotmat, gt_betas)
loss_real = ((disc_real_out - 1.0) ** 2).sum() / batch_size
loss_disc = loss_fake + loss_real
loss = self.cfg.LOSS_WEIGHTS.ADVERSARIAL * loss_disc
optimizer.zero_grad()
self.manual_backward(loss)
optimizer.step()
return loss_disc.detach()
def training_step(self, joint_batch: Dict, batch_idx: int) -> Dict:
"""
Run a full training step
Args:
joint_batch (Dict): Dictionary containing image and mocap batch data
batch_idx (int): Unused.
batch_idx (torch.Tensor): Unused.
Returns:
Dict: Dictionary containing regression output.
"""
batch = joint_batch['img']
mocap_batch = joint_batch['mocap']
optimizer = self.optimizers(use_pl_optimizer=True)
if self.cfg.LOSS_WEIGHTS.ADVERSARIAL > 0:
optimizer, optimizer_disc = optimizer
batch_size = batch['img'].shape[0]
output = self.forward_step(batch, train=True)
pred_smpl_params = output['pred_smpl_params']
if self.cfg.get('UPDATE_GT_SPIN', False):
self.update_batch_gt_spin(batch, output)
loss = self.compute_loss(batch, output, train=True)
if self.cfg.LOSS_WEIGHTS.ADVERSARIAL > 0:
disc_out = self.discriminator(pred_smpl_params['body_pose'].reshape(batch_size, -1), pred_smpl_params['betas'].reshape(batch_size, -1))
loss_adv = ((disc_out - 1.0) ** 2).sum() / batch_size
loss = loss + self.cfg.LOSS_WEIGHTS.ADVERSARIAL * loss_adv
# Error if Nan
if torch.isnan(loss):
raise ValueError('Loss is NaN')
optimizer.zero_grad()
self.manual_backward(loss)
# Clip gradient
if self.cfg.TRAIN.get('GRAD_CLIP_VAL', 0) > 0:
gn = torch.nn.utils.clip_grad_norm_(self.get_parameters(), self.cfg.TRAIN.GRAD_CLIP_VAL, error_if_nonfinite=True)
self.log('train/grad_norm', gn, on_step=True, on_epoch=True, prog_bar=True, logger=True)
optimizer.step()
if self.cfg.LOSS_WEIGHTS.ADVERSARIAL > 0:
loss_disc = self.training_step_discriminator(mocap_batch, pred_smpl_params['body_pose'].reshape(batch_size, -1), pred_smpl_params['betas'].reshape(batch_size, -1), optimizer_disc)
output['losses']['loss_gen'] = loss_adv
output['losses']['loss_disc'] = loss_disc
if self.global_step > 0 and self.global_step % self.cfg.GENERAL.LOG_STEPS == 0:
self.tensorboard_logging(batch, output, self.global_step, train=True)
self.log('train/loss', output['losses']['loss'], on_step=True, on_epoch=True, prog_bar=True, logger=False)
return output
def validation_step(self, batch: Dict, batch_idx: int, dataloader_idx=0) -> Dict:
"""
Run a validation step and log to Tensorboard
Args:
batch (Dict): Dictionary containing batch data
batch_idx (int): Unused.
Returns:
Dict: Dictionary containing regression output.
"""
# batch_size = batch['img'].shape[0]
output = self.forward_step(batch, train=False)
loss = self.compute_loss(batch, output, train=False)
output['loss'] = loss
self.tensorboard_logging(batch, output, self.global_step, train=False)
return output
+92
View File
@@ -0,0 +1,92 @@
import torch
import torch.nn as nn
class Keypoint2DLoss(nn.Module):
def __init__(self, loss_type: str = 'l1'):
"""
2D keypoint loss module.
Args:
loss_type (str): Choose between l1 and l2 losses.
"""
super(Keypoint2DLoss, self).__init__()
if loss_type == 'l1':
self.loss_fn = nn.L1Loss(reduction='none')
elif loss_type == 'l2':
self.loss_fn = nn.MSELoss(reduction='none')
else:
raise NotImplementedError('Unsupported loss function')
def forward(self, pred_keypoints_2d: torch.Tensor, gt_keypoints_2d: torch.Tensor) -> torch.Tensor:
"""
Compute 2D reprojection loss on the keypoints.
Args:
pred_keypoints_2d (torch.Tensor): Tensor of shape [B, S, N, 2] containing projected 2D keypoints (B: batch_size, S: num_samples, N: num_keypoints)
gt_keypoints_2d (torch.Tensor): Tensor of shape [B, S, N, 3] containing the ground truth 2D keypoints and confidence.
Returns:
torch.Tensor: 2D keypoint loss.
"""
conf = gt_keypoints_2d[:, :, -1].unsqueeze(-1).clone()
batch_size = conf.shape[0]
loss = (conf * self.loss_fn(pred_keypoints_2d, gt_keypoints_2d[:, :, :-1])).sum(dim=(1,2))
return loss.sum()
class Keypoint3DLoss(nn.Module):
def __init__(self, loss_type: str = 'l1'):
"""
3D keypoint loss module.
Args:
loss_type (str): Choose between l1 and l2 losses.
"""
super(Keypoint3DLoss, self).__init__()
if loss_type == 'l1':
self.loss_fn = nn.L1Loss(reduction='none')
elif loss_type == 'l2':
self.loss_fn = nn.MSELoss(reduction='none')
else:
raise NotImplementedError('Unsupported loss function')
def forward(self, pred_keypoints_3d: torch.Tensor, gt_keypoints_3d: torch.Tensor, pelvis_id: int = 39):
"""
Compute 3D keypoint loss.
Args:
pred_keypoints_3d (torch.Tensor): Tensor of shape [B, S, N, 3] containing the predicted 3D keypoints (B: batch_size, S: num_samples, N: num_keypoints)
gt_keypoints_3d (torch.Tensor): Tensor of shape [B, S, N, 4] containing the ground truth 3D keypoints and confidence.
Returns:
torch.Tensor: 3D keypoint loss.
"""
batch_size = pred_keypoints_3d.shape[0]
gt_keypoints_3d = gt_keypoints_3d.clone()
pred_keypoints_3d = pred_keypoints_3d - pred_keypoints_3d[:, pelvis_id, :].unsqueeze(dim=1)
gt_keypoints_3d[:, :, :-1] = gt_keypoints_3d[:, :, :-1] - gt_keypoints_3d[:, pelvis_id, :-1].unsqueeze(dim=1)
conf = gt_keypoints_3d[:, :, -1].unsqueeze(-1).clone()
gt_keypoints_3d = gt_keypoints_3d[:, :, :-1]
loss = (conf * self.loss_fn(pred_keypoints_3d, gt_keypoints_3d)).sum(dim=(1,2))
return loss.sum()
class ParameterLoss(nn.Module):
def __init__(self):
"""
SMPL parameter loss module.
"""
super(ParameterLoss, self).__init__()
self.loss_fn = nn.MSELoss(reduction='none')
def forward(self, pred_param: torch.Tensor, gt_param: torch.Tensor, has_param: torch.Tensor):
"""
Compute SMPL parameter loss.
Args:
pred_param (torch.Tensor): Tensor of shape [B, S, ...] containing the predicted parameters (body pose / global orientation / betas)
gt_param (torch.Tensor): Tensor of shape [B, S, ...] containing the ground truth SMPL parameters.
Returns:
torch.Tensor: L2 parameter loss loss.
"""
batch_size = pred_param.shape[0]
num_dims = len(pred_param.shape)
mask_dimension = [batch_size] + [1] * (num_dims-1)
has_param = has_param.type(pred_param.type()).view(*mask_dimension)
loss_param = (has_param * self.loss_fn(pred_param, gt_param))
return loss_param.sum()
@@ -0,0 +1,41 @@
import torch
import numpy as np
import pickle
from typing import Optional
import smplx
from smplx.lbs import vertices2joints
from smplx.utils import SMPLOutput
class SMPL(smplx.SMPLLayer):
def __init__(self, *args, joint_regressor_extra: Optional[str] = None, update_hips: bool = False, **kwargs):
"""
Extension of the official SMPL implementation to support more joints.
Args:
Same as SMPLLayer.
joint_regressor_extra (str): Path to extra joint regressor.
"""
super(SMPL, self).__init__(*args, **kwargs)
smpl_to_openpose = [24, 12, 17, 19, 21, 16, 18, 20, 0, 2, 5, 8, 1, 4,
7, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34]
if joint_regressor_extra is not None:
self.register_buffer('joint_regressor_extra', torch.tensor(pickle.load(open(joint_regressor_extra, 'rb'), encoding='latin1'), dtype=torch.float32))
self.register_buffer('joint_map', torch.tensor(smpl_to_openpose, dtype=torch.long))
self.update_hips = update_hips
def forward(self, *args, **kwargs) -> SMPLOutput:
"""
Run forward pass. Same as SMPL and also append an extra set of joints if joint_regressor_extra is specified.
"""
smpl_output = super(SMPL, self).forward(*args, **kwargs)
joints = smpl_output.joints[:, self.joint_map, :]
if self.update_hips:
joints[:,[9,12]] = joints[:,[9,12]] + \
0.25*(joints[:,[9,12]]-joints[:,[12,9]]) + \
0.5*(joints[:,[8]] - 0.5*(joints[:,[9,12]] + joints[:,[12,9]]))
if hasattr(self, 'joint_regressor_extra'):
extra_joints = vertices2joints(self.joint_regressor_extra, smpl_output.vertices)
joints = torch.cat([joints, extra_joints], dim=1)
smpl_output.joints = joints
return smpl_output
+25
View File
@@ -0,0 +1,25 @@
import torch
from typing import Any
from .renderer import Renderer
from .mesh_renderer import MeshRenderer
from .skeleton_renderer import SkeletonRenderer
from .pose_utils import eval_pose, Evaluator
def recursive_to(x: Any, target: torch.device):
"""
Recursively transfer a batch of data to the target device
Args:
x (Any): Batch of data.
target (torch.device): Target device.
Returns:
Batch of data where all tensors are transfered to the target device.
"""
if isinstance(x, dict):
return {k: recursive_to(v, target) for k, v in x.items()}
elif isinstance(x, torch.Tensor):
return x.to(target)
elif isinstance(x, list):
return [recursive_to(i, target) for i in x]
else:
return x
+67
View File
@@ -0,0 +1,67 @@
import os
import re
import sys
from urllib import request as urlrequest
def _progress_bar(count, total):
"""Report download progress. Credit:
https://stackoverflow.com/questions/3173320/text-progress-bar-in-the-console/27871113
"""
bar_len = 60
filled_len = int(round(bar_len * count / float(total)))
percents = round(100.0 * count / float(total), 1)
bar = "=" * filled_len + "-" * (bar_len - filled_len)
sys.stdout.write(
" [{}] {}% of {:.1f}MB file \r".format(bar, percents, total / 1024 / 1024)
)
sys.stdout.flush()
if count >= total:
sys.stdout.write("\n")
def download_url(url, dst_file_path, chunk_size=8192, progress_hook=_progress_bar):
"""Download url and write it to dst_file_path. Credit:
https://stackoverflow.com/questions/2028517/python-urllib2-progress-hook
"""
# url = url + "?dl=1" if "dropbox" in url else url
req = urlrequest.Request(url)
response = urlrequest.urlopen(req)
total_size = response.info().get("Content-Length")
if total_size is None:
raise ValueError("Cannot determine size of download from {}".format(url))
total_size = int(total_size.strip())
bytes_so_far = 0
with open(dst_file_path, "wb") as f:
while 1:
chunk = response.read(chunk_size)
bytes_so_far += len(chunk)
if not chunk:
break
if progress_hook:
progress_hook(bytes_so_far, total_size)
f.write(chunk)
return bytes_so_far
def cache_url(url_or_file, cache_file_path, download=True, log=True):
"""Download the file specified by the URL to the cache_dir and return the path to
the cached file. If the argument is not a URL, simply return it as is.
"""
is_url = re.match(r"^(?:http)s?://", url_or_file, re.IGNORECASE) is not None
if not is_url:
return url_or_file
url = url_or_file
if os.path.exists(cache_file_path):
return cache_file_path
cache_file_dir = os.path.dirname(cache_file_path)
if not os.path.exists(cache_file_dir):
os.makedirs(cache_file_dir)
if download:
if log:
print("Downloading remote file {} to {}".format(url, cache_file_path))
download_url(url, cache_file_path)
return cache_file_path
+102
View File
@@ -0,0 +1,102 @@
from typing import Optional
import torch
from torch.nn import functional as F
def aa_to_rotmat(theta: torch.Tensor):
"""
Convert axis-angle representation to rotation matrix.
Works by first converting it to a quaternion.
Args:
theta (torch.Tensor): Tensor of shape (B, 3) containing axis-angle representations.
Returns:
torch.Tensor: Corresponding rotation matrices with shape (B, 3, 3).
"""
norm = torch.norm(theta + 1e-8, p = 2, dim = 1)
angle = torch.unsqueeze(norm, -1)
normalized = torch.div(theta, angle)
angle = angle * 0.5
v_cos = torch.cos(angle)
v_sin = torch.sin(angle)
quat = torch.cat([v_cos, v_sin * normalized], dim = 1)
return quat_to_rotmat(quat)
def quat_to_rotmat(quat: torch.Tensor) -> torch.Tensor:
"""
Convert quaternion representation to rotation matrix.
Args:
quat (torch.Tensor) of shape (B, 4); 4 <===> (w, x, y, z).
Returns:
torch.Tensor: Corresponding rotation matrices with shape (B, 3, 3).
"""
norm_quat = quat
norm_quat = norm_quat/norm_quat.norm(p=2, dim=1, keepdim=True)
w, x, y, z = norm_quat[:,0], norm_quat[:,1], norm_quat[:,2], norm_quat[:,3]
B = quat.size(0)
w2, x2, y2, z2 = w.pow(2), x.pow(2), y.pow(2), z.pow(2)
wx, wy, wz = w*x, w*y, w*z
xy, xz, yz = x*y, x*z, y*z
rotMat = torch.stack([w2 + x2 - y2 - z2, 2*xy - 2*wz, 2*wy + 2*xz,
2*wz + 2*xy, w2 - x2 + y2 - z2, 2*yz - 2*wx,
2*xz - 2*wy, 2*wx + 2*yz, w2 - x2 - y2 + z2], dim=1).view(B, 3, 3)
return rotMat
def rot6d_to_rotmat(x: torch.Tensor) -> torch.Tensor:
"""
Convert 6D rotation representation to 3x3 rotation matrix.
Based on Zhou et al., "On the Continuity of Rotation Representations in Neural Networks", CVPR 2019
Args:
x (torch.Tensor): (B,6) Batch of 6-D rotation representations.
Returns:
torch.Tensor: Batch of corresponding rotation matrices with shape (B,3,3).
"""
x = x.reshape(-1,2,3).permute(0, 2, 1).contiguous()
a1 = x[:, :, 0]
a2 = x[:, :, 1]
b1 = F.normalize(a1)
b2 = F.normalize(a2 - torch.einsum('bi,bi->b', b1, a2).unsqueeze(-1) * b1)
b3 = torch.cross(b1, b2)
return torch.stack((b1, b2, b3), dim=-1)
def perspective_projection(points: torch.Tensor,
translation: torch.Tensor,
focal_length: torch.Tensor,
camera_center: Optional[torch.Tensor] = None,
rotation: Optional[torch.Tensor] = None) -> torch.Tensor:
"""
Computes the perspective projection of a set of 3D points.
Args:
points (torch.Tensor): Tensor of shape (B, N, 3) containing the input 3D points.
translation (torch.Tensor): Tensor of shape (B, 3) containing the 3D camera translation.
focal_length (torch.Tensor): Tensor of shape (B, 2) containing the focal length in pixels.
camera_center (torch.Tensor): Tensor of shape (B, 2) containing the camera center in pixels.
rotation (torch.Tensor): Tensor of shape (B, 3, 3) containing the camera rotation.
Returns:
torch.Tensor: Tensor of shape (B, N, 2) containing the projection of the input points.
"""
batch_size = points.shape[0]
if rotation is None:
rotation = torch.eye(3, device=points.device, dtype=points.dtype).unsqueeze(0).expand(batch_size, -1, -1)
if camera_center is None:
camera_center = torch.zeros(batch_size, 2, device=points.device, dtype=points.dtype)
# Populate intrinsic camera matrix K.
K = torch.zeros([batch_size, 3, 3], device=points.device, dtype=points.dtype)
K[:,0,0] = focal_length[:,0]
K[:,1,1] = focal_length[:,1]
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]
@@ -0,0 +1,149 @@
import os
if os.name == 'posix' and "DISPLAY" not in os.environ:
os.environ['PYOPENGL_PLATFORM'] = 'egl'
import torch
from torchvision.utils import make_grid
import numpy as np
import pyrender
import trimesh
import cv2
import torch.nn.functional as F
from .render_openpose import render_openpose
def create_raymond_lights():
import pyrender
thetas = np.pi * np.array([1.0 / 6.0, 1.0 / 6.0, 1.0 / 6.0])
phis = np.pi * np.array([0.0, 2.0 / 3.0, 4.0 / 3.0])
nodes = []
for phi, theta in zip(phis, thetas):
xp = np.sin(theta) * np.cos(phi)
yp = np.sin(theta) * np.sin(phi)
zp = np.cos(theta)
z = np.array([xp, yp, zp])
z = z / np.linalg.norm(z)
x = np.array([-z[1], z[0], 0.0])
if np.linalg.norm(x) == 0:
x = np.array([1.0, 0.0, 0.0])
x = x / np.linalg.norm(x)
y = np.cross(z, x)
matrix = np.eye(4)
matrix[:3,:3] = np.c_[x,y,z]
nodes.append(pyrender.Node(
light=pyrender.DirectionalLight(color=np.ones(3), intensity=1.0),
matrix=matrix
))
return nodes
class MeshRenderer:
def __init__(self, cfg, faces=None):
self.cfg = cfg
self.focal_length = cfg.EXTRA.FOCAL_LENGTH
self.img_res = cfg.MODEL.IMAGE_SIZE
self.renderer = pyrender.OffscreenRenderer(viewport_width=self.img_res,
viewport_height=self.img_res,
point_size=1.0)
self.camera_center = [self.img_res // 2, self.img_res // 2]
self.faces = faces
def visualize(self, vertices, camera_translation, images, focal_length=None, nrow=3, padding=2):
images_np = np.transpose(images, (0,2,3,1))
rend_imgs = []
for i in range(vertices.shape[0]):
fl = self.focal_length
rend_img = torch.from_numpy(np.transpose(self.__call__(vertices[i], camera_translation[i], images_np[i], focal_length=fl, side_view=False), (2,0,1))).float()
rend_img_side = torch.from_numpy(np.transpose(self.__call__(vertices[i], camera_translation[i], images_np[i], focal_length=fl, side_view=True), (2,0,1))).float()
rend_imgs.append(torch.from_numpy(images[i]))
rend_imgs.append(rend_img)
rend_imgs.append(rend_img_side)
rend_imgs = make_grid(rend_imgs, nrow=nrow, padding=padding)
return rend_imgs
def visualize_tensorboard(self, vertices, camera_translation, images, pred_keypoints, gt_keypoints, focal_length=None, nrow=5, padding=2):
images_np = np.transpose(images, (0,2,3,1))
rend_imgs = []
pred_keypoints = np.concatenate((pred_keypoints, np.ones_like(pred_keypoints)[:, :, [0]]), axis=-1)
pred_keypoints = self.img_res * (pred_keypoints + 0.5)
gt_keypoints[:, :, :-1] = self.img_res * (gt_keypoints[:, :, :-1] + 0.5)
keypoint_matches = [(1, 12), (2, 8), (3, 7), (4, 6), (5, 9), (6, 10), (7, 11), (8, 14), (9, 2), (10, 1), (11, 0), (12, 3), (13, 4), (14, 5)]
for i in range(vertices.shape[0]):
fl = self.focal_length
rend_img = torch.from_numpy(np.transpose(self.__call__(vertices[i], camera_translation[i], images_np[i], focal_length=fl, side_view=False), (2,0,1))).float()
rend_img_side = torch.from_numpy(np.transpose(self.__call__(vertices[i], camera_translation[i], images_np[i], focal_length=fl, side_view=True), (2,0,1))).float()
body_keypoints = pred_keypoints[i, :25]
extra_keypoints = pred_keypoints[i, -19:]
for pair in keypoint_matches:
body_keypoints[pair[0], :] = extra_keypoints[pair[1], :]
pred_keypoints_img = render_openpose(255 * images_np[i].copy(), body_keypoints) / 255
body_keypoints = gt_keypoints[i, :25]
extra_keypoints = gt_keypoints[i, -19:]
for pair in keypoint_matches:
if extra_keypoints[pair[1], -1] > 0 and body_keypoints[pair[0], -1] == 0:
body_keypoints[pair[0], :] = extra_keypoints[pair[1], :]
gt_keypoints_img = render_openpose(255*images_np[i].copy(), body_keypoints) / 255
rend_imgs.append(torch.from_numpy(images[i]))
rend_imgs.append(rend_img)
rend_imgs.append(rend_img_side)
rend_imgs.append(torch.from_numpy(pred_keypoints_img).permute(2,0,1))
rend_imgs.append(torch.from_numpy(gt_keypoints_img).permute(2,0,1))
rend_imgs = make_grid(rend_imgs, nrow=nrow, padding=padding)
return rend_imgs
def __call__(self, vertices, camera_translation, image, focal_length=5000, text=None, resize=None, side_view=False, baseColorFactor=(1.0, 1.0, 0.9, 1.0), rot_angle=90):
renderer = pyrender.OffscreenRenderer(viewport_width=image.shape[1],
viewport_height=image.shape[0],
point_size=1.0)
material = pyrender.MetallicRoughnessMaterial(
metallicFactor=0.0,
alphaMode='OPAQUE',
baseColorFactor=baseColorFactor)
camera_translation[0] *= -1.
mesh = trimesh.Trimesh(vertices.copy(), self.faces.copy())
if side_view:
rot = trimesh.transformations.rotation_matrix(
np.radians(rot_angle), [0, 1, 0])
mesh.apply_transform(rot)
rot = trimesh.transformations.rotation_matrix(
np.radians(180), [1, 0, 0])
mesh.apply_transform(rot)
mesh = pyrender.Mesh.from_trimesh(mesh, material=material)
scene = pyrender.Scene(bg_color=[0.0, 0.0, 0.0, 0.0],
ambient_light=(0.3, 0.3, 0.3))
scene.add(mesh, 'mesh')
camera_pose = np.eye(4)
camera_pose[:3, 3] = camera_translation
camera_center = [image.shape[1] / 2., image.shape[0] / 2.]
camera = pyrender.IntrinsicsCamera(fx=focal_length, fy=focal_length,
cx=camera_center[0], cy=camera_center[1])
scene.add(camera, pose=camera_pose)
light_nodes = create_raymond_lights()
for node in light_nodes:
scene.add_node(node)
color, rend_depth = renderer.render(scene, flags=pyrender.RenderFlags.RGBA)
color = color.astype(np.float32) / 255.0
valid_mask = (color[:, :, -1] > 0)[:, :, np.newaxis]
if not side_view:
output_img = (color[:, :, :3] * valid_mask +
(1 - valid_mask) * image)
else:
output_img = color[:, :, :3]
if resize is not None:
output_img = cv2.resize(output_img, resize)
output_img = output_img.astype(np.float32)
renderer.delete()
return output_img
+203
View File
@@ -0,0 +1,203 @@
import time
import warnings
from importlib.util import find_spec
from pathlib import Path
from typing import Callable, List
import hydra
from omegaconf import DictConfig, OmegaConf
from pytorch_lightning import Callback
from pytorch_lightning.loggers import Logger
from pytorch_lightning.utilities import rank_zero_only
from . import pylogger, rich_utils
log = pylogger.get_pylogger(__name__)
def task_wrapper(task_func: Callable) -> Callable:
"""Optional decorator that wraps the task function in extra utilities.
Makes multirun more resistant to failure.
Utilities:
- Calling the `utils.extras()` before the task is started
- Calling the `utils.close_loggers()` after the task is finished
- Logging the exception if occurs
- Logging the task total execution time
- Logging the output dir
"""
def wrap(cfg: DictConfig):
# apply extra utilities
extras(cfg)
# execute the task
try:
start_time = time.time()
ret = task_func(cfg=cfg)
except Exception as ex:
log.exception("") # save exception to `.log` file
raise ex
finally:
path = Path(cfg.paths.output_dir, "exec_time.log")
content = f"'{cfg.task_name}' execution time: {time.time() - start_time} (s)"
save_file(path, content) # save task execution time (even if exception occurs)
close_loggers() # close loggers (even if exception occurs so multirun won't fail)
log.info(f"Output dir: {cfg.paths.output_dir}")
return ret
return wrap
def extras(cfg: DictConfig) -> None:
"""Applies optional utilities before the task is started.
Utilities:
- Ignoring python warnings
- Setting tags from command line
- Rich config printing
"""
# return if no `extras` config
if not cfg.get("extras"):
log.warning("Extras config not found! <cfg.extras=null>")
return
# disable python warnings
if cfg.extras.get("ignore_warnings"):
log.info("Disabling python warnings! <cfg.extras.ignore_warnings=True>")
warnings.filterwarnings("ignore")
# prompt user to input tags from command line if none are provided in the config
if cfg.extras.get("enforce_tags"):
log.info("Enforcing tags! <cfg.extras.enforce_tags=True>")
rich_utils.enforce_tags(cfg, save_to_file=True)
# pretty print config tree using Rich library
if cfg.extras.get("print_config"):
log.info("Printing config tree with Rich! <cfg.extras.print_config=True>")
rich_utils.print_config_tree(cfg, resolve=True, save_to_file=True)
@rank_zero_only
def save_file(path: str, content: str) -> None:
"""Save file in rank zero mode (only on one process in multi-GPU setup)."""
with open(path, "w+") as file:
file.write(content)
def instantiate_callbacks(callbacks_cfg: DictConfig) -> List[Callback]:
"""Instantiates callbacks from config."""
callbacks: List[Callback] = []
if not callbacks_cfg:
log.warning("Callbacks config is empty.")
return callbacks
if not isinstance(callbacks_cfg, DictConfig):
raise TypeError("Callbacks config must be a DictConfig!")
for _, cb_conf in callbacks_cfg.items():
if isinstance(cb_conf, DictConfig) and "_target_" in cb_conf:
log.info(f"Instantiating callback <{cb_conf._target_}>")
callbacks.append(hydra.utils.instantiate(cb_conf))
return callbacks
def instantiate_loggers(logger_cfg: DictConfig) -> List[Logger]:
"""Instantiates loggers from config."""
logger: List[Logger] = []
if not logger_cfg:
log.warning("Logger config is empty.")
return logger
if not isinstance(logger_cfg, DictConfig):
raise TypeError("Logger config must be a DictConfig!")
for _, lg_conf in logger_cfg.items():
if isinstance(lg_conf, DictConfig) and "_target_" in lg_conf:
log.info(f"Instantiating logger <{lg_conf._target_}>")
logger.append(hydra.utils.instantiate(lg_conf))
return logger
@rank_zero_only
def log_hyperparameters(object_dict: dict) -> None:
"""Controls which config parts are saved by lightning loggers.
Additionally saves:
- Number of model parameters
"""
hparams = {}
cfg = object_dict["cfg"]
model = object_dict["model"]
trainer = object_dict["trainer"]
if not trainer.logger:
log.warning("Logger not found! Skipping hyperparameter logging...")
return
# save number of model parameters
hparams["model/params/total"] = sum(p.numel() for p in model.parameters())
hparams["model/params/trainable"] = sum(
p.numel() for p in model.parameters() if p.requires_grad
)
hparams["model/params/non_trainable"] = sum(
p.numel() for p in model.parameters() if not p.requires_grad
)
for k in cfg.keys():
hparams[k] = cfg.get(k)
# Resolve all interpolations
def _resolve(_cfg):
if isinstance(_cfg, DictConfig):
_cfg = OmegaConf.to_container(_cfg, resolve=True)
return _cfg
hparams = {k: _resolve(v) for k, v in hparams.items()}
# send hparams to all loggers
trainer.logger.log_hyperparams(hparams)
def get_metric_value(metric_dict: dict, metric_name: str) -> float:
"""Safely retrieves value of the metric logged in LightningModule."""
if not metric_name:
log.info("Metric name is None! Skipping metric value retrieval...")
return None
if metric_name not in metric_dict:
raise Exception(
f"Metric value not found! <metric_name={metric_name}>\n"
"Make sure metric name logged in LightningModule is correct!\n"
"Make sure `optimized_metric` name in `hparams_search` config is correct!"
)
metric_value = metric_dict[metric_name].item()
log.info(f"Retrieved metric value! <{metric_name}={metric_value}>")
return metric_value
def close_loggers() -> None:
"""Makes sure all loggers closed properly (prevents logging failure during multirun)."""
log.info("Closing loggers...")
if find_spec("wandb"): # if wandb is installed
import wandb
if wandb.run:
log.info("Closing wandb!")
wandb.finish()
@@ -0,0 +1,94 @@
# Copyright (c) OpenMMLab. All rights reserved.
import numpy as np
def _calc_distances(preds, targets, mask, normalize):
"""Calculate the normalized distances between preds and target.
Note:
batch_size: N
num_keypoints: K
dimension of keypoints: D (normally, D=2 or D=3)
Args:
preds (np.ndarray[N, K, D]): Predicted keypoint location.
targets (np.ndarray[N, K, D]): Groundtruth keypoint location.
mask (np.ndarray[N, K]): Visibility of the target. False for invisible
joints, and True for visible. Invisible joints will be ignored for
accuracy calculation.
normalize (np.ndarray[N, D]): Typical value is heatmap_size
Returns:
np.ndarray[K, N]: The normalized distances. \
If target keypoints are missing, the distance is -1.
"""
N, K, _ = preds.shape
# set mask=0 when normalize==0
_mask = mask.copy()
_mask[np.where((normalize == 0).sum(1))[0], :] = False
distances = np.full((N, K), -1, dtype=np.float32)
# handle invalid values
normalize[np.where(normalize <= 0)] = 1e6
distances[_mask] = np.linalg.norm(
((preds - targets) / normalize[:, None, :])[_mask], axis=-1)
return distances.T
def _distance_acc(distances, thr=0.5):
"""Return the percentage below the distance threshold, while ignoring
distances values with -1.
Note:
batch_size: N
Args:
distances (np.ndarray[N, ]): The normalized distances.
thr (float): Threshold of the distances.
Returns:
float: Percentage of distances below the threshold. \
If all target keypoints are missing, return -1.
"""
distance_valid = distances != -1
num_distance_valid = distance_valid.sum()
if num_distance_valid > 0:
return (distances[distance_valid] < thr).sum() / num_distance_valid
return -1
def keypoint_pck_accuracy(pred, gt, mask, thr, normalize):
"""Calculate the pose accuracy of PCK for each individual keypoint and the
averaged accuracy across all keypoints for coordinates.
Note:
PCK metric measures accuracy of the localization of the body joints.
The distances between predicted positions and the ground-truth ones
are typically normalized by the bounding box size.
The threshold (thr) of the normalized distance is commonly set
as 0.05, 0.1 or 0.2 etc.
- batch_size: N
- num_keypoints: K
Args:
pred (np.ndarray[N, K, 2]): Predicted keypoint location.
gt (np.ndarray[N, K, 2]): Groundtruth keypoint location.
mask (np.ndarray[N, K]): Visibility of the target. False for invisible
joints, and True for visible. Invisible joints will be ignored for
accuracy calculation.
thr (float): Threshold of PCK calculation.
normalize (np.ndarray[N, 2]): Normalization factor for H&W.
Returns:
tuple: A tuple containing keypoint accuracy.
- acc (np.ndarray[K]): Accuracy of each keypoint.
- avg_acc (float): Averaged accuracy across all keypoints.
- cnt (int): Number of valid keypoints.
"""
distances = _calc_distances(pred, gt, mask, normalize)
acc = np.array([_distance_acc(d, thr) for d in distances])
valid_acc = acc[acc >= 0]
cnt = len(valid_acc)
avg_acc = valid_acc.mean() if cnt > 0 else 0
return acc, avg_acc, cnt
+310
View File
@@ -0,0 +1,310 @@
"""
Code adapted from: https://github.com/akanazawa/hmr/blob/master/src/benchmark/eval_util.py
"""
import torch
import numpy as np
from typing import Optional, Dict, List, Tuple
def compute_similarity_transform(S1: torch.Tensor, S2: torch.Tensor) -> torch.Tensor:
"""
Computes a similarity transform (sR, t) in a batched way that takes
a set of 3D points S1 (B, N, 3) closest to a set of 3D points S2 (B, N, 3),
where R is a 3x3 rotation matrix, t 3x1 translation, s scale.
i.e. solves the orthogonal Procrutes problem.
Args:
S1 (torch.Tensor): First set of points of shape (B, N, 3).
S2 (torch.Tensor): Second set of points of shape (B, N, 3).
Returns:
(torch.Tensor): The first set of points after applying the similarity transformation.
"""
batch_size = S1.shape[0]
S1 = S1.permute(0, 2, 1)
S2 = S2.permute(0, 2, 1)
# 1. Remove mean.
mu1 = S1.mean(dim=2, keepdim=True)
mu2 = S2.mean(dim=2, keepdim=True)
X1 = S1 - mu1
X2 = S2 - mu2
# 2. Compute variance of X1 used for scale.
var1 = (X1**2).sum(dim=(1,2))
# 3. The outer product of X1 and X2.
K = torch.matmul(X1, X2.permute(0, 2, 1))
# 4. Solution that Maximizes trace(R'K) is R=U*V', where U, V are singular vectors of K.
U, s, V = torch.svd(K)
Vh = V.permute(0, 2, 1)
# Construct Z that fixes the orientation of R to get det(R)=1.
Z = torch.eye(U.shape[1], device=U.device).unsqueeze(0).repeat(batch_size, 1, 1)
Z[:, -1, -1] *= torch.sign(torch.linalg.det(torch.matmul(U, Vh)))
# Construct R.
R = torch.matmul(torch.matmul(V, Z), U.permute(0, 2, 1))
# 5. Recover scale.
trace = torch.matmul(R, K).diagonal(offset=0, dim1=-1, dim2=-2).sum(dim=-1)
scale = (trace / var1).unsqueeze(dim=-1).unsqueeze(dim=-1)
# 6. Recover translation.
t = mu2 - scale*torch.matmul(R, mu1)
# 7. Error:
S1_hat = scale*torch.matmul(R, S1) + t
return S1_hat.permute(0, 2, 1)
def reconstruction_error(S1, S2) -> np.array:
"""
Computes the mean Euclidean distance of 2 set of points S1, S2 after performing Procrustes alignment.
Args:
S1 (torch.Tensor): First set of points of shape (B, N, 3).
S2 (torch.Tensor): Second set of points of shape (B, N, 3).
Returns:
(np.array): Reconstruction error.
"""
S1_hat = compute_similarity_transform(S1, S2)
re = torch.sqrt( ((S1_hat - S2)** 2).sum(dim=-1)).mean(dim=-1)
return re
def eval_pose(pred_joints, gt_joints) -> Tuple[np.array, np.array]:
"""
Compute joint errors in mm before and after Procrustes alignment.
Args:
pred_joints (torch.Tensor): Predicted 3D joints of shape (B, N, 3).
gt_joints (torch.Tensor): Ground truth 3D joints of shape (B, N, 3).
Returns:
Tuple[np.array, np.array]: Joint errors in mm before and after alignment.
"""
# Absolute error (MPJPE)
mpjpe = torch.sqrt(((pred_joints - gt_joints) ** 2).sum(dim=-1)).mean(dim=-1).cpu().numpy()
# Reconstruction_error
r_error = reconstruction_error(pred_joints, gt_joints).cpu().numpy()
return 1000 * mpjpe, 1000 * r_error
class Evaluator:
def __init__(self,
dataset_length: int,
keypoint_list: List,
pelvis_ind: int,
metrics: List = ['mode_mpjpe', 'mode_re', 'min_mpjpe', 'min_re'],
pck_thresholds: Optional[List] = None):
"""
Class used for evaluating trained models on different 3D pose datasets.
Args:
dataset_length (int): Total dataset length.
keypoint_list [List]: List of keypoints used for evaluation.
pelvis_ind (int): Index of pelvis keypoint; used for aligning the predictions and ground truth.
metrics [List]: List of evaluation metrics to record.
"""
self.dataset_length = dataset_length
self.keypoint_list = keypoint_list
self.pelvis_ind = pelvis_ind
self.metrics = metrics
for metric in self.metrics:
setattr(self, metric, np.zeros((dataset_length,)))
self.counter = 0
if pck_thresholds is None:
self.pck_evaluator = None
else:
self.pck_evaluator = EvaluatorPCK(pck_thresholds)
def log(self):
"""
Print current evaluation metrics
"""
if self.counter == 0:
print('Evaluation has not started')
return
print(f'{self.counter} / {self.dataset_length} samples')
if self.pck_evaluator is not None:
self.pck_evaluator.log()
for metric in self.metrics:
if metric in ['mode_mpjpe', 'mode_re', 'min_mpjpe', 'min_re']:
unit = 'mm'
else:
unit = ''
print(f'{metric}: {getattr(self, metric)[:self.counter].mean()} {unit}')
print('***')
def get_metrics_dict(self) -> Dict:
"""
Returns:
Dict: Dictionary of evaluation metrics.
"""
d1 = {metric: getattr(self, metric)[:self.counter].mean() for metric in self.metrics}
if self.pck_evaluator is not None:
d2 = self.pck_evaluator.get_metrics_dict()
d1.update(d2)
return d1
def __call__(self, output: Dict, batch: Dict, opt_output: Optional[Dict] = None):
"""
Evaluate current batch.
Args:
output (Dict): Regression output.
batch (Dict): Dictionary containing images and their corresponding annotations.
opt_output (Dict): Optimization output.
"""
if self.pck_evaluator is not None:
self.pck_evaluator(output, batch, opt_output)
pred_keypoints_3d = output['pred_keypoints_3d'].detach()
pred_keypoints_3d = pred_keypoints_3d[:,None,:,:]
batch_size = pred_keypoints_3d.shape[0]
num_samples = pred_keypoints_3d.shape[1]
gt_keypoints_3d = batch['keypoints_3d'][:, :, :-1].unsqueeze(1).repeat(1, num_samples, 1, 1)
# Align predictions and ground truth such that the pelvis location is at the origin
pred_keypoints_3d -= pred_keypoints_3d[:, :, [self.pelvis_ind]]
gt_keypoints_3d -= gt_keypoints_3d[:, :, [self.pelvis_ind]]
# Compute joint errors
mpjpe, re = eval_pose(pred_keypoints_3d.reshape(batch_size * num_samples, -1, 3)[:, self.keypoint_list], gt_keypoints_3d.reshape(batch_size * num_samples, -1 ,3)[:, self.keypoint_list])
mpjpe = mpjpe.reshape(batch_size, num_samples)
re = re.reshape(batch_size, num_samples)
# Compute 2d keypoint errors
pred_keypoints_2d = output['pred_keypoints_2d'].detach()
pred_keypoints_2d = pred_keypoints_2d[:,None,:,:]
gt_keypoints_2d = batch['keypoints_2d'][:,None,:,:].repeat(1, num_samples, 1, 1)
conf = gt_keypoints_2d[:, :, :, -1].clone()
kp_err = torch.nn.functional.mse_loss(
pred_keypoints_2d,
gt_keypoints_2d[:, :, :, :-1],
reduction='none'
).sum(dim=3)
kp_l2_loss = (conf * kp_err).mean(dim=2)
kp_l2_loss = kp_l2_loss.detach().cpu().numpy()
# Compute joint errors after optimization, if available.
if opt_output is not None:
opt_keypoints_3d = opt_output['model_joints']
opt_keypoints_3d -= opt_keypoints_3d[:, [self.pelvis_ind]]
opt_mpjpe, opt_re = eval_pose(opt_keypoints_3d[:, self.keypoint_list], gt_keypoints_3d[:, 0, self.keypoint_list])
# The 0-th sample always corresponds to the mode
if hasattr(self, 'mode_mpjpe'):
mode_mpjpe = mpjpe[:, 0]
self.mode_mpjpe[self.counter:self.counter+batch_size] = mode_mpjpe
if hasattr(self, 'mode_re'):
mode_re = re[:, 0]
self.mode_re[self.counter:self.counter+batch_size] = mode_re
if hasattr(self, 'mode_kpl2'):
mode_kpl2 = kp_l2_loss[:, 0]
self.mode_kpl2[self.counter:self.counter+batch_size] = mode_kpl2
if hasattr(self, 'min_mpjpe'):
min_mpjpe = mpjpe.min(axis=-1)
self.min_mpjpe[self.counter:self.counter+batch_size] = min_mpjpe
if hasattr(self, 'min_re'):
min_re = re.min(axis=-1)
self.min_re[self.counter:self.counter+batch_size] = min_re
if hasattr(self, 'min_kpl2'):
min_kpl2 = kp_l2_loss.min(axis=-1)
self.min_kpl2[self.counter:self.counter+batch_size] = min_kpl2
if hasattr(self, 'opt_mpjpe'):
self.opt_mpjpe[self.counter:self.counter+batch_size] = opt_mpjpe
if hasattr(self, 'opt_re'):
self.opt_re[self.counter:self.counter+batch_size] = opt_re
self.counter += batch_size
if hasattr(self, 'mode_mpjpe') and hasattr(self, 'mode_re'):
return {
'mode_mpjpe': mode_mpjpe,
'mode_re': mode_re,
}
else:
return {}
class EvaluatorPCK:
def __init__(self, thresholds: List = [0.05, 0.1, 0.2, 0.3, 0.4, 0.5],):
"""
Class used for evaluating trained models on different 3D pose datasets.
Args:
thresholds [List]: List of PCK thresholds to evaluate.
metrics [List]: List of evaluation metrics to record.
"""
self.thresholds = thresholds
self.pred_kp_2d = []
self.gt_kp_2d = []
self.gt_conf_2d = []
self.counter = 0
def log(self):
"""
Print current evaluation metrics
"""
if self.counter == 0:
print('Evaluation has not started')
return
print(f'{self.counter} samples')
metrics_dict = self.get_metrics_dict()
for metric in metrics_dict:
print(f'{metric}: {metrics_dict[metric]}')
print('***')
def get_metrics_dict(self) -> Dict:
"""
Returns:
Dict: Dictionary of evaluation metrics.
"""
pcks = self.compute_pcks()
metrics = {}
for thr, (acc,avg_acc,cnt) in zip(self.thresholds, pcks):
metrics.update({f'kp{i}_pck_{thr}': float(a) for i, a in enumerate(acc) if a>=0})
metrics.update({f'kpAvg_pck_{thr}': float(avg_acc)})
return metrics
def compute_pcks(self):
pred_kp_2d = np.concatenate(self.pred_kp_2d, axis=0)
gt_kp_2d = np.concatenate(self.gt_kp_2d, axis=0)
gt_conf_2d = np.concatenate(self.gt_conf_2d, axis=0)
assert pred_kp_2d.shape == gt_kp_2d.shape
assert pred_kp_2d[..., 0].shape == gt_conf_2d.shape
assert pred_kp_2d.shape[1] == 1 # num_samples
from .pck_accuracy import keypoint_pck_accuracy
pcks = [
keypoint_pck_accuracy(
pred_kp_2d[:, 0, :, :],
gt_kp_2d[:, 0, :, :],
gt_conf_2d[:, 0, :]>0.5,
thr=thr,
normalize = np.ones((len(pred_kp_2d),2)) # Already in [-0.5,0.5] range. No need to normalize
)
for thr in self.thresholds
]
return pcks
def __call__(self, output: Dict, batch: Dict, opt_output: Optional[Dict] = None):
"""
Evaluate current batch.
Args:
output (Dict): Regression output.
batch (Dict): Dictionary containing images and their corresponding annotations.
opt_output (Dict): Optimization output.
"""
pred_keypoints_2d = output['pred_keypoints_2d'].detach()
num_samples = 1
batch_size = pred_keypoints_2d.shape[0]
pred_keypoints_2d = pred_keypoints_2d[:,None,:,:]
gt_keypoints_2d = batch['keypoints_2d'][:,None,:,:].repeat(1, num_samples, 1, 1)
gt_bbox_expand_factor = (batch['box_size']/(batch['_scale']*200).max(dim=-1).values)
gt_bbox_expand_factor = gt_bbox_expand_factor[:,None,None,None].repeat(1, num_samples, 1, 1)
gt_bbox_expand_factor = gt_bbox_expand_factor.detach().cpu().numpy()
self.pred_kp_2d.append(pred_keypoints_2d[:, :, :, :2].detach().cpu().numpy() * gt_bbox_expand_factor)
self.gt_conf_2d.append(gt_keypoints_2d[:, :, :, -1].detach().cpu().numpy())
self.gt_kp_2d.append(gt_keypoints_2d[:, :, :, :2].detach().cpu().numpy() * gt_bbox_expand_factor)
self.counter += batch_size
+17
View File
@@ -0,0 +1,17 @@
import logging
from pytorch_lightning.utilities import rank_zero_only
def get_pylogger(name=__name__) -> logging.Logger:
"""Initializes multi-GPU-friendly python command line logger."""
logger = logging.getLogger(name)
# this ensures all logging levels get marked with the rank zero decorator
# otherwise logs would get multiplied for each GPU process in multi-GPU setup
logging_levels = ("debug", "info", "warning", "error", "exception", "fatal", "critical")
for level in logging_levels:
setattr(logger, level, rank_zero_only(getattr(logger, level)))
return logger
@@ -0,0 +1,149 @@
"""
Render OpenPose keypoints.
Code was ported to Python from the official C++ implementation https://github.com/CMU-Perceptual-Computing-Lab/openpose/blob/master/src/openpose/utilities/keypoint.cpp
"""
import cv2
import math
import numpy as np
from typing import List, Tuple
def get_keypoints_rectangle(keypoints: np.array, threshold: float) -> Tuple[float, float, float]:
"""
Compute rectangle enclosing keypoints above the threshold.
Args:
keypoints (np.array): Keypoint array of shape (N, 3).
threshold (float): Confidence visualization threshold.
Returns:
Tuple[float, float, float]: Rectangle width, height and area.
"""
valid_ind = keypoints[:, -1] > threshold
if valid_ind.sum() > 0:
valid_keypoints = keypoints[valid_ind][:, :-1]
max_x = valid_keypoints[:,0].max()
max_y = valid_keypoints[:,1].max()
min_x = valid_keypoints[:,0].min()
min_y = valid_keypoints[:,1].min()
width = max_x - min_x
height = max_y - min_y
area = width * height
return width, height, area
else:
return 0,0,0
def render_keypoints(img: np.array,
keypoints: np.array,
pairs: List,
colors: List,
thickness_circle_ratio: float,
thickness_line_ratio_wrt_circle: float,
pose_scales: List,
threshold: float = 0.1) -> np.array:
"""
Render keypoints on input image.
Args:
img (np.array): Input image of shape (H, W, 3) with pixel values in the [0,255] range.
keypoints (np.array): Keypoint array of shape (N, 3).
pairs (List): List of keypoint pairs per limb.
colors: (List): List of colors per keypoint.
thickness_circle_ratio (float): Circle thickness ratio.
thickness_line_ratio_wrt_circle (float): Line thickness ratio wrt the circle.
pose_scales (List): List of pose scales.
threshold (float): Only visualize keypoints with confidence above the threshold.
Returns:
(np.array): Image of shape (H, W, 3) with keypoints drawn on top of the original image.
"""
img_orig = img.copy()
width, height = img.shape[1], img.shape[2]
area = width * height
lineType = 8
shift = 0
numberColors = len(colors)
thresholdRectangle = 0.1
person_width, person_height, person_area = get_keypoints_rectangle(keypoints, thresholdRectangle)
if person_area > 0:
ratioAreas = min(1, max(person_width / width, person_height / height))
thicknessRatio = np.maximum(np.round(math.sqrt(area) * thickness_circle_ratio * ratioAreas), 2)
thicknessCircle = np.maximum(1, thicknessRatio if ratioAreas > 0.05 else -np.ones_like(thicknessRatio))
thicknessLine = np.maximum(1, np.round(thicknessRatio * thickness_line_ratio_wrt_circle))
radius = thicknessRatio / 2
img = np.ascontiguousarray(img.copy())
for i, pair in enumerate(pairs):
index1, index2 = pair
if keypoints[index1, -1] > threshold and keypoints[index2, -1] > threshold:
thicknessLineScaled = int(round(min(thicknessLine[index1], thicknessLine[index2]) * pose_scales[0]))
colorIndex = index2
color = colors[colorIndex % numberColors]
keypoint1 = keypoints[index1, :-1].astype(np.int)
keypoint2 = keypoints[index2, :-1].astype(np.int)
cv2.line(img, tuple(keypoint1.tolist()), tuple(keypoint2.tolist()), tuple(color.tolist()), thicknessLineScaled, lineType, shift)
for part in range(len(keypoints)):
faceIndex = part
if keypoints[faceIndex, -1] > threshold:
radiusScaled = int(round(radius[faceIndex] * pose_scales[0]))
thicknessCircleScaled = int(round(thicknessCircle[faceIndex] * pose_scales[0]))
colorIndex = part
color = colors[colorIndex % numberColors]
center = keypoints[faceIndex, :-1].astype(np.int)
cv2.circle(img, tuple(center.tolist()), radiusScaled, tuple(color.tolist()), thicknessCircleScaled, lineType, shift)
return img
def render_body_keypoints(img: np.array,
body_keypoints: np.array) -> np.array:
"""
Render OpenPose body keypoints on input image.
Args:
img (np.array): Input image of shape (H, W, 3) with pixel values in the [0,255] range.
body_keypoints (np.array): Keypoint array of shape (N, 3); 3 <====> (x, y, confidence).
Returns:
(np.array): Image of shape (H, W, 3) with keypoints drawn on top of the original image.
"""
thickness_circle_ratio = 1./75. * np.ones(body_keypoints.shape[0])
thickness_line_ratio_wrt_circle = 0.75
pairs = []
pairs = [1,8,1,2,1,5,2,3,3,4,5,6,6,7,8,9,9,10,10,11,8,12,12,13,13,14,1,0,0,15,15,17,0,16,16,18,14,19,19,20,14,21,11,22,22,23,11,24]
pairs = np.array(pairs).reshape(-1,2)
colors = [255., 0., 85.,
255., 0., 0.,
255., 85., 0.,
255., 170., 0.,
255., 255., 0.,
170., 255., 0.,
85., 255., 0.,
0., 255., 0.,
255., 0., 0.,
0., 255., 85.,
0., 255., 170.,
0., 255., 255.,
0., 170., 255.,
0., 85., 255.,
0., 0., 255.,
255., 0., 170.,
170., 0., 255.,
255., 0., 255.,
85., 0., 255.,
0., 0., 255.,
0., 0., 255.,
0., 0., 255.,
0., 255., 255.,
0., 255., 255.,
0., 255., 255.]
colors = np.array(colors).reshape(-1,3)
pose_scales = [1]
return render_keypoints(img, body_keypoints, pairs, colors, thickness_circle_ratio, thickness_line_ratio_wrt_circle, pose_scales, 0.1)
def render_openpose(img: np.array,
body_keypoints: np.array) -> np.array:
"""
Render keypoints in the OpenPose format on input image.
Args:
img (np.array): Input image of shape (H, W, 3) with pixel values in the [0,255] range.
body_keypoints (np.array): Keypoint array of shape (N, 3); 3 <====> (x, y, confidence).
Returns:
(np.array): Image of shape (H, W, 3) with keypoints drawn on top of the original image.
"""
img = render_body_keypoints(img, body_keypoints)
return img
+399
View File
@@ -0,0 +1,399 @@
import os
if os.name == 'posix' and "DISPLAY" not in os.environ:
os.environ['PYOPENGL_PLATFORM'] = 'egl'
import torch
import numpy as np
import pyrender
import trimesh
import cv2
from yacs.config import CfgNode
from typing import List, Optional
def cam_crop_to_full(cam_bbox, box_center, box_size, img_size, focal_length=5000.):
# Convert cam_bbox to full image
img_w, img_h = img_size[:, 0], img_size[:, 1]
cx, cy, b = box_center[:, 0], box_center[:, 1], box_size
w_2, h_2 = img_w / 2., img_h / 2.
bs = b * cam_bbox[:, 0] + 1e-9
tz = 2 * focal_length / bs
tx = (2 * (cx - w_2) / bs) + cam_bbox[:, 1]
ty = (2 * (cy - h_2) / bs) + cam_bbox[:, 2]
full_cam = torch.stack([tx, ty, tz], dim=-1)
return full_cam
def get_light_poses(n_lights=5, elevation=np.pi / 3, dist=12):
# get lights in a circle around origin at elevation
thetas = elevation * np.ones(n_lights)
phis = 2 * np.pi * np.arange(n_lights) / n_lights
poses = []
trans = make_translation(torch.tensor([0, 0, dist]))
for phi, theta in zip(phis, thetas):
rot = make_rotation(rx=-theta, ry=phi, order="xyz")
poses.append((rot @ trans).numpy())
return poses
def make_translation(t):
return make_4x4_pose(torch.eye(3), t)
def make_rotation(rx=0, ry=0, rz=0, order="xyz"):
Rx = rotx(rx)
Ry = roty(ry)
Rz = rotz(rz)
if order == "xyz":
R = Rz @ Ry @ Rx
elif order == "xzy":
R = Ry @ Rz @ Rx
elif order == "yxz":
R = Rz @ Rx @ Ry
elif order == "yzx":
R = Rx @ Rz @ Ry
elif order == "zyx":
R = Rx @ Ry @ Rz
elif order == "zxy":
R = Ry @ Rx @ Rz
return make_4x4_pose(R, torch.zeros(3))
def make_4x4_pose(R, t):
"""
:param R (*, 3, 3)
:param t (*, 3)
return (*, 4, 4)
"""
dims = R.shape[:-2]
pose_3x4 = torch.cat([R, t.view(*dims, 3, 1)], dim=-1)
bottom = (
torch.tensor([0, 0, 0, 1], device=R.device)
.reshape(*(1,) * len(dims), 1, 4)
.expand(*dims, 1, 4)
)
return torch.cat([pose_3x4, bottom], dim=-2)
def rotx(theta):
return torch.tensor(
[
[1, 0, 0],
[0, np.cos(theta), -np.sin(theta)],
[0, np.sin(theta), np.cos(theta)],
],
dtype=torch.float32,
)
def roty(theta):
return torch.tensor(
[
[np.cos(theta), 0, np.sin(theta)],
[0, 1, 0],
[-np.sin(theta), 0, np.cos(theta)],
],
dtype=torch.float32,
)
def rotz(theta):
return torch.tensor(
[
[np.cos(theta), -np.sin(theta), 0],
[np.sin(theta), np.cos(theta), 0],
[0, 0, 1],
],
dtype=torch.float32,
)
def create_raymond_lights() -> List[pyrender.Node]:
"""
Return raymond light nodes for the scene.
"""
thetas = np.pi * np.array([1.0 / 6.0, 1.0 / 6.0, 1.0 / 6.0])
phis = np.pi * np.array([0.0, 2.0 / 3.0, 4.0 / 3.0])
nodes = []
for phi, theta in zip(phis, thetas):
xp = np.sin(theta) * np.cos(phi)
yp = np.sin(theta) * np.sin(phi)
zp = np.cos(theta)
z = np.array([xp, yp, zp])
z = z / np.linalg.norm(z)
x = np.array([-z[1], z[0], 0.0])
if np.linalg.norm(x) == 0:
x = np.array([1.0, 0.0, 0.0])
x = x / np.linalg.norm(x)
y = np.cross(z, x)
matrix = np.eye(4)
matrix[:3,:3] = np.c_[x,y,z]
nodes.append(pyrender.Node(
light=pyrender.DirectionalLight(color=np.ones(3), intensity=1.0),
matrix=matrix
))
return nodes
class Renderer:
def __init__(self, cfg: CfgNode, faces: np.array):
"""
Wrapper around the pyrender renderer to render SMPL meshes.
Args:
cfg (CfgNode): Model config file.
faces (np.array): Array of shape (F, 3) containing the mesh faces.
"""
self.cfg = cfg
self.focal_length = cfg.EXTRA.FOCAL_LENGTH
self.img_res = cfg.MODEL.IMAGE_SIZE
self.camera_center = [self.img_res // 2, self.img_res // 2]
self.faces = faces
def __call__(self,
vertices: np.array,
camera_translation: np.array,
image: torch.Tensor,
full_frame: bool = False,
imgname: Optional[str] = None,
side_view=False, top_view=False,
rot_angle=90,
mesh_base_color=(1.0, 1.0, 0.9),
scene_bg_color=(0,0,0),
return_rgba=False,
) -> np.array:
"""
Render meshes on input image
Args:
vertices (np.array): Array of shape (V, 3) containing the mesh vertices.
camera_translation (np.array): Array of shape (3,) with the camera translation.
image (torch.Tensor): Tensor of shape (3, H, W) containing the image crop with normalized pixel values.
full_frame (bool): If True, then render on the full image.
imgname (Optional[str]): Contains the original image filenamee. Used only if full_frame == True.
"""
if full_frame:
image = cv2.imread(imgname).astype(np.float32)[:, :, ::-1] / 255.
else:
image = image.clone() * torch.tensor(self.cfg.MODEL.IMAGE_STD, device=image.device).reshape(3,1,1)
image = image + torch.tensor(self.cfg.MODEL.IMAGE_MEAN, device=image.device).reshape(3,1,1)
image = image.permute(1, 2, 0).cpu().numpy()
renderer = pyrender.OffscreenRenderer(viewport_width=image.shape[1],
viewport_height=image.shape[0],
point_size=1.0)
material = pyrender.MetallicRoughnessMaterial(
metallicFactor=0.0,
alphaMode='OPAQUE',
baseColorFactor=(*mesh_base_color, 1.0))
camera_translation[0] *= -1.
mesh = trimesh.Trimesh(vertices.copy(), self.faces.copy())
if side_view:
rot = trimesh.transformations.rotation_matrix(
np.radians(rot_angle), [0, 1, 0])
mesh.apply_transform(rot)
elif top_view:
rot = trimesh.transformations.rotation_matrix(
np.radians(rot_angle), [1, 0, 0])
mesh.apply_transform(rot)
rot = trimesh.transformations.rotation_matrix(
np.radians(180), [1, 0, 0])
mesh.apply_transform(rot)
mesh = pyrender.Mesh.from_trimesh(mesh, material=material)
scene = pyrender.Scene(bg_color=[*scene_bg_color, 0.0],
ambient_light=(0.3, 0.3, 0.3))
scene.add(mesh, 'mesh')
camera_pose = np.eye(4)
camera_pose[:3, 3] = camera_translation
camera_center = [image.shape[1] / 2., image.shape[0] / 2.]
camera = pyrender.IntrinsicsCamera(fx=self.focal_length, fy=self.focal_length,
cx=camera_center[0], cy=camera_center[1], zfar=1e12)
scene.add(camera, pose=camera_pose)
light_nodes = create_raymond_lights()
for node in light_nodes:
scene.add_node(node)
color, rend_depth = renderer.render(scene, flags=pyrender.RenderFlags.RGBA)
color = color.astype(np.float32) / 255.0
renderer.delete()
if return_rgba:
return color
valid_mask = (color[:, :, -1])[:, :, np.newaxis]
if not side_view and not top_view:
output_img = (color[:, :, :3] * valid_mask + (1 - valid_mask) * image)
else:
output_img = color[:, :, :3]
output_img = output_img.astype(np.float32)
return output_img
def vertices_to_trimesh(self, vertices, camera_translation, mesh_base_color=(1.0, 1.0, 0.9),
rot_axis=[1,0,0], rot_angle=0,):
# material = pyrender.MetallicRoughnessMaterial(
# metallicFactor=0.0,
# alphaMode='OPAQUE',
# baseColorFactor=(*mesh_base_color, 1.0))
vertex_colors = np.array([(*mesh_base_color, 1.0)] * vertices.shape[0])
mesh = trimesh.Trimesh(vertices.copy() + camera_translation, self.faces.copy(), vertex_colors=vertex_colors)
# mesh = trimesh.Trimesh(vertices.copy(), self.faces.copy())
rot = trimesh.transformations.rotation_matrix(
np.radians(rot_angle), rot_axis)
mesh.apply_transform(rot)
rot = trimesh.transformations.rotation_matrix(
np.radians(180), [1, 0, 0])
mesh.apply_transform(rot)
return mesh
def render_rgba(
self,
vertices: np.array,
cam_t = None,
rot=None,
rot_axis=[1,0,0],
rot_angle=0,
camera_z=3,
# camera_translation: np.array,
mesh_base_color=(1.0, 1.0, 0.9),
scene_bg_color=(0,0,0),
render_res=[256, 256],
):
renderer = pyrender.OffscreenRenderer(viewport_width=render_res[0],
viewport_height=render_res[1],
point_size=1.0)
# material = pyrender.MetallicRoughnessMaterial(
# metallicFactor=0.0,
# alphaMode='OPAQUE',
# baseColorFactor=(*mesh_base_color, 1.0))
if cam_t is not None:
camera_translation = cam_t.copy()
# camera_translation[0] *= -1.
else:
camera_translation = np.array([0, 0, camera_z * self.focal_length/render_res[1]])
mesh = self.vertices_to_trimesh(vertices, camera_translation, mesh_base_color, rot_axis, rot_angle)
mesh = pyrender.Mesh.from_trimesh(mesh)
# mesh = pyrender.Mesh.from_trimesh(mesh, material=material)
scene = pyrender.Scene(bg_color=[*scene_bg_color, 0.0],
ambient_light=(0.3, 0.3, 0.3))
scene.add(mesh, 'mesh')
camera_pose = np.eye(4)
# camera_pose[:3, 3] = camera_translation
camera_center = [render_res[0] / 2., render_res[1] / 2.]
camera = pyrender.IntrinsicsCamera(fx=self.focal_length, fy=self.focal_length,
cx=camera_center[0], cy=camera_center[1], zfar=1e12)
# Create camera node and add it to pyRender scene
camera_node = pyrender.Node(camera=camera, matrix=camera_pose)
scene.add_node(camera_node)
self.add_point_lighting(scene, camera_node)
self.add_lighting(scene, camera_node)
light_nodes = create_raymond_lights()
for node in light_nodes:
scene.add_node(node)
color, rend_depth = renderer.render(scene, flags=pyrender.RenderFlags.RGBA)
color = color.astype(np.float32) / 255.0
renderer.delete()
return color
def render_rgba_multiple(
self,
vertices: List[np.array],
cam_t: List[np.array],
rot_axis=[1,0,0],
rot_angle=0,
mesh_base_color=(1.0, 1.0, 0.9),
scene_bg_color=(0,0,0),
render_res=[256, 256],
focal_length=None,
):
renderer = pyrender.OffscreenRenderer(viewport_width=render_res[0],
viewport_height=render_res[1],
point_size=1.0)
# material = pyrender.MetallicRoughnessMaterial(
# metallicFactor=0.0,
# alphaMode='OPAQUE',
# baseColorFactor=(*mesh_base_color, 1.0))
mesh_list = [pyrender.Mesh.from_trimesh(self.vertices_to_trimesh(vvv, ttt.copy(), mesh_base_color, rot_axis, rot_angle)) for vvv,ttt in zip(vertices, cam_t)]
scene = pyrender.Scene(bg_color=[*scene_bg_color, 0.0],
ambient_light=(0.3, 0.3, 0.3))
for i,mesh in enumerate(mesh_list):
scene.add(mesh, f'mesh_{i}')
camera_pose = np.eye(4)
# camera_pose[:3, 3] = camera_translation
camera_center = [render_res[0] / 2., render_res[1] / 2.]
focal_length = focal_length if focal_length is not None else self.focal_length
camera = pyrender.IntrinsicsCamera(fx=focal_length, fy=focal_length,
cx=camera_center[0], cy=camera_center[1], zfar=1e12)
# Create camera node and add it to pyRender scene
camera_node = pyrender.Node(camera=camera, matrix=camera_pose)
scene.add_node(camera_node)
self.add_point_lighting(scene, camera_node)
self.add_lighting(scene, camera_node)
light_nodes = create_raymond_lights()
for node in light_nodes:
scene.add_node(node)
color, rend_depth = renderer.render(scene, flags=pyrender.RenderFlags.RGBA)
color = color.astype(np.float32) / 255.0
renderer.delete()
return color
def add_lighting(self, scene, cam_node, color=np.ones(3), intensity=1.0):
# from phalp.visualize.py_renderer import get_light_poses
light_poses = get_light_poses()
light_poses.append(np.eye(4))
cam_pose = scene.get_pose(cam_node)
for i, pose in enumerate(light_poses):
matrix = cam_pose @ pose
node = pyrender.Node(
name=f"light-{i:02d}",
light=pyrender.DirectionalLight(color=color, intensity=intensity),
matrix=matrix,
)
if scene.has_node(node):
continue
scene.add_node(node)
def add_point_lighting(self, scene, cam_node, color=np.ones(3), intensity=1.0):
# from phalp.visualize.py_renderer import get_light_poses
light_poses = get_light_poses(dist=0.5)
light_poses.append(np.eye(4))
cam_pose = scene.get_pose(cam_node)
for i, pose in enumerate(light_poses):
matrix = cam_pose @ pose
# node = pyrender.Node(
# name=f"light-{i:02d}",
# light=pyrender.DirectionalLight(color=color, intensity=intensity),
# matrix=matrix,
# )
node = pyrender.Node(
name=f"plight-{i:02d}",
light=pyrender.PointLight(color=color, intensity=intensity),
matrix=matrix,
)
if scene.has_node(node):
continue
scene.add_node(node)
+105
View File
@@ -0,0 +1,105 @@
from pathlib import Path
from typing import Sequence
import rich
import rich.syntax
import rich.tree
from hydra.core.hydra_config import HydraConfig
from omegaconf import DictConfig, OmegaConf, open_dict
from pytorch_lightning.utilities import rank_zero_only
from rich.prompt import Prompt
from . import pylogger
log = pylogger.get_pylogger(__name__)
@rank_zero_only
def print_config_tree(
cfg: DictConfig,
print_order: Sequence[str] = (
"datamodule",
"model",
"callbacks",
"logger",
"trainer",
"paths",
"extras",
),
resolve: bool = False,
save_to_file: bool = False,
) -> None:
"""Prints content of DictConfig using Rich library and its tree structure.
Args:
cfg (DictConfig): Configuration composed by Hydra.
print_order (Sequence[str], optional): Determines in what order config components are printed.
resolve (bool, optional): Whether to resolve reference fields of DictConfig.
save_to_file (bool, optional): Whether to export config to the hydra output folder.
"""
style = "dim"
tree = rich.tree.Tree("CONFIG", style=style, guide_style=style)
queue = []
# add fields from `print_order` to queue
for field in print_order:
queue.append(field) if field in cfg else log.warning(
f"Field '{field}' not found in config. Skipping '{field}' config printing..."
)
# add all the other fields to queue (not specified in `print_order`)
for field in cfg:
if field not in queue:
queue.append(field)
# generate config tree from queue
for field in queue:
branch = tree.add(field, style=style, guide_style=style)
config_group = cfg[field]
if isinstance(config_group, DictConfig):
branch_content = OmegaConf.to_yaml(config_group, resolve=resolve)
else:
branch_content = str(config_group)
branch.add(rich.syntax.Syntax(branch_content, "yaml"))
# print config tree
rich.print(tree)
# save config tree to file
if save_to_file:
with open(Path(cfg.paths.output_dir, "config_tree.log"), "w") as file:
rich.print(tree, file=file)
@rank_zero_only
def enforce_tags(cfg: DictConfig, save_to_file: bool = False) -> None:
"""Prompts user to input tags from command line if no tags are provided in config."""
if not cfg.get("tags"):
if "id" in HydraConfig().cfg.hydra.job:
raise ValueError("Specify tags before launching a multirun!")
log.warning("No tags provided in config. Prompting user to input tags...")
tags = Prompt.ask("Enter a list of comma separated tags", default="dev")
tags = [t.strip() for t in tags.split(",") if t != ""]
with open_dict(cfg):
cfg.tags = tags
log.info(f"Tags: {cfg.tags}")
if save_to_file:
with open(Path(cfg.paths.output_dir, "tags.log"), "w") as file:
rich.print(cfg.tags, file=file)
if __name__ == "__main__":
from hydra import compose, initialize
with initialize(version_base="1.2", config_path="../../configs"):
cfg = compose(config_name="train.yaml", return_hydra_config=False, overrides=[])
print_config_tree(cfg, resolve=False, save_to_file=False)
@@ -0,0 +1,122 @@
import torch
import numpy as np
import trimesh
from typing import Optional
from yacs.config import CfgNode
from .geometry import perspective_projection
from .render_openpose import render_openpose
class SkeletonRenderer:
def __init__(self, cfg: CfgNode):
"""
Object used to render 3D keypoints. Faster for use during training.
Args:
cfg (CfgNode): Model config file.
"""
self.cfg = cfg
def __call__(self,
pred_keypoints_3d: torch.Tensor,
gt_keypoints_3d: torch.Tensor,
gt_keypoints_2d: torch.Tensor,
images: Optional[np.array] = None,
camera_translation: Optional[torch.Tensor] = None) -> np.array:
"""
Render batch of 3D keypoints.
Args:
pred_keypoints_3d (torch.Tensor): Tensor of shape (B, S, N, 3) containing a batch of predicted 3D keypoints, with S samples per image.
gt_keypoints_3d (torch.Tensor): Tensor of shape (B, N, 4) containing corresponding ground truth 3D keypoints; last value is the confidence.
gt_keypoints_2d (torch.Tensor): Tensor of shape (B, N, 3) containing corresponding ground truth 2D keypoints.
images (torch.Tensor): Tensor of shape (B, H, W, 3) containing images with values in the [0,255] range.
camera_translation (torch.Tensor): Tensor of shape (B, 3) containing the camera translation.
Returns:
np.array : Image with the following layout. Each row contains the a) input image,
b) image with gt 2D keypoints,
c) image with projected gt 3D keypoints,
d_1, ... , d_S) image with projected predicted 3D keypoints,
e) gt 3D keypoints rendered from a side view,
f_1, ... , f_S) predicted 3D keypoints frorm a side view
"""
batch_size = pred_keypoints_3d.shape[0]
# num_samples = pred_keypoints_3d.shape[1]
pred_keypoints_3d = pred_keypoints_3d.clone().cpu().float()
gt_keypoints_3d = gt_keypoints_3d.clone().cpu().float()
gt_keypoints_3d[:, :, :-1] = gt_keypoints_3d[:, :, :-1] - gt_keypoints_3d[:, [25+14], :-1] + pred_keypoints_3d[:, [25+14]]
gt_keypoints_2d = gt_keypoints_2d.clone().cpu().float().numpy()
gt_keypoints_2d[:, :, :-1] = self.cfg.MODEL.IMAGE_SIZE * (gt_keypoints_2d[:, :, :-1] + 1.0) / 2.0
openpose_indices = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14]
gt_indices = [12, 8, 7, 6, 9, 10, 11, 14, 2, 1, 0, 3, 4, 5]
gt_indices = [25 + i for i in gt_indices]
keypoints_to_render = torch.ones(batch_size, gt_keypoints_3d.shape[1], 1)
rotation = torch.eye(3).unsqueeze(0)
if camera_translation is None:
camera_translation = torch.tensor([0.0, 0.0, 2 * self.cfg.EXTRA.FOCAL_LENGTH / (0.8 * self.cfg.MODEL.IMAGE_SIZE)]).unsqueeze(0).repeat(batch_size, 1)
else:
camera_translation = camera_translation.cpu()
if images is None:
images = np.zeros((batch_size, self.cfg.MODEL.IMAGE_SIZE, self.cfg.MODEL.IMAGE_SIZE, 3))
focal_length = torch.tensor([self.cfg.EXTRA.FOCAL_LENGTH, self.cfg.EXTRA.FOCAL_LENGTH]).reshape(1, 2)
camera_center = torch.tensor([self.cfg.MODEL.IMAGE_SIZE, self.cfg.MODEL.IMAGE_SIZE], dtype=torch.float).reshape(1, 2) / 2.
gt_keypoints_3d_proj = perspective_projection(gt_keypoints_3d[:, :, :-1], rotation=rotation.repeat(batch_size, 1, 1), translation=camera_translation[:, :], focal_length=focal_length.repeat(batch_size, 1), camera_center=camera_center.repeat(batch_size, 1))
pred_keypoints_3d_proj = perspective_projection(pred_keypoints_3d.reshape(batch_size, -1, 3), rotation=rotation.repeat(batch_size, 1, 1), translation=camera_translation.reshape(batch_size, -1), focal_length=focal_length.repeat(batch_size, 1), camera_center=camera_center.repeat(batch_size, 1)).reshape(batch_size, -1, 2)
gt_keypoints_3d_proj = torch.cat([gt_keypoints_3d_proj, gt_keypoints_3d[:, :, [-1]]], dim=-1).cpu().numpy()
pred_keypoints_3d_proj = torch.cat([pred_keypoints_3d_proj, keypoints_to_render.reshape(batch_size, -1, 1)], dim=-1).cpu().numpy()
rows = []
# Rotate keypoints to visualize side view
R = torch.tensor(trimesh.transformations.rotation_matrix(np.radians(90), [0, 1, 0])[:3, :3]).float()
gt_keypoints_3d_side = gt_keypoints_3d.clone()
gt_keypoints_3d_side[:, :, :-1] = torch.einsum('bni,ij->bnj', gt_keypoints_3d_side[:, :, :-1], R)
pred_keypoints_3d_side = pred_keypoints_3d.clone()
pred_keypoints_3d_side = torch.einsum('bni,ij->bnj', pred_keypoints_3d_side, R)
gt_keypoints_3d_proj_side = perspective_projection(gt_keypoints_3d_side[:, :, :-1], rotation=rotation.repeat(batch_size, 1, 1), translation=camera_translation[:, :], focal_length=focal_length.repeat(batch_size, 1), camera_center=camera_center.repeat(batch_size, 1))
pred_keypoints_3d_proj_side = perspective_projection(pred_keypoints_3d_side.reshape(batch_size, -1, 3), rotation=rotation.repeat(batch_size, 1, 1), translation=camera_translation.reshape(batch_size, -1), focal_length=focal_length.repeat(batch_size, 1), camera_center=camera_center.repeat(batch_size, 1)).reshape(batch_size, -1, 2)
gt_keypoints_3d_proj_side = torch.cat([gt_keypoints_3d_proj_side, gt_keypoints_3d_side[:, :, [-1]]], dim=-1).cpu().numpy()
pred_keypoints_3d_proj_side = torch.cat([pred_keypoints_3d_proj_side, keypoints_to_render.reshape(batch_size, -1, 1)], dim=-1).cpu().numpy()
for i in range(batch_size):
img = images[i]
side_img = np.zeros((self.cfg.MODEL.IMAGE_SIZE, self.cfg.MODEL.IMAGE_SIZE, 3))
# gt 2D keypoints
body_keypoints_2d = gt_keypoints_2d[i, :25].copy()
for op, gt in zip(openpose_indices, gt_indices):
if gt_keypoints_2d[i, gt, -1] > body_keypoints_2d[op, -1]:
body_keypoints_2d[op] = gt_keypoints_2d[i, gt]
gt_keypoints_img = render_openpose(img, body_keypoints_2d) / 255.
# gt 3D keypoints
body_keypoints_3d_proj = gt_keypoints_3d_proj[i, :25].copy()
for op, gt in zip(openpose_indices, gt_indices):
if gt_keypoints_3d_proj[i, gt, -1] > body_keypoints_3d_proj[op, -1]:
body_keypoints_3d_proj[op] = gt_keypoints_3d_proj[i, gt]
gt_keypoints_3d_proj_img = render_openpose(img, body_keypoints_3d_proj) / 255.
# gt 3D keypoints from the side
body_keypoints_3d_proj = gt_keypoints_3d_proj_side[i, :25].copy()
for op, gt in zip(openpose_indices, gt_indices):
if gt_keypoints_3d_proj_side[i, gt, -1] > body_keypoints_3d_proj[op, -1]:
body_keypoints_3d_proj[op] = gt_keypoints_3d_proj_side[i, gt]
gt_keypoints_3d_proj_img_side = render_openpose(side_img, body_keypoints_3d_proj) / 255.
# pred 3D keypoints
pred_keypoints_3d_proj_imgs = []
body_keypoints_3d_proj = pred_keypoints_3d_proj[i, :25].copy()
for op, gt in zip(openpose_indices, gt_indices):
if pred_keypoints_3d_proj[i, gt, -1] >= body_keypoints_3d_proj[op, -1]:
body_keypoints_3d_proj[op] = pred_keypoints_3d_proj[i, gt]
pred_keypoints_3d_proj_imgs.append(render_openpose(img, body_keypoints_3d_proj) / 255.)
pred_keypoints_3d_proj_img = np.concatenate(pred_keypoints_3d_proj_imgs, axis=1)
# gt 3D keypoints from the side
pred_keypoints_3d_proj_imgs_side = []
body_keypoints_3d_proj = pred_keypoints_3d_proj_side[i, :25].copy()
for op, gt in zip(openpose_indices, gt_indices):
if pred_keypoints_3d_proj_side[i, gt, -1] >= body_keypoints_3d_proj[op, -1]:
body_keypoints_3d_proj[op] = pred_keypoints_3d_proj_side[i, gt]
pred_keypoints_3d_proj_imgs_side.append(render_openpose(side_img, body_keypoints_3d_proj) / 255.)
pred_keypoints_3d_proj_img_side = np.concatenate(pred_keypoints_3d_proj_imgs_side, axis=1)
rows.append(np.concatenate((gt_keypoints_img, gt_keypoints_3d_proj_img, pred_keypoints_3d_proj_img, gt_keypoints_3d_proj_img_side, pred_keypoints_3d_proj_img_side), axis=1))
# Concatenate images
img = np.concatenate(rows, axis=0)
img[:, ::self.cfg.MODEL.IMAGE_SIZE, :] = 1.0
img[::self.cfg.MODEL.IMAGE_SIZE, :, :] = 1.0
img[:, (1+1+1)*self.cfg.MODEL.IMAGE_SIZE, :] = 0.5
return img
@@ -0,0 +1,85 @@
import numpy as np
import torch
from torch.nn import functional as F
# from psbody.mesh.visibility import visibility_compute
def uv_to_xyz_and_normals(verts, f, fmap, bmap, ftov):
vn = estimate_vertex_normals(verts, f, ftov)
pixels_to_set = torch.nonzero(fmap+1)
x_to_set = pixels_to_set[:,0]
y_to_set = pixels_to_set[:,1]
b_coords = bmap[x_to_set, y_to_set, :]
f_coords = fmap[x_to_set, y_to_set]
v_ids = f[f_coords]
points = (b_coords[:,0,None]*verts[:,v_ids[:,0]]
+ b_coords[:,1,None]*verts[:,v_ids[:,1]]
+ b_coords[:,2,None]*verts[:,v_ids[:,2]])
normals = (b_coords[:,0,None]*vn[:,v_ids[:,0]]
+ b_coords[:,1,None]*vn[:,v_ids[:,1]]
+ b_coords[:,2,None]*vn[:,v_ids[:,2]])
return points, normals, vn, f_coords
def estimate_vertex_normals(v, f, ftov):
face_normals = TriNormalsScaled(v, f)
non_scaled_normals = torch.einsum('ij,bjk->bik', ftov, face_normals)
norms = torch.sum(non_scaled_normals ** 2.0, 2) ** 0.5
norms[norms == 0] = 1.0
return torch.div(non_scaled_normals, norms[:,:,None])
def TriNormalsScaled(v, f):
return torch.cross(_edges_for(v, f, 1, 0), _edges_for(v, f, 2, 0))
def _edges_for(v, f, cplus, cminus):
return v[:,f[:,cplus]] - v[:,f[:,cminus]]
def psbody_get_face_visibility(v, n, f, cams, normal_threshold=0.5):
bn, nverts, _ = v.shape
nfaces, _ = f.shape
vis_f = np.zeros([bn, nfaces], dtype='float32')
for i in range(bn):
vis, n_dot_cam = visibility_compute(v=v[i], n=n[i], f=f, cams=cams)
vis_v = (vis == 1) & (n_dot_cam > normal_threshold)
vis_f[i] = np.all(vis_v[0,f],1)
return vis_f
def compute_uvsampler(vt, ft, tex_size=6):
"""
For this mesh, pre-computes the UV coordinates for
F x T x T points.
Returns F x T x T x 2
"""
uv = obj2nmr_uvmap(ft, vt, tex_size=tex_size)
uv = uv.reshape(-1, tex_size, tex_size, 2)
return uv
def obj2nmr_uvmap(ft, vt, tex_size=6):
"""
Converts obj uv_map to NMR uv_map (F x T x T x 2),
where tex_size (T) is the sample rate on each face.
"""
# This is F x 3 x 2
uv_map_for_verts = vt[ft]
# obj's y coordinate is [1-0], but image is [0-1]
uv_map_for_verts[:, :, 1] = 1 - uv_map_for_verts[:, :, 1]
# range [0, 1] -> [-1, 1]
uv_map_for_verts = (2 * uv_map_for_verts) - 1
alpha = np.arange(tex_size, dtype=float) / (tex_size - 1)
beta = np.arange(tex_size, dtype=float) / (tex_size - 1)
import itertools
# Barycentric coordinate values
coords = np.stack([p for p in itertools.product(*[alpha, beta])])
# Compute alpha, beta (this is the same order as NMR)
v2 = uv_map_for_verts[:, 2]
v0v2 = uv_map_for_verts[:, 0] - uv_map_for_verts[:, 2]
v1v2 = uv_map_for_verts[:, 1] - uv_map_for_verts[:, 2]
# Interpolate the vertex uv values: F x 2 x T*2
uv_map = np.dstack([v0v2, v1v2]).dot(coords.T) + v2.reshape(-1, 2, 1)
# F x T*2 x 2 -> F x T x T x 2
uv_map = np.transpose(uv_map, (0, 2, 1)).reshape(-1, tex_size, tex_size, 2)
return uv_map
@@ -0,0 +1,93 @@
import detectron2.data.transforms as T
import torch
from detectron2.checkpoint import DetectionCheckpointer
from detectron2.config import CfgNode, instantiate
from detectron2.data import MetadataCatalog
from omegaconf import OmegaConf
class DefaultPredictor_Lazy:
"""Create a simple end-to-end predictor with the given config that runs on single device for a
single input image.
Compared to using the model directly, this class does the following additions:
1. Load checkpoint from the weights specified in config (cfg.MODEL.WEIGHTS).
2. Always take BGR image as the input and apply format conversion internally.
3. Apply resizing defined by the config (`cfg.INPUT.{MIN,MAX}_SIZE_TEST`).
4. Take one input image and produce a single output, instead of a batch.
This is meant for simple demo purposes, so it does the above steps automatically.
This is not meant for benchmarks or running complicated inference logic.
If you'd like to do anything more complicated, please refer to its source code as
examples to build and use the model manually.
Attributes:
metadata (Metadata): the metadata of the underlying dataset, obtained from
test dataset name in the config.
Examples:
::
pred = DefaultPredictor(cfg)
inputs = cv2.imread("input.jpg")
outputs = pred(inputs)
"""
def __init__(self, cfg):
"""
Args:
cfg: a yacs CfgNode or a omegaconf dict object.
"""
if isinstance(cfg, CfgNode):
self.cfg = cfg.clone() # cfg can be modified by model
self.model = build_model(self.cfg) # noqa: F821
if len(cfg.DATASETS.TEST):
test_dataset = cfg.DATASETS.TEST[0]
checkpointer = DetectionCheckpointer(self.model)
checkpointer.load(cfg.MODEL.WEIGHTS)
self.aug = T.ResizeShortestEdge(
[cfg.INPUT.MIN_SIZE_TEST, cfg.INPUT.MIN_SIZE_TEST], cfg.INPUT.MAX_SIZE_TEST
)
self.input_format = cfg.INPUT.FORMAT
else: # new LazyConfig
self.cfg = cfg
self.model = instantiate(cfg.model)
test_dataset = OmegaConf.select(cfg, "dataloader.test.dataset.names", default=None)
if isinstance(test_dataset, (list, tuple)):
test_dataset = test_dataset[0]
checkpointer = DetectionCheckpointer(self.model)
checkpointer.load(OmegaConf.select(cfg, "train.init_checkpoint", default=""))
mapper = instantiate(cfg.dataloader.test.mapper)
self.aug = mapper.augmentations
self.input_format = mapper.image_format
self.model.eval().cuda()
if test_dataset:
self.metadata = MetadataCatalog.get(test_dataset)
assert self.input_format in ["RGB", "BGR"], self.input_format
def __call__(self, original_image):
"""
Args:
original_image (np.ndarray): an image of shape (H, W, C) (in BGR order).
Returns:
predictions (dict):
the output of the model for one image only.
See :doc:`/tutorials/models` for details about the format.
"""
with torch.no_grad():
if self.input_format == "RGB":
original_image = original_image[:, :, ::-1]
height, width = original_image.shape[:2]
image = self.aug(T.AugInput(original_image)).apply_image(original_image)
image = torch.as_tensor(image.astype("float32").transpose(2, 0, 1))
inputs = {"image": image, "height": height, "width": width}
predictions = self.model([inputs])[0]
return predictions
View File
+592
View File
@@ -0,0 +1,592 @@
import os
from typing import List, Union
import numpy as np
import math
import time
import heapq
import torch
from torch import Tensor, nn
from torch.distributions.distribution import Distribution
from transformers import AutoModelForSeq2SeqLM, T5ForConditionalGeneration, T5Tokenizer, AutoTokenizer, GPT2LMHeadModel, GPT2Tokenizer
import random
from typing import Optional
from .tools.token_emb import NewTokenEmb
class MLM(nn.Module):
def __init__(
self,
model_path: str,
model_type: str = "t5",
stage: str = "lm_pretrain",
new_token_type: str = "insert",
motion_codebook_size: int = 512,
framerate: float = 20.0,
down_t: int = 4,
predict_ratio: float = 0.2,
inbetween_ratio: float = 0.25,
max_length: int = 256,
lora: bool = False,
quota_ratio: float = 0.5,
noise_density: float = 0.15,
mean_noise_span_length: int = 3,
**kwargs,
) -> None:
super().__init__()
# Parameters
self.m_codebook_size = motion_codebook_size
self.max_length = max_length
self.framerate = framerate
self.down_t = down_t
self.predict_ratio = predict_ratio
self.inbetween_ratio = inbetween_ratio
self.noise_density = noise_density
self.mean_noise_span_length = mean_noise_span_length
self.quota_ratio = quota_ratio
self.stage = stage
# Instantiate language model
self.tokenizer = AutoTokenizer.from_pretrained(model_path, legacy=True)
if model_type == "t5":
self.language_model = T5ForConditionalGeneration.from_pretrained(
model_path)
self.lm_type = 'encdec'
elif model_type == "gpt2":
self.language_model = GPT2LMHeadModel.from_pretrained(model_path)
self.lm_type = 'dec'
else:
raise ValueError("type must be either seq2seq or conditional")
if self.lm_type == 'dec':
self.tokenizer.pad_token = self.tokenizer.eos_token
# Add motion tokens
self.tokenizer.add_tokens(
[f'<motion_id_{i}>' for i in range(self.m_codebook_size + 3)])
if new_token_type == "insert":
self.language_model.resize_token_embeddings(len(self.tokenizer))
elif new_token_type == "mlp":
shared = NewTokenEmb(self.language_model.shared,
self.m_codebook_size + 3)
# lm_head = NewTokenEmb(self.language_model.lm_head,
# self.m_codebook_size + 3)
self.language_model.resize_token_embeddings(len(self.tokenizer))
self.language_model.shared = shared
# self.language_model.lm_head = lm_head
# Lora
if lora:
from peft import LoraConfig, TaskType, get_peft_model, get_peft_model_state_dict
from peft.utils.other import fsdp_auto_wrap_policy
peft_config = LoraConfig(
bias="none",
task_type="CAUSAL_LM",
# inference_mode=False,
r=8,
lora_alpha=16,
lora_dropout=0.05)
self.language_model = get_peft_model(self.language_model,
peft_config)
def forward(self, texts: List[str], motion_tokens: Tensor,
lengths: List[int], tasks: dict):
if self.lm_type == 'encdec':
return self.forward_encdec(texts, motion_tokens, lengths, tasks)
elif self.lm_type == 'dec':
return self.forward_dec(texts, motion_tokens, lengths, tasks)
else:
raise NotImplementedError("Only conditional_multitask supported")
def forward_encdec(
self,
texts: List[str],
motion_tokens: Tensor,
lengths: List[int],
tasks: dict,
):
# Tensor to string
motion_strings = self.motion_token_to_string(motion_tokens, lengths)
# Supervised or unsupervised
# condition = random.choice(
# ['text', 'motion', 'supervised', 'supervised', 'supervised'])
condition = random.choice(['supervised', 'supervised', 'supervised'])
if condition == 'text':
inputs = texts
outputs = texts
elif condition == 'motion':
inputs = motion_strings
outputs = motion_strings
else:
inputs, outputs = self.template_fulfill(tasks, lengths,
motion_strings, texts)
# Tokenize
source_encoding = self.tokenizer(inputs,
padding='max_length',
max_length=self.max_length,
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt")
source_attention_mask = source_encoding.attention_mask.to(
motion_tokens.device)
source_input_ids = source_encoding.input_ids.to(motion_tokens.device)
if condition in ['text', 'motion']:
batch_size, expandend_input_length = source_input_ids.shape
mask_indices = np.asarray([
self.random_spans_noise_mask(expandend_input_length)
for i in range(batch_size)
])
target_mask = ~mask_indices
input_ids_sentinel = self.create_sentinel_ids(
mask_indices.astype(np.int8))
target_sentinel = self.create_sentinel_ids(
target_mask.astype(np.int8))
labels_input_ids = self.filter_input_ids(source_input_ids,
target_sentinel)
source_input_ids = self.filter_input_ids(source_input_ids,
input_ids_sentinel)
else:
target_inputs = self.tokenizer(outputs,
padding='max_length',
max_length=self.max_length,
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt")
labels_input_ids = target_inputs.input_ids.to(motion_tokens.device)
lables_attention_mask = target_inputs.attention_mask.to(
motion_tokens.device)
labels_input_ids[labels_input_ids == 0] = -100
outputs = self.language_model(
input_ids=source_input_ids,
attention_mask=source_attention_mask
if condition == 'supervised' else None,
labels=labels_input_ids,
decoder_attention_mask=lables_attention_mask
if condition == 'supervised' else None,
)
return outputs
def forward_dec(
self,
texts: List[str],
motion_tokens: Tensor,
lengths: List[int],
tasks: dict,
):
self.tokenizer.padding_side = "right"
# Tensor to string
motion_strings = self.motion_token_to_string(motion_tokens, lengths)
# Supervised or unsupervised
condition = random.choice(
['text', 'motion', 'supervised', 'supervised', 'supervised'])
if condition == 'text':
labels = texts
elif condition == 'motion':
labels = motion_strings
else:
inputs, outputs = self.template_fulfill(tasks, lengths,
motion_strings, texts)
labels = []
for i in range(len(inputs)):
labels.append(inputs[i] + ' \n ' + outputs[i] +
self.tokenizer.eos_token)
# Tokenize
inputs = self.tokenizer(labels,
padding='max_length',
max_length=self.max_length,
truncation=True,
return_attention_mask=True,
return_tensors="pt")
labels_input_ids = inputs.input_ids.to(motion_tokens.device)
lables_attention_mask = inputs.attention_mask.to(motion_tokens.device)
# print(labels_input_ids[0:5])
outputs = self.language_model(input_ids=labels_input_ids,
attention_mask=lables_attention_mask,
labels=inputs["input_ids"])
return outputs
def generate_direct(self,
texts: List[str],
max_length: int = 256,
num_beams: int = 1,
do_sample: bool = True,
bad_words_ids: List[int] = None):
# Device
self.device = self.language_model.device
# Tokenize
if self.lm_type == 'dec':
texts = [text + " \n " for text in texts]
source_encoding = self.tokenizer(texts,
padding='max_length',
max_length=self.max_length,
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt")
source_input_ids = source_encoding.input_ids.to(self.device)
source_attention_mask = source_encoding.attention_mask.to(self.device)
if self.lm_type == 'encdec':
outputs = self.language_model.generate(
source_input_ids,
max_length=max_length,
num_beams=num_beams,
do_sample=do_sample,
bad_words_ids=bad_words_ids,
)
elif self.lm_type == 'dec':
outputs = self.language_model.generate(
input_ids=source_input_ids,
attention_mask=source_attention_mask,
pad_token_id=self.tokenizer.pad_token_id,
do_sample=do_sample,
max_new_tokens=max_length)
self.tokenizer.padding_side = 'left'
outputs_string = self.tokenizer.batch_decode(outputs,
skip_special_tokens=True)
print(texts[:2])
print(outputs_string[:2])
outputs_tokens, cleaned_text = self.motion_string_to_token(
outputs_string)
return outputs_tokens, cleaned_text
def generate_conditional(self,
texts: Optional[List[str]] = None,
motion_tokens: Optional[Tensor] = None,
lengths: Optional[List[int]] = None,
task: str = "t2m",
with_len: bool = False,
stage: str = 'train',
tasks: dict = None):
self.device = self.language_model.device
if task in ["t2m", "m2m", "pred", "inbetween"]:
if task == "t2m":
assert texts is not None
motion_strings = [''] * len(texts)
if not with_len:
if tasks is None:
tasks = [{
'input':
['Generate motion: <Caption_Placeholder>'],
'output': ['']
}] * len(texts)
lengths = [0] * len(texts)
else:
tasks = [{
'input': [
'Generate motion with <Frame_Placeholder> frames: <Caption_Placeholder>'
],
'output': ['']
}] * len(texts)
elif task == "pred":
assert motion_tokens is not None and lengths is not None
texts = [''] * len(lengths)
tasks = [{
'input': ['Predict motion: <Motion_Placeholder_s1>'],
'output': ['']
}] * len(lengths)
motion_strings_old = self.motion_token_to_string(
motion_tokens, lengths)
motion_strings = []
for i, length in enumerate(lengths):
split = length // 5
motion_strings.append(
'>'.join(motion_strings_old[i].split('>')[:split]) +
'>')
elif task == "inbetween":
assert motion_tokens is not None and lengths is not None
texts = [''] * len(lengths)
tasks = [{
'input': [
"Complete the masked motion: <Motion_Placeholder_Masked>"
],
'output': ['']
}] * len(lengths)
motion_strings = self.motion_token_to_string(
motion_tokens, lengths)
inputs, outputs = self.template_fulfill(tasks, lengths,
motion_strings, texts,
stage)
outputs_tokens, cleaned_text = self.generate_direct(inputs,
max_length=128,
num_beams=1,
do_sample=True)
return outputs_tokens
elif task == "m2t":
assert motion_tokens is not None and lengths is not None
motion_strings = self.motion_token_to_string(
motion_tokens, lengths)
if not with_len:
tasks = [{
'input': ['Generate text: <Motion_Placeholder>'],
'output': ['']
}] * len(lengths)
else:
tasks = [{
'input': [
'Generate text with <Frame_Placeholder> frames: <Motion_Placeholder>'
],
'output': ['']
}] * len(lengths)
texts = [''] * len(lengths)
inputs, outputs = self.template_fulfill(tasks, lengths,
motion_strings, texts)
outputs_tokens, cleaned_text = self.generate_direct(
inputs,
max_length=40,
num_beams=1,
do_sample=False,
# bad_words_ids=self.bad_words_ids
)
return cleaned_text
def motion_token_to_string(self, motion_token: Tensor, lengths: List[int]):
motion_string = []
for i in range(len(motion_token)):
motion_i = motion_token[i].cpu(
) if motion_token[i].device.type == 'cuda' else motion_token[i]
motion_list = motion_i.tolist()[:lengths[i]]
motion_string.append(
(f'<motion_id_{self.m_codebook_size}>' +
''.join([f'<motion_id_{int(i)}>' for i in motion_list]) +
f'<motion_id_{self.m_codebook_size + 1}>'))
return motion_string
def motion_token_list_to_string(self, motion_token: Tensor):
motion_string = []
for i in range(len(motion_token)):
motion_i = motion_token[i].cpu(
) if motion_token[i].device.type == 'cuda' else motion_token[i]
motion_list = motion_i.tolist()
motion_string.append(
(f'<motion_id_{self.m_codebook_size}>' +
''.join([f'<motion_id_{int(i)}>' for i in motion_list]) +
f'<motion_id_{self.m_codebook_size + 1}>'))
return motion_string
def motion_string_to_token(self, motion_string: List[str]):
motion_tokens = []
output_string = []
for i in range(len(motion_string)):
string = self.get_middle_str(
motion_string[i], f'<motion_id_{self.m_codebook_size}>',
f'<motion_id_{self.m_codebook_size + 1}>')
string_list = string.split('><')
token_list = [
int(i.split('_')[-1].replace('>', ''))
for i in string_list[1:-1]
]
if len(token_list) == 0:
token_list = [0]
token_list_padded = torch.tensor(token_list,
dtype=int).to(self.device)
motion_tokens.append(token_list_padded)
output_string.append(motion_string[i].replace(
string, '<Motion_Placeholder>'))
return motion_tokens, output_string
def placeholder_fulfill(self, prompt: str, length: int, motion_string: str,
text: str):
seconds = math.floor(length / self.framerate)
motion_splited = motion_string.split('>')
token_length = length / self.down_t
predict_head = int(token_length * self.predict_ratio + 1)
masked_head = int(token_length * self.inbetween_ratio + 1)
masked_tail = int(token_length * (1 - self.inbetween_ratio) + 1)
motion_predict_head = '>'.join(
motion_splited[:predict_head]
) + f'><motion_id_{self.m_codebook_size+1}>'
motion_predict_last = f'<motion_id_{self.m_codebook_size}>' + '>'.join(
motion_splited[predict_head:])
motion_masked = '>'.join(
motion_splited[:masked_head]
) + '>' + f'<motion_id_{self.m_codebook_size+2}>' * (
masked_tail - masked_head) + '>'.join(motion_splited[masked_tail:])
if random.random() < self.quota_ratio:
text = f'\"{text}\"'
prompt = prompt.replace('<Caption_Placeholder>', text).replace(
'<Motion_Placeholder>',
motion_string).replace('<Frame_Placeholder>', f'{length}').replace(
'<Second_Placeholder>', '%.1f' % seconds).replace(
'<Motion_Placeholder_s1>', motion_predict_head).replace(
'<Motion_Placeholder_s2>',
motion_predict_last).replace(
'<Motion_Placeholder_Masked>', motion_masked)
return prompt
def template_fulfill(self,
tasks,
lengths,
motion_strings,
texts,
stage='test'):
inputs = []
outputs = []
for i in range(len(lengths)):
input_template = random.choice(tasks[i]['input'])
output_template = random.choice(tasks[i]['output'])
length = lengths[i]
inputs.append(
self.placeholder_fulfill(input_template, length,
motion_strings[i], texts[i]))
outputs.append(
self.placeholder_fulfill(output_template, length,
motion_strings[i], texts[i]))
return inputs, outputs
def get_middle_str(self, content, startStr, endStr):
try:
startIndex = content.index(startStr)
if startIndex >= 0:
startIndex += len(startStr)
endIndex = content.index(endStr)
except:
return f'<motion_id_{self.m_codebook_size}><motion_id_0><motion_id_{self.m_codebook_size+1}>'
return f'<motion_id_{self.m_codebook_size}>' + content[
startIndex:endIndex] + f'<motion_id_{self.m_codebook_size+1}>'
def random_spans_noise_mask(self, length):
# From https://github.com/google-research/text-to-text-transfer-transformer/blob/84f8bcc14b5f2c03de51bd3587609ba8f6bbd1cd/t5/data/preprocessors.py
orig_length = length
num_noise_tokens = int(np.round(length * self.noise_density))
# avoid degeneracy by ensuring positive numbers of noise and nonnoise tokens.
num_noise_tokens = min(max(num_noise_tokens, 1), length - 1)
num_noise_spans = int(
np.round(num_noise_tokens / self.mean_noise_span_length))
# avoid degeneracy by ensuring positive number of noise spans
num_noise_spans = max(num_noise_spans, 1)
num_nonnoise_tokens = length - num_noise_tokens
# pick the lengths of the noise spans and the non-noise spans
def _random_segmentation(num_items, num_segments):
"""Partition a sequence of items randomly into non-empty segments.
Args:
num_items: an integer scalar > 0
num_segments: an integer scalar in [1, num_items]
Returns:
a Tensor with shape [num_segments] containing positive integers that add
up to num_items
"""
mask_indices = np.arange(num_items - 1) < (num_segments - 1)
np.random.shuffle(mask_indices)
first_in_segment = np.pad(mask_indices, [[1, 0]])
segment_id = np.cumsum(first_in_segment)
# count length of sub segments assuming that list is sorted
_, segment_length = np.unique(segment_id, return_counts=True)
return segment_length
noise_span_lengths = _random_segmentation(num_noise_tokens,
num_noise_spans)
nonnoise_span_lengths = _random_segmentation(num_nonnoise_tokens,
num_noise_spans)
interleaved_span_lengths = np.reshape(
np.stack([nonnoise_span_lengths, noise_span_lengths], axis=1),
[num_noise_spans * 2],
)
span_starts = np.cumsum(interleaved_span_lengths)[:-1]
span_start_indicator = np.zeros((length, ), dtype=np.int8)
span_start_indicator[span_starts] = True
span_num = np.cumsum(span_start_indicator)
is_noise = np.equal(span_num % 2, 1)
return is_noise[:orig_length]
def create_sentinel_ids(self, mask_indices):
# From https://github.com/huggingface/transformers/blob/main/examples/flax/language-modeling/run_t5_mlm_flax.py
start_indices = mask_indices - np.roll(mask_indices, 1,
axis=-1) * mask_indices
start_indices[:, 0] = mask_indices[:, 0]
sentinel_ids = np.where(start_indices != 0,
np.cumsum(start_indices, axis=-1),
start_indices)
sentinel_ids = np.where(sentinel_ids != 0,
(len(self.tokenizer) - sentinel_ids), 0)
sentinel_ids -= mask_indices - start_indices
return sentinel_ids
def filter_input_ids(self, input_ids, sentinel_ids):
# From https://github.com/huggingface/transformers/blob/main/examples/flax/language-modeling/run_t5_mlm_flax.py
batch_size = input_ids.shape[0]
input_ids_full = np.where(sentinel_ids != 0, sentinel_ids,
input_ids.to('cpu'))
# input_ids tokens and sentinel tokens are >= 0, tokens < 0 are
# masked tokens coming after sentinel tokens and should be removed
input_ids = input_ids_full[input_ids_full >= 0].reshape(
(batch_size, -1))
input_ids = np.concatenate(
[
input_ids,
np.full((batch_size, 1),
self.tokenizer.eos_token_id,
dtype=np.int32),
],
axis=-1,
)
input_ids = torch.tensor(input_ids, device=self.device)
return input_ids
+190
View File
@@ -0,0 +1,190 @@
# Partially from https://github.com/Mael-zys/T2M-GPT
from typing import List, Optional, Union
import torch
import torch.nn as nn
from torch import Tensor, nn
from torch.distributions.distribution import Distribution
from .tools.resnet import Resnet1D
from .tools.quantize_cnn import QuantizeEMAReset, Quantizer, QuantizeEMA, QuantizeReset
from collections import OrderedDict
class VQVae(nn.Module):
def __init__(self,
nfeats: int,
quantizer: str = "ema_reset",
code_num=512,
code_dim=512,
output_emb_width=512,
down_t=3,
stride_t=2,
width=512,
depth=3,
dilation_growth_rate=3,
norm=None,
activation: str = "relu",
**kwargs) -> None:
super().__init__()
self.code_dim = code_dim
self.encoder = Encoder(nfeats,
output_emb_width,
down_t,
stride_t,
width,
depth,
dilation_growth_rate,
activation=activation,
norm=norm)
self.decoder = Decoder(nfeats,
output_emb_width,
down_t,
stride_t,
width,
depth,
dilation_growth_rate,
activation=activation,
norm=norm)
if quantizer == "ema_reset":
self.quantizer = QuantizeEMAReset(code_num, code_dim, mu=0.99)
elif quantizer == "orig":
self.quantizer = Quantizer(code_num, code_dim, beta=1.0)
elif quantizer == "ema":
self.quantizer = QuantizeEMA(code_num, code_dim, mu=0.99)
elif quantizer == "reset":
self.quantizer = QuantizeReset(code_num, code_dim)
def preprocess(self, x):
# (bs, T, Jx3) -> (bs, Jx3, T)
x = x.permute(0, 2, 1)
return x
def postprocess(self, x):
# (bs, Jx3, T) -> (bs, T, Jx3)
x = x.permute(0, 2, 1)
return x
def forward(self, features: Tensor):
# Preprocess
x_in = self.preprocess(features)
# 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 encode(
self,
features: Tensor,
) -> Union[Tensor, Distribution]:
N, T, _ = features.shape
x_in = self.preprocess(features)
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)
# latent, dist
return code_idx, None
def decode(self, z: Tensor):
x_d = self.quantizer.dequantize(z)
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 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)
@@ -0,0 +1,111 @@
import torch
import torch.nn as nn
from torch.nn.utils.rnn import pack_padded_sequence
class MovementConvEncoder(nn.Module):
def __init__(self, input_size, hidden_size, output_size):
super(MovementConvEncoder, self).__init__()
self.main = nn.Sequential(
nn.Conv1d(input_size, hidden_size, 4, 2, 1),
nn.Dropout(0.2, inplace=True),
nn.LeakyReLU(0.2, inplace=True),
nn.Conv1d(hidden_size, output_size, 4, 2, 1),
nn.Dropout(0.2, inplace=True),
nn.LeakyReLU(0.2, inplace=True),
)
self.out_net = nn.Linear(output_size, output_size)
# self.main.apply(init_weight)
# self.out_net.apply(init_weight)
def forward(self, inputs):
inputs = inputs.permute(0, 2, 1)
outputs = self.main(inputs).permute(0, 2, 1)
# print(outputs.shape)
return self.out_net(outputs)
class MotionEncoderBiGRUCo(nn.Module):
def __init__(self, input_size, hidden_size, output_size):
super(MotionEncoderBiGRUCo, self).__init__()
self.input_emb = nn.Linear(input_size, hidden_size)
self.gru = nn.GRU(
hidden_size, hidden_size, batch_first=True, bidirectional=True
)
self.output_net = nn.Sequential(
nn.Linear(hidden_size * 2, hidden_size),
nn.LayerNorm(hidden_size),
nn.LeakyReLU(0.2, inplace=True),
nn.Linear(hidden_size, output_size),
)
# self.input_emb.apply(init_weight)
# self.output_net.apply(init_weight)
self.hidden_size = hidden_size
self.hidden = nn.Parameter(
torch.randn((2, 1, self.hidden_size), requires_grad=True)
)
# input(batch_size, seq_len, dim)
def forward(self, inputs, m_lens):
num_samples = inputs.shape[0]
input_embs = self.input_emb(inputs)
hidden = self.hidden.repeat(1, num_samples, 1)
cap_lens = m_lens.data.tolist()
# emb = pack_padded_sequence(input=input_embs, lengths=cap_lens, batch_first=True)
emb = input_embs
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 TextEncoderBiGRUCo(nn.Module):
def __init__(self, word_size, pos_size, hidden_size, output_size):
super(TextEncoderBiGRUCo, self).__init__()
self.pos_emb = nn.Linear(pos_size, word_size)
self.input_emb = nn.Linear(word_size, hidden_size)
self.gru = nn.GRU(
hidden_size, hidden_size, batch_first=True, bidirectional=True
)
self.output_net = nn.Sequential(
nn.Linear(hidden_size * 2, hidden_size),
nn.LayerNorm(hidden_size),
nn.LeakyReLU(0.2, inplace=True),
nn.Linear(hidden_size, output_size),
)
# self.input_emb.apply(init_weight)
# self.pos_emb.apply(init_weight)
# self.output_net.apply(init_weight)
# self.linear2.apply(init_weight)
# self.batch_size = batch_size
self.hidden_size = hidden_size
self.hidden = nn.Parameter(
torch.randn((2, 1, self.hidden_size), requires_grad=True)
)
# input(batch_size, seq_len, dim)
def forward(self, word_embs, pos_onehot, cap_lens):
num_samples = word_embs.shape[0]
pos_embs = self.pos_emb(pos_onehot)
inputs = word_embs + pos_embs
input_embs = self.input_emb(inputs)
hidden = self.hidden.repeat(1, num_samples, 1)
cap_lens = cap_lens.data.tolist()
emb = pack_padded_sequence(input=input_embs, lengths=cap_lens, batch_first=True)
gru_seq, gru_last = self.gru(emb, hidden)
gru_last = torch.cat([gru_last[0], gru_last[1]], dim=-1)
return self.output_net(gru_last)
@@ -0,0 +1,322 @@
# This file is taken from signjoey repository
import math
import torch
from torch import Tensor, nn
def get_activation(activation_type):
if activation_type == "relu":
return nn.ReLU()
elif activation_type == "relu6":
return nn.ReLU6()
elif activation_type == "prelu":
return nn.PReLU()
elif activation_type == "selu":
return nn.SELU()
elif activation_type == "celu":
return nn.CELU()
elif activation_type == "gelu":
return nn.GELU()
elif activation_type == "sigmoid":
return nn.Sigmoid()
elif activation_type == "softplus":
return nn.Softplus()
elif activation_type == "softshrink":
return nn.Softshrink()
elif activation_type == "softsign":
return nn.Softsign()
elif activation_type == "tanh":
return nn.Tanh()
elif activation_type == "tanhshrink":
return nn.Tanhshrink()
else:
raise ValueError("Unknown activation type {}".format(activation_type))
class MaskedNorm(nn.Module):
"""
Original Code from:
https://discuss.pytorch.org/t/batchnorm-for-different-sized-samples-in-batch/44251/8
"""
def __init__(self, norm_type, num_groups, num_features):
super().__init__()
self.norm_type = norm_type
if self.norm_type == "batch":
self.norm = nn.BatchNorm1d(num_features=num_features)
elif self.norm_type == "group":
self.norm = nn.GroupNorm(num_groups=num_groups, num_channels=num_features)
elif self.norm_type == "layer":
self.norm = nn.LayerNorm(normalized_shape=num_features)
else:
raise ValueError("Unsupported Normalization Layer")
self.num_features = num_features
def forward(self, x: Tensor, mask: Tensor):
if self.training:
reshaped = x.reshape([-1, self.num_features])
reshaped_mask = mask.reshape([-1, 1]) > 0
selected = torch.masked_select(reshaped, reshaped_mask).reshape(
[-1, self.num_features]
)
batch_normed = self.norm(selected)
scattered = reshaped.masked_scatter(reshaped_mask, batch_normed)
return scattered.reshape([x.shape[0], -1, self.num_features])
else:
reshaped = x.reshape([-1, self.num_features])
batched_normed = self.norm(reshaped)
return batched_normed.reshape([x.shape[0], -1, self.num_features])
# TODO (Cihan): Spatial and Word Embeddings are pretty much the same
# We might as well convert them into a single module class.
# Only difference is the lut vs linear layers.
class Embeddings(nn.Module):
"""
Simple embeddings class
"""
# pylint: disable=unused-argument
def __init__(
self,
embedding_dim: int = 64,
num_heads: int = 8,
scale: bool = False,
scale_factor: float = None,
norm_type: str = None,
activation_type: str = None,
vocab_size: int = 0,
padding_idx: int = 1,
freeze: bool = False,
**kwargs
):
"""
Create new embeddings for the vocabulary.
Use scaling for the Transformer.
:param embedding_dim:
:param scale:
:param vocab_size:
:param padding_idx:
:param freeze: freeze the embeddings during training
"""
super().__init__()
self.embedding_dim = embedding_dim
self.vocab_size = vocab_size
self.lut = nn.Embedding(vocab_size, self.embedding_dim, padding_idx=padding_idx)
self.norm_type = norm_type
if self.norm_type:
self.norm = MaskedNorm(
norm_type=norm_type, num_groups=num_heads, num_features=embedding_dim
)
self.activation_type = activation_type
if self.activation_type:
self.activation = get_activation(activation_type)
self.scale = scale
if self.scale:
if scale_factor:
self.scale_factor = scale_factor
else:
self.scale_factor = math.sqrt(self.embedding_dim)
if freeze:
freeze_params(self)
# pylint: disable=arguments-differ
def forward(self, x: Tensor, mask: Tensor = None) -> Tensor:
"""
Perform lookup for input `x` in the embedding table.
:param mask: token masks
:param x: index in the vocabulary
:return: embedded representation for `x`
"""
x = self.lut(x)
if self.norm_type:
x = self.norm(x, mask)
if self.activation_type:
x = self.activation(x)
if self.scale:
return x * self.scale_factor
else:
return x
def __repr__(self):
return "%s(embedding_dim=%d, vocab_size=%d)" % (
self.__class__.__name__,
self.embedding_dim,
self.vocab_size,
)
class SpatialEmbeddings(nn.Module):
"""
Simple Linear Projection Layer
(For encoder outputs to predict glosses)
"""
# pylint: disable=unused-argument
def __init__(
self,
embedding_dim: int,
input_size: int,
num_heads: int,
freeze: bool = False,
norm_type: str = "batch",
activation_type: str = "softsign",
scale: bool = False,
scale_factor: float = None,
**kwargs
):
"""
Create new embeddings for the vocabulary.
Use scaling for the Transformer.
:param embedding_dim:
:param input_size:
:param freeze: freeze the embeddings during training
"""
super().__init__()
self.embedding_dim = embedding_dim
self.input_size = input_size
self.ln = nn.Linear(self.input_size, self.embedding_dim)
self.norm_type = norm_type
if self.norm_type:
self.norm = MaskedNorm(
norm_type=norm_type, num_groups=num_heads, num_features=embedding_dim
)
self.activation_type = activation_type
if self.activation_type:
self.activation = get_activation(activation_type)
self.scale = scale
if self.scale:
if scale_factor:
self.scale_factor = scale_factor
else:
self.scale_factor = math.sqrt(self.embedding_dim)
if freeze:
freeze_params(self)
# pylint: disable=arguments-differ
def forward(self, x: Tensor, mask: Tensor) -> Tensor:
"""
:param mask: frame masks
:param x: input frame features
:return: embedded representation for `x`
"""
x = self.ln(x)
if self.norm_type:
x = self.norm(x, mask)
if self.activation_type:
x = self.activation(x)
if self.scale:
return x * self.scale_factor
else:
return x
def __repr__(self):
return "%s(embedding_dim=%d, input_size=%d)" % (
self.__class__.__name__,
self.embedding_dim,
self.input_size,
)
def get_timestep_embedding(
timesteps: torch.Tensor,
embedding_dim: int,
flip_sin_to_cos: bool = False,
downscale_freq_shift: float = 1,
scale: float = 1,
max_period: int = 10000,
):
"""
This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.
:param timesteps: a 1-D Tensor of N indices, one per batch element.
These may be fractional.
:param embedding_dim: the dimension of the output. :param max_period: controls the minimum frequency of the
embeddings. :return: an [N x dim] Tensor of positional embeddings.
"""
assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array"
half_dim = embedding_dim // 2
exponent = -math.log(max_period) * torch.arange(
start=0, end=half_dim, dtype=torch.float32, device=timesteps.device
)
exponent = exponent / (half_dim - downscale_freq_shift)
emb = torch.exp(exponent)
emb = timesteps[:, None].float() * emb[None, :]
# scale embeddings
emb = scale * emb
# concat sine and cosine embeddings
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
# flip sine and cosine embeddings
if flip_sin_to_cos:
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
# zero pad
if embedding_dim % 2 == 1:
emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))
return emb
class TimestepEmbedding(nn.Module):
def __init__(self, channel: int, time_embed_dim: int, act_fn: str = "silu"):
super().__init__()
self.linear_1 = nn.Linear(channel, time_embed_dim)
self.act = None
if act_fn == "silu":
self.act = nn.SiLU()
self.linear_2 = nn.Linear(time_embed_dim, time_embed_dim)
def forward(self, sample):
sample = self.linear_1(sample)
if self.act is not None:
sample = self.act(sample)
sample = self.linear_2(sample)
return sample
class Timesteps(nn.Module):
def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float):
super().__init__()
self.num_channels = num_channels
self.flip_sin_to_cos = flip_sin_to_cos
self.downscale_freq_shift = downscale_freq_shift
def forward(self, timesteps):
t_emb = get_timestep_embedding(
timesteps,
self.num_channels,
flip_sin_to_cos=self.flip_sin_to_cos,
downscale_freq_shift=self.downscale_freq_shift,
)
return t_emb
@@ -0,0 +1,414 @@
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, mu):
super().__init__()
self.nb_code = nb_code
self.code_dim = code_dim
self.mu = mu
self.reset_codebook()
def reset_codebook(self):
self.init = False
self.code_sum = None
self.code_count = None
device = "cuda" if torch.cuda.is_available() else "cpu"
self.register_buffer('codebook', torch.zeros(self.nb_code, self.code_dim).to(device))
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):
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, mu):
super().__init__()
self.nb_code = nb_code
self.code_dim = code_dim
self.mu = 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
# 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
@@ -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)
@@ -0,0 +1,73 @@
from torch import Tensor, nn
class NewTokenEmb(nn.Module):
"""
For adding new tokens to a pretrained model
"""
def __init__(self,
old_embeddings: nn.Embedding,
new_num_tokens: int = None) -> None:
super().__init__()
self.num_tokens = old_embeddings.num_embeddings + new_num_tokens
self.old_num_tokens = old_embeddings.num_embeddings
self.new_num_tokens = new_num_tokens
self.embedding_dim = old_embeddings.embedding_dim
# For text embeddings
self.text_embeddings = nn.Embedding(
self.num_tokens,
self.embedding_dim,
device=old_embeddings.weight.device,
dtype=old_embeddings.weight.dtype)
with torch.no_grad():
self.text_embeddings.weight.data[:old_embeddings.
num_embeddings] = old_embeddings.weight.data
self.text_embeddings.weight.data[
self.old_num_tokens:] = torch.zeros(
self.new_num_tokens,
self.embedding_dim,
dtype=old_embeddings.weight.dtype,
device=old_embeddings.weight.device)
self.text_embeddings.weight.requires_grad_(False)
# For motion embeddings
self.motion_embeddings = nn.Embedding(
new_num_tokens,
self.embedding_dim,
device=old_embeddings.weight.device,
dtype=old_embeddings.weight.dtype)
with torch.no_grad():
self.motion_embeddings.weight.data[:self.
old_num_tokens] = torch.zeros(
new_num_tokens,
self.embedding_dim,
dtype=old_embeddings.weight.
dtype,
device=old_embeddings.
weight.device)
self.word2motionProj = nn.Linear(self.old_num_tokens, new_num_tokens)
def forward(self, input: Tensor) -> Tensor:
with torch.no_grad():
self.motion_embeddings.weight.data[:self.
old_num_tokens] = torch.zeros(
self.new_num_tokens,
self.embedding_dim,
dtype=self.motion_embeddings
.weight.dtype,
device=self.
motion_embeddings.weight.
device)
self.motion_embeddings.weight.data[
self.old_num_tokens:] = self.word2motionProj(
self.text_embeddings.weight.data[:self.old_num_tokens].permute(
1, 0)).permute(1, 0)
return self.text_embeddings(input) + self.motion_embeddings(input)
@@ -0,0 +1,285 @@
# -*- coding: utf-8 -*-
import math
import torch
import torch.nn as nn
from torch import Tensor
# Took from https://github.com/joeynmt/joeynmt/blob/fb66afcbe1beef9acd59283bcc084c4d4c1e6343/joeynmt/transformer_layers.py
# pylint: disable=arguments-differ
class MultiHeadedAttention(nn.Module):
"""
Multi-Head Attention module from "Attention is All You Need"
Implementation modified from OpenNMT-py.
https://github.com/OpenNMT/OpenNMT-py
"""
def __init__(self, num_heads: int, size: int, dropout: float = 0.1):
"""
Create a multi-headed attention layer.
:param num_heads: the number of heads
:param size: model size (must be divisible by num_heads)
:param dropout: probability of dropping a unit
"""
super().__init__()
assert size % num_heads == 0
self.head_size = head_size = size // num_heads
self.model_size = size
self.num_heads = num_heads
self.k_layer = nn.Linear(size, num_heads * head_size)
self.v_layer = nn.Linear(size, num_heads * head_size)
self.q_layer = nn.Linear(size, num_heads * head_size)
self.output_layer = nn.Linear(size, size)
self.softmax = nn.Softmax(dim=-1)
self.dropout = nn.Dropout(dropout)
def forward(self, k: Tensor, v: Tensor, q: Tensor, mask: Tensor = None):
"""
Computes multi-headed attention.
:param k: keys [B, M, D] with M being the sentence length.
:param v: values [B, M, D]
:param q: query [B, M, D]
:param mask: optional mask [B, 1, M] or [B, M, M]
:return:
"""
batch_size = k.size(0)
num_heads = self.num_heads
# project the queries (q), keys (k), and values (v)
k = self.k_layer(k)
v = self.v_layer(v)
q = self.q_layer(q)
# reshape q, k, v for our computation to [batch_size, num_heads, ..]
k = k.view(batch_size, -1, num_heads, self.head_size).transpose(1, 2)
v = v.view(batch_size, -1, num_heads, self.head_size).transpose(1, 2)
q = q.view(batch_size, -1, num_heads, self.head_size).transpose(1, 2)
# compute scores
q = q / math.sqrt(self.head_size)
# batch x num_heads x query_len x key_len
scores = torch.matmul(q, k.transpose(2, 3))
# torch.Size([48, 8, 183, 183])
# apply the mask (if we have one)
# we add a dimension for the heads to it below: [B, 1, 1, M]
if mask is not None:
scores = scores.masked_fill(~mask.unsqueeze(1), float('-inf'))
# apply attention dropout and compute context vectors.
attention = self.softmax(scores)
attention = self.dropout(attention)
# torch.Size([48, 8, 183, 183]) [bs, nheads, time, time] (for decoding)
# v: torch.Size([48, 8, 183, 32]) (32 is 256/8)
# get context vector (select values with attention) and reshape
# back to [B, M, D]
context = torch.matmul(attention, v) # torch.Size([48, 8, 183, 32])
context = context.transpose(1, 2).contiguous().view(
batch_size, -1, num_heads * self.head_size)
# torch.Size([48, 183, 256]) put back to 256 (combine the heads)
output = self.output_layer(context)
# torch.Size([48, 183, 256]): 1 output per time step
return output
# pylint: disable=arguments-differ
class PositionwiseFeedForward(nn.Module):
"""
Position-wise Feed-forward layer
Projects to ff_size and then back down to input_size.
"""
def __init__(self, input_size, ff_size, dropout=0.1):
"""
Initializes position-wise feed-forward layer.
:param input_size: dimensionality of the input.
:param ff_size: dimensionality of intermediate representation
:param dropout:
"""
super().__init__()
self.layer_norm = nn.LayerNorm(input_size, eps=1e-6)
self.pwff_layer = nn.Sequential(
nn.Linear(input_size, ff_size),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(ff_size, input_size),
nn.Dropout(dropout),
)
def forward(self, x):
x_norm = self.layer_norm(x)
return self.pwff_layer(x_norm) + x
# pylint: disable=arguments-differ
class PositionalEncoding(nn.Module):
"""
Pre-compute position encodings (PE).
In forward pass, this adds the position-encodings to the
input for as many time steps as necessary.
Implementation based on OpenNMT-py.
https://github.com/OpenNMT/OpenNMT-py
"""
def __init__(self, size: int = 0, max_len: int = 5000):
"""
Positional Encoding with maximum length max_len
:param size:
:param max_len:
:param dropout:
"""
if size % 2 != 0:
raise ValueError("Cannot use sin/cos positional encoding with "
"odd dim (got dim={:d})".format(size))
pe = torch.zeros(max_len, size)
position = torch.arange(0, max_len).unsqueeze(1)
div_term = torch.exp((torch.arange(0, size, 2, dtype=torch.float) *
-(math.log(10000.0) / size)))
pe[:, 0::2] = torch.sin(position.float() * div_term)
pe[:, 1::2] = torch.cos(position.float() * div_term)
pe = pe.unsqueeze(0) # shape: [1, size, max_len]
super().__init__()
self.register_buffer('pe', pe)
self.dim = size
def forward(self, emb):
"""Embed inputs.
Args:
emb (FloatTensor): Sequence of word vectors
``(seq_len, batch_size, self.dim)``
"""
# Add position encodings
return emb + self.pe[:, :emb.size(1)]
class TransformerEncoderLayer(nn.Module):
"""
One Transformer encoder layer has a Multi-head attention layer plus
a position-wise feed-forward layer.
"""
def __init__(self,
size: int = 0,
ff_size: int = 0,
num_heads: int = 0,
dropout: float = 0.1):
"""
A single Transformer layer.
:param size:
:param ff_size:
:param num_heads:
:param dropout:
"""
super().__init__()
self.layer_norm = nn.LayerNorm(size, eps=1e-6)
self.src_src_att = MultiHeadedAttention(num_heads,
size,
dropout=dropout)
self.feed_forward = PositionwiseFeedForward(size,
ff_size=ff_size,
dropout=dropout)
self.dropout = nn.Dropout(dropout)
self.size = size
# pylint: disable=arguments-differ
def forward(self, x: Tensor, mask: Tensor) -> Tensor:
"""
Forward pass for a single transformer encoder layer.
First applies layer norm, then self attention,
then dropout with residual connection (adding the input to the result),
and then a position-wise feed-forward layer.
:param x: layer input
:param mask: input mask
:return: output tensor
"""
x_norm = self.layer_norm(x)
h = self.src_src_att(x_norm, x_norm, x_norm, mask)
h = self.dropout(h) + x
o = self.feed_forward(h)
return o
class TransformerDecoderLayer(nn.Module):
"""
Transformer decoder layer.
Consists of self-attention, source-attention, and feed-forward.
"""
def __init__(self,
size: int = 0,
ff_size: int = 0,
num_heads: int = 0,
dropout: float = 0.1):
"""
Represents a single Transformer decoder layer.
It attends to the source representation and the previous decoder states.
:param size: model dimensionality
:param ff_size: size of the feed-forward intermediate layer
:param num_heads: number of heads
:param dropout: dropout to apply to input
"""
super().__init__()
self.size = size
self.trg_trg_att = MultiHeadedAttention(num_heads,
size,
dropout=dropout)
self.src_trg_att = MultiHeadedAttention(num_heads,
size,
dropout=dropout)
self.feed_forward = PositionwiseFeedForward(size,
ff_size=ff_size,
dropout=dropout)
self.x_layer_norm = nn.LayerNorm(size, eps=1e-6)
self.dec_layer_norm = nn.LayerNorm(size, eps=1e-6)
self.dropout = nn.Dropout(dropout)
# pylint: disable=arguments-differ
def forward(self,
x: Tensor = None,
memory: Tensor = None,
src_mask: Tensor = None,
trg_mask: Tensor = None) -> Tensor:
"""
Forward pass of a single Transformer decoder layer.
:param x: inputs
:param memory: source representations
:param src_mask: source mask
:param trg_mask: target mask (so as to not condition on future steps)
:return: output tensor
"""
# decoder/target self-attention
x_norm = self.x_layer_norm(x) # torch.Size([48, 183, 256])
h1 = self.trg_trg_att(x_norm, x_norm, x_norm, mask=trg_mask)
h1 = self.dropout(h1) + x
# source-target attention
h1_norm = self.dec_layer_norm(
h1) # torch.Size([48, 183, 256]) (same for memory)
h2 = self.src_trg_att(memory, memory, h1_norm, mask=src_mask)
# final position-wise feed-forward layer
o = self.feed_forward(self.dropout(h2) + h1)
return o
+200
View File
@@ -0,0 +1,200 @@
import os
from pytorch_lightning import LightningModule, Trainer
from pytorch_lightning.callbacks import Callback, RichProgressBar, ModelCheckpoint
def build_callbacks(cfg, logger=None, phase='test', **kwargs):
callbacks = []
logger = logger
# Rich Progress Bar
callbacks.append(progressBar())
# Checkpoint Callback
if phase == 'train':
callbacks.extend(getCheckpointCallback(cfg, logger=logger, **kwargs))
return callbacks
def getCheckpointCallback(cfg, logger=None, **kwargs):
callbacks = []
# Logging
metric_monitor = {
"loss_total": "total/train",
"Train_jf": "recons/text2jfeats/train",
"Val_jf": "recons/text2jfeats/val",
"Train_rf": "recons/text2rfeats/train",
"Val_rf": "recons/text2rfeats/val",
"APE root": "Metrics/APE_root",
"APE mean pose": "Metrics/APE_mean_pose",
"AVE root": "Metrics/AVE_root",
"AVE mean pose": "Metrics/AVE_mean_pose",
"R_TOP_1": "Metrics/R_precision_top_1",
"R_TOP_2": "Metrics/R_precision_top_2",
"R_TOP_3": "Metrics/R_precision_top_3",
"gt_R_TOP_3": "Metrics/gt_R_precision_top_3",
"FID": "Metrics/FID",
"gt_FID": "Metrics/gt_FID",
"Diversity": "Metrics/Diversity",
"MM dist": "Metrics/Matching_score",
"Accuracy": "Metrics/accuracy",
}
callbacks.append(
progressLogger(logger,metric_monitor=metric_monitor,log_every_n_steps=1))
# Save 10 latest checkpoints
checkpointParams = {
'dirpath': os.path.join(cfg.FOLDER_EXP, "checkpoints"),
'filename': "{epoch}",
'monitor': "step",
'mode': "max",
'every_n_epochs': cfg.LOGGER.VAL_EVERY_STEPS,
'save_top_k': 8,
'save_last': True,
'save_on_train_epoch_end': True
}
callbacks.append(ModelCheckpoint(**checkpointParams))
# Save checkpoint every n*10 epochs
checkpointParams.update({
'every_n_epochs':
cfg.LOGGER.VAL_EVERY_STEPS * 10,
'save_top_k':
-1,
'save_last':
False
})
callbacks.append(ModelCheckpoint(**checkpointParams))
metrics = cfg.METRIC.TYPE
metric_monitor_map = {
'TemosMetric': {
'Metrics/APE_root': {
'abbr': 'APEroot',
'mode': 'min'
},
},
'TM2TMetrics': {
'Metrics/FID': {
'abbr': 'FID',
'mode': 'min'
},
'Metrics/R_precision_top_3': {
'abbr': 'R3',
'mode': 'max'
}
},
'MRMetrics': {
'Metrics/MPJPE': {
'abbr': 'MPJPE',
'mode': 'min'
}
},
'HUMANACTMetrics': {
'Metrics/Accuracy': {
'abbr': 'Accuracy',
'mode': 'max'
}
},
'UESTCMetrics': {
'Metrics/Accuracy': {
'abbr': 'Accuracy',
'mode': 'max'
}
},
'UncondMetrics': {
'Metrics/FID': {
'abbr': 'FID',
'mode': 'min'
}
}
}
checkpointParams.update({
'every_n_epochs': cfg.LOGGER.VAL_EVERY_STEPS,
'save_top_k': 1,
})
for metric in metrics:
if metric in metric_monitor_map.keys():
metric_monitors = metric_monitor_map[metric]
# Delete R3 if training VAE
if cfg.TRAIN.STAGE == 'vae' and metric == 'TM2TMetrics':
del metric_monitors['Metrics/R_precision_top_3']
for metric_monitor in metric_monitors:
checkpointParams.update({
'filename':
metric_monitor_map[metric][metric_monitor]['mode']
+ "-" +
metric_monitor_map[metric][metric_monitor]['abbr']
+ "{ep}",
'monitor':
metric_monitor,
'mode':
metric_monitor_map[metric][metric_monitor]['mode'],
})
callbacks.append(
ModelCheckpoint(**checkpointParams))
return callbacks
class progressBar(RichProgressBar):
def __init__(self, ):
super().__init__()
def get_metrics(self, trainer, model):
# Don't show the version number
items = super().get_metrics(trainer, model)
items.pop("v_num", None)
return items
class progressLogger(Callback):
def __init__(self,
logger,
metric_monitor: dict,
precision: int = 3,
log_every_n_steps: int = 1):
# Metric to monitor
self.logger = logger
self.metric_monitor = metric_monitor
self.precision = precision
self.log_every_n_steps = log_every_n_steps
def on_train_start(self, trainer: Trainer, pl_module: LightningModule,
**kwargs) -> None:
self.logger.info("Training started")
def on_train_end(self, trainer: Trainer, pl_module: LightningModule,
**kwargs) -> None:
self.logger.info("Training done")
def on_validation_epoch_end(self, trainer: Trainer,
pl_module: LightningModule, **kwargs) -> None:
if trainer.sanity_checking:
self.logger.info("Sanity checking ok.")
def on_train_epoch_end(self,
trainer: Trainer,
pl_module: LightningModule,
padding=False,
**kwargs) -> None:
metric_format = f"{{:.{self.precision}e}}"
line = f"Epoch {trainer.current_epoch}"
if padding:
line = f"{line:>{len('Epoch xxxx')}}" # Right padding
if trainer.current_epoch % self.log_every_n_steps == 0:
metrics_str = []
losses_dict = trainer.callback_metrics
for metric_name, dico_name in self.metric_monitor.items():
if dico_name in losses_dict:
metric = losses_dict[dico_name].item()
metric = metric_format.format(metric)
metric = f"{metric_name} {metric}"
metrics_str.append(metric)
line = line + ": " + " ".join(metrics_str)
self.logger.info(line)
+217
View File
@@ -0,0 +1,217 @@
import importlib
from argparse import ArgumentParser
from omegaconf import OmegaConf
from os.path import join as pjoin
import os
import glob
def get_module_config(cfg, filepath="./configs"):
"""
Load yaml config files from subfolders
"""
yamls = glob.glob(pjoin(filepath, '*', '*.yaml'))
yamls = [y.replace(filepath, '') for y in yamls]
for yaml in yamls:
nodes = yaml.replace('.yaml', '').replace('/', '.')
nodes = nodes[1:] if nodes[0] == '.' else nodes
OmegaConf.update(cfg, nodes, OmegaConf.load('./configs' + yaml))
return cfg
def get_obj_from_str(string, reload=False):
"""
Get object from string
"""
package_directory_name = os.path.basename(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
module, cls = string.rsplit(".", 1)
if reload:
module_imp = importlib.import_module(module)
importlib.reload(module_imp)
return getattr(importlib.import_module(module, package=package_directory_name), cls)
def instantiate_from_config(config):
"""
Instantiate object from config
"""
if not "target" in config:
raise KeyError("Expected key `target` to instantiate.")
return get_obj_from_str(config["target"])(**config.get("params", dict()))
def resume_config(cfg: OmegaConf):
"""
Resume model and wandb
"""
if cfg.TRAIN.RESUME:
resume = cfg.TRAIN.RESUME
if os.path.exists(resume):
# Checkpoints
cfg.TRAIN.PRETRAINED = pjoin(resume, "checkpoints", "last.ckpt")
# Wandb
wandb_files = os.listdir(pjoin(resume, "wandb", "latest-run"))
wandb_run = [item for item in wandb_files if "run-" in item][0]
cfg.LOGGER.WANDB.params.id = wandb_run.replace("run-","").replace(".wandb", "")
else:
raise ValueError("Resume path is not right.")
return cfg
def parse_args(phase="train"):
"""
Parse arguments and load config files
"""
parser = ArgumentParser()
group = parser.add_argument_group("Training options")
# Assets
group.add_argument(
"--cfg_assets",
type=str,
required=False,
default="./configs/assets.yaml",
help="config file for asset paths",
)
# Default config
if phase in ["train", "test"]:
cfg_defualt = "./configs/default.yaml"
elif phase == "render":
cfg_defualt = "./configs/render.yaml"
elif phase == "webui":
cfg_defualt = "./configs/webui.yaml"
group.add_argument(
"--cfg",
type=str,
required=False,
default=cfg_defualt,
help="config file",
)
# Parse for each phase
if phase in ["train", "test"]:
group.add_argument("--batch_size",
type=int,
required=False,
help="training batch size")
group.add_argument("--num_nodes",
type=int,
required=False,
help="number of nodes")
group.add_argument("--device",
type=int,
nargs="+",
required=False,
help="training device")
group.add_argument("--task",
type=str,
required=False,
help="evaluation task type")
group.add_argument("--nodebug",
action="store_true",
required=False,
help="debug or not")
if phase == "demo":
group.add_argument(
"--example",
type=str,
required=False,
help="input text and lengths with txt format",
)
group.add_argument(
"--out_dir",
type=str,
required=False,
help="output dir",
)
group.add_argument("--task",
type=str,
required=False,
help="evaluation task type")
if phase == "render":
group.add_argument("--npy",
type=str,
required=False,
default=None,
help="npy motion files")
group.add_argument("--dir",
type=str,
required=False,
default=None,
help="npy motion folder")
group.add_argument("--fps",
type=int,
required=False,
default=30,
help="render fps")
group.add_argument(
"--mode",
type=str,
required=False,
default="sequence",
help="render target: video, sequence, frame",
)
params = parser.parse_args()
# Load yaml config files
OmegaConf.register_new_resolver("eval", eval)
cfg_assets = OmegaConf.load(params.cfg_assets)
cfg_base = OmegaConf.load(pjoin(cfg_assets.CONFIG_FOLDER, 'default.yaml'))
cfg_exp = OmegaConf.merge(cfg_base, OmegaConf.load(params.cfg))
if not cfg_exp.FULL_CONFIG:
cfg_exp = get_module_config(cfg_exp, cfg_assets.CONFIG_FOLDER)
cfg = OmegaConf.merge(cfg_exp, cfg_assets)
# Update config with arguments
if phase in ["train", "test"]:
cfg.TRAIN.BATCH_SIZE = params.batch_size if params.batch_size else cfg.TRAIN.BATCH_SIZE
cfg.DEVICE = params.device if params.device else cfg.DEVICE
cfg.NUM_NODES = params.num_nodes if params.num_nodes else cfg.NUM_NODES
cfg.model.params.task = params.task if params.task else cfg.model.params.task
cfg.DEBUG = not params.nodebug if params.nodebug is not None else cfg.DEBUG
# Force no debug in test
if phase == "test":
cfg.DEBUG = False
cfg.DEVICE = [0]
print("Force no debugging and one gpu when testing")
if phase == "demo":
cfg.DEMO.RENDER = params.render
cfg.DEMO.FRAME_RATE = params.frame_rate
cfg.DEMO.EXAMPLE = params.example
cfg.DEMO.TASK = params.task
cfg.TEST.FOLDER = params.out_dir if params.out_dir else cfg.TEST.FOLDER
os.makedirs(cfg.TEST.FOLDER, exist_ok=True)
if phase == "render":
if params.npy:
cfg.RENDER.NPY = params.npy
cfg.RENDER.INPUT_MODE = "npy"
if params.dir:
cfg.RENDER.DIR = params.dir
cfg.RENDER.INPUT_MODE = "dir"
if params.fps:
cfg.RENDER.FPS = float(params.fps)
cfg.RENDER.MODE = params.mode
# Debug mode
if cfg.DEBUG:
cfg.NAME = "debug--" + cfg.NAME
cfg.LOGGER.WANDB.params.offline = True
cfg.LOGGER.VAL_EVERY_STEPS = 1
# Resume config
cfg = resume_config(cfg)
return cfg
@@ -0,0 +1,175 @@
NAME: Webui # Experiment name
DEBUG: False # Debug mode
ACCELERATOR: 'cpu' # Devices optioncal: “cpu”, “gpu”, “tpu”, “ipu”, “hpu”, “mps, “auto”
DEVICE: [0] # Index of gpus eg. [0] or [0,1,2,3]
# Training configuration
TRAIN:
#---------------------------------
STAGE: lm_instruct
DATASETS: ['humanml3d'] # Training datasets
SPLIT: 'train' # Training split name
NUM_WORKERS: 32 # Number of workers
BATCH_SIZE: 16 # Size of batches
START_EPOCH: 0 # Start epochMMOTIONENCODER
END_EPOCH: 99999 # End epoch
ABLATION:
pkeep: 0.5
OPTIM:
TYPE: AdamW # Optimizer type
LR: 2e-4 # Learning rate
WEIGHT_DECAY: 0.0
LR_SCHEDULER: [100, 200, 300, 400]
GAMMA: 0.8
LR_SCHEDULER:
target: CosineAnnealingLR
params:
T_max: ${eval:${LOGGER.VAL_EVERY_STEPS} * 100}
eta_min: 1e-6
# Evaluating Configuration
EVAL:
DATASETS: ['humanml3d'] # Evaluating datasets
BATCH_SIZE: 32 # Evaluating Batch size
NUM_WORKERS: 8 # Validation Batch size
SPLIT: test
# Test Configuration
TEST:
CHECKPOINTS: checkpoints/MotionGPT-base/motiongpt_s3_h3d.ckpt
DATASETS: ['humanml3d'] # training datasets
SPLIT: test
BATCH_SIZE: 32 # training Batch size
MEAN: False
NUM_SAMPLES: 1
FACT: 1
# Datasets Configuration
DATASET:
target: .mGPT.data.HumanML3D.HumanML3DDataModule
JOINT_TYPE: 'humanml3d' # join type
CODE_PATH: 'VQBEST'
TASK_ROOT: .deps/mGPT_instructions
TASK_PATH: ''
SMPL_PATH: .deps/smpl
TRANSFORM_PATH: .deps/transforms/
WORD_VERTILIZER_PATH: .deps/glove/
NFEATS: 263
KIT:
ROOT: .datasets/kit-ml # KIT directory
SPLIT_ROOT: .datasets/kit-ml # KIT splits directory
MEAN_STD_PATH: .deps/t2m/
MAX_MOTION_LEN: 196
MIN_MOTION_LEN: 24
MAX_TEXT_LEN: 20
PICK_ONE_TEXT: true
FRAME_RATE: 12.5
UNIT_LEN: 4
HUMANML3D:
ROOT: .datasets/humanml3d # HumanML3D directory
SPLIT_ROOT: .datasets/humanml3d # HumanML3D splits directory
ASSETS_ROOT: .assets/meta/
MEAN_STD_PATH: .deps/t2m/
MAX_MOTION_LEN: 196
MIN_MOTION_LEN: 40
MAX_TEXT_LEN: 20
PICK_ONE_TEXT: true
FRAME_RATE: 20.0
UNIT_LEN: 4
STD_TEXT: False
ABLATION:
# For MotionGPT
use_length: False
predict_ratio: 0.2
inbetween_ratio: 0.25
image_size: 256
# For Motion-latent-diffusion
VAE_TYPE: 'actor' # vae ablation: actor or mcross
VAE_ARCH: 'encoder_decoder' # mdiffusion vae architecture
PE_TYPE: 'actor' # mdiffusion mld or actor
DIFF_PE_TYPE: 'actor' # mdiffusion mld or actor
SKIP_CONNECT: False # skip connection for denoiser va
MLP_DIST: False # use linear to expand mean and std rather expand token nums
IS_DIST: False # Mcross distribution kl
PREDICT_EPSILON: True # noise or motion
METRIC:
TYPE: ['TM2TMetrics']
TM2T:
t2m_path: .deps/t2m/ # path for tm2t evaluator
TASK: 't2m'
FORCE_IN_METER: True
DIST_SYNC_ON_STEP: True
MM_NUM_SAMPLES: 100 # Number of samples for multimodal test
MM_NUM_REPEATS: 30 # Number of repeats for multimodal test
MM_NUM_TIMES: 10 # Number of times to repeat the multimodal test
DIVERSITY_TIMES: 300 # Number of times to repeat the diversity test
# Losses Configuration
LOSS:
TYPE: t2mgpt # Losses type
LAMBDA_FEATURE: 1.0
LAMBDA_VELOCITY: 0.5
LAMBDA_COMMIT: 0.02
LAMBDA_CLS: 1.0
LAMBDA_M2T2M: 1.0
LAMBDA_T2M2T: 10.0
ABLATION:
RECONS_LOSS: 'l1_smooth'
LAMBDA_REC: 1.0 # Lambda for reconstruction losses
LAMBDA_JOINT: 1.0 # Lambda for joint losses
LAMBDA_LATENT: 1e-5 # Lambda for latent losses
LAMBDA_KL: 1e-5 # Lambda for kl losses
LAMBDA_GEN: 1.0 # Lambda for text-motion generation losses
LAMBDA_CROSS: 1.0 # Lambda for cross-reconstruction losses
LAMBDA_CYCLE: 1.0 # Lambda for cycle losses
LAMBDA_PRIOR: 0.0 # Lambda for diffusion prior losses
# Model Configuration
model:
target: .mGPT.models.mgpt.MotionGPT
params:
condition: 'text'
task: 't2m'
lm:
target: .mGPT.archs.mgpt_lm.MLM
params:
model_type: t5
model_path: google/flan-t5-base
stage: ${TRAIN.STAGE}
motion_codebook_size: 512
ablation: ${ABLATION}
motion_vae:
target: .mGPT.archs.mgpt_vq.VQVae
params:
quantizer: 'ema_reset'
code_num: 512
code_dim: 512
output_emb_width: 512
down_t: 2
stride_t: 2
width: 512
depth: 3
dilation_growth_rate: 3
norm: None
activation: 'relu'
nfeats: ${DATASET.NFEATS}
ablation: ${ABLATION}
# Related parameters
stage: ${TRAIN.STAGE}
debug: ${DEBUG}
codebook_size: 512
metrics_dict: ${METRIC.TYPE}
# Logger configuration
LOGGER:
LOG_EVERY_STEPS: 5
VAL_EVERY_STEPS: 10
TENSORBOARD: True
wandb:
params:
project: null
+118
View File
@@ -0,0 +1,118 @@
import numpy as np
import torch
from os.path import join as pjoin
import os
from .humanml.utils.word_vectorizer import WordVectorizer
from .humanml.scripts.motion_process import (process_file, recover_from_ric)
from . import BASEDataModule
from .humanml import Text2MotionDatasetEval, Text2MotionDataset, Text2MotionDatasetCB, MotionDataset, MotionDatasetVQ, Text2MotionDatasetToken, Text2MotionDatasetM2T
from .utils import humanml3d_collate
script_directory = os.path.dirname(os.path.abspath(__file__))
class HumanML3DDataModule(BASEDataModule):
def __init__(self, cfg, **kwargs):
super().__init__(collate_fn=humanml3d_collate)
self.cfg = cfg
self.save_hyperparameters(logger=False)
# Basic info of the dataset
cfg.DATASET.JOINT_TYPE = 'humanml3d'
self.name = "humanml3d"
self.njoints = 22
# Path to the dataset
data_root = cfg.DATASET.HUMANML3D.ROOT
assets_root = cfg.DATASET.HUMANML3D.ASSETS_ROOT
self.hparams.data_root = data_root
self.hparams.text_dir = pjoin(data_root, "texts")
self.hparams.motion_dir = pjoin(data_root, 'new_joint_vecs')
# Mean and std of the dataset
self.hparams.mean = np.load(pjoin(script_directory, "mean.npy"))
self.hparams.std = np.load(pjoin(script_directory, "std.npy"))
# Mean and std for fair evaluation
self.hparams.mean_eval = np.load(pjoin(script_directory, "mean_eval.npy"))
self.hparams.std_eval = np.load(pjoin(script_directory, "std_eval.npy"))
# Length of the dataset
self.hparams.max_motion_length = cfg.DATASET.HUMANML3D.MAX_MOTION_LEN
self.hparams.min_motion_length = cfg.DATASET.HUMANML3D.MIN_MOTION_LEN
self.hparams.max_text_len = cfg.DATASET.HUMANML3D.MAX_TEXT_LEN
self.hparams.unit_length = cfg.DATASET.HUMANML3D.UNIT_LEN
# Additional parameters
self.hparams.debug = cfg.DEBUG
self.hparams.stage = cfg.TRAIN.STAGE
# Dataset switch
self.DatasetEval = Text2MotionDatasetEval
if cfg.TRAIN.STAGE == "vae":
if cfg.model.params.motion_vae.target.split('.')[-1].lower() == "vqvae":
self.hparams.win_size = 64
self.Dataset = MotionDatasetVQ
else:
self.Dataset = MotionDataset
elif 'lm' in cfg.TRAIN.STAGE:
self.hparams.code_path = cfg.DATASET.CODE_PATH
self.hparams.task_path = cfg.DATASET.TASK_PATH
self.hparams.std_text = cfg.DATASET.HUMANML3D.STD_TEXT
self.Dataset = Text2MotionDatasetCB
elif cfg.TRAIN.STAGE == "token":
self.Dataset = Text2MotionDatasetToken
self.DatasetEval = Text2MotionDatasetToken
elif cfg.TRAIN.STAGE == "m2t":
self.Dataset = Text2MotionDatasetM2T
self.DatasetEval = Text2MotionDatasetM2T
else:
self.Dataset = Text2MotionDataset
# Get additional info of the dataset
self.nfeats = 263
cfg.DATASET.NFEATS = self.nfeats
def feats2joints(self, features):
mean = torch.tensor(self.hparams.mean).to(features)
std = torch.tensor(self.hparams.std).to(features)
features = features * std + mean
return recover_from_ric(features, self.njoints)
def joints2feats(self, features):
features = process_file(features, self.njoints)[0]
return features
def normalize(self, features):
mean = torch.tensor(self.hparams.mean).to(features)
std = torch.tensor(self.hparams.std).to(features)
features = (features - mean) / std
return features
def denormalize(self, features):
mean = torch.tensor(self.hparams.mean).to(features)
std = torch.tensor(self.hparams.std).to(features)
features = features * std + mean
return features
def renorm4t2m(self, features):
# renorm to t2m norms for using t2m evaluators
ori_mean = torch.tensor(self.hparams.mean).to(features)
ori_std = torch.tensor(self.hparams.std).to(features)
eval_mean = torch.tensor(self.hparams.mean_eval).to(features)
eval_std = torch.tensor(self.hparams.std_eval).to(features)
features = features * ori_std + ori_mean
features = (features - eval_mean) / eval_std
return features
def mm_mode(self, mm_on=True):
if mm_on:
self.is_mm = True
self.name_list = self.test_dataset.name_list
self.mm_list = np.random.choice(self.name_list,
self.cfg.METRIC.MM_NUM_SAMPLES,
replace=False)
self.test_dataset.name_list = self.mm_list
else:
self.is_mm = False
self.test_dataset.name_list = self.name_list
+88
View File
@@ -0,0 +1,88 @@
import numpy as np
import torch
from os.path import join as pjoin
from .humanml.utils.word_vectorizer import WordVectorizer
from .humanml.scripts.motion_process import (process_file, recover_from_ric)
from .HumanML3D import HumanML3DDataModule
from .humanml import Text2MotionDatasetEval, Text2MotionDataset, Text2MotionDatasetCB, MotionDataset, MotionDatasetVQ, Text2MotionDatasetToken
class KitDataModule(HumanML3DDataModule):
def __init__(self, cfg, **kwargs):
super().__init__(cfg, **kwargs)
# Basic info of the dataset
self.name = "kit"
self.njoints = 21
# Path to the dataset
data_root = cfg.DATASET.KIT.ROOT
self.hparams.data_root = data_root
self.hparams.text_dir = pjoin(data_root, "texts")
self.hparams.motion_dir = pjoin(data_root, 'new_joint_vecs')
# Mean and std of the dataset
dis_data_root = pjoin(cfg.DATASET.KIT.MEAN_STD_PATH, 'kit',
"VQVAEV3_CB1024_CMT_H1024_NRES3", "meta")
self.hparams.mean = np.load(pjoin(dis_data_root, "mean.npy"))
self.hparams.std = np.load(pjoin(dis_data_root, "std.npy"))
# Mean and std for fair evaluation
dis_data_root_eval = pjoin(cfg.DATASET.KIT.MEAN_STD_PATH, 't2m',
"Comp_v6_KLD005", "meta")
self.hparams.mean_eval = np.load(pjoin(dis_data_root_eval, "mean.npy"))
self.hparams.std_eval = np.load(pjoin(dis_data_root_eval, "std.npy"))
# Length of the dataset
self.hparams.max_motion_length = cfg.DATASET.KIT.MAX_MOTION_LEN
self.hparams.min_motion_length = cfg.DATASET.KIT.MIN_MOTION_LEN
self.hparams.max_text_len = cfg.DATASET.KIT.MAX_TEXT_LEN
self.hparams.unit_length = cfg.DATASET.KIT.UNIT_LEN
# Get additional info of the dataset
self._sample_set = self.get_sample_set(overrides={"split": "test", "tiny": True})
self.nfeats = self._sample_set.nfeats
cfg.DATASET.NFEATS = self.nfeats
def feats2joints(self, features):
mean = torch.tensor(self.hparams.mean).to(features)
std = torch.tensor(self.hparams.std).to(features)
features = features * std + mean
return recover_from_ric(features, self.njoints)
def joints2feats(self, features):
features = process_file(features, self.njoints)[0]
# mean = torch.tensor(self.hparams.mean).to(features)
# std = torch.tensor(self.hparams.std).to(features)
# features = (features - mean) / std
return features
def normalize(self, features):
mean = torch.tensor(self.hparams.mean).to(features)
std = torch.tensor(self.hparams.std).to(features)
features = (features - mean) / std
return features
def renorm4t2m(self, features):
# renorm to t2m norms for using t2m evaluators
ori_mean = torch.tensor(self.hparams.mean).to(features)
ori_std = torch.tensor(self.hparams.std).to(features)
eval_mean = torch.tensor(self.hparams.mean_eval).to(features)
eval_std = torch.tensor(self.hparams.std_eval).to(features)
features = features * ori_std + ori_mean
features = (features - eval_mean) / eval_std
return features
def mm_mode(self, mm_on=True):
# random select samples for mm
if mm_on:
self.is_mm = True
self.name_list = self.test_dataset.name_list
self.mm_list = np.random.choice(self.name_list,
self.cfg.METRIC.MM_NUM_SAMPLES,
replace=False)
self.test_dataset.name_list = self.mm_list
else:
self.is_mm = False
self.test_dataset.name_list = self.name_list
+103
View File
@@ -0,0 +1,103 @@
import pytorch_lightning as pl
from torch.utils.data import DataLoader
class BASEDataModule(pl.LightningDataModule):
def __init__(self, collate_fn):
super().__init__()
self.dataloader_options = {"collate_fn": collate_fn}
self.persistent_workers = True
self.is_mm = False
self._train_dataset = None
self._val_dataset = None
self._test_dataset = None
def get_sample_set(self, overrides={}):
sample_params = self.hparams.copy()
sample_params.update(overrides)
return self.DatasetEval(**sample_params)
@property
def train_dataset(self):
if self._train_dataset is None:
self._train_dataset = self.Dataset(split=self.cfg.TRAIN.SPLIT,
**self.hparams)
return self._train_dataset
@property
def val_dataset(self):
if self._val_dataset is None:
params = self.hparams.copy()
params['code_path'] = None
params['split'] = self.cfg.EVAL.SPLIT
self._val_dataset = self.DatasetEval(**params)
return self._val_dataset
@property
def test_dataset(self):
if self._test_dataset is None:
# self._test_dataset = self.DatasetEval(split=self.cfg.TEST.SPLIT,
# **self.hparams)
params = self.hparams.copy()
params['code_path'] = None
params['split'] = self.cfg.TEST.SPLIT
self._test_dataset = self.DatasetEval( **params)
return self._test_dataset
def setup(self, stage=None):
# Use the getter the first time to load the data
if stage in (None, "fit"):
_ = self.train_dataset
_ = self.val_dataset
if stage in (None, "test"):
_ = self.test_dataset
def train_dataloader(self):
dataloader_options = self.dataloader_options.copy()
dataloader_options["batch_size"] = self.cfg.TRAIN.BATCH_SIZE
dataloader_options["num_workers"] = self.cfg.TRAIN.NUM_WORKERS
return DataLoader(
self.train_dataset,
shuffle=False,
persistent_workers=True,
**dataloader_options,
)
def predict_dataloader(self):
dataloader_options = self.dataloader_options.copy()
dataloader_options[
"batch_size"] = 1 if self.is_mm else self.cfg.TEST.BATCH_SIZE
dataloader_options["num_workers"] = self.cfg.TEST.NUM_WORKERS
dataloader_options["shuffle"] = False
return DataLoader(
self.test_dataset,
persistent_workers=True,
**dataloader_options,
)
def val_dataloader(self):
# overrides batch_size and num_workers
dataloader_options = self.dataloader_options.copy()
dataloader_options["batch_size"] = self.cfg.EVAL.BATCH_SIZE
dataloader_options["num_workers"] = self.cfg.EVAL.NUM_WORKERS
dataloader_options["shuffle"] = False
return DataLoader(
self.val_dataset,
persistent_workers=True,
**dataloader_options,
)
def test_dataloader(self):
# overrides batch_size and num_workers
dataloader_options = self.dataloader_options.copy()
dataloader_options[
"batch_size"] = 1 if self.is_mm else self.cfg.TEST.BATCH_SIZE
dataloader_options["num_workers"] = self.cfg.TEST.NUM_WORKERS
dataloader_options["shuffle"] = False
return DataLoader(
self.test_dataset,
persistent_workers=True,
**dataloader_options,
)
@@ -0,0 +1,15 @@
from omegaconf import OmegaConf
from os.path import join as pjoin
from ..config import instantiate_from_config
def build_data(cfg, phase="train"):
data_config = OmegaConf.to_container(cfg.DATASET, resolve=True)
data_config['params'] = {'cfg': cfg, 'phase': phase}
if isinstance(data_config['target'], str):
return instantiate_from_config(data_config)
elif isinstance(data_config['target'], list):
data_config_tmp = data_config.copy()
data_config_tmp['params']['dataModules'] = data_config['target']
data_config_tmp['target'] = 'mGPT.data.Concat.ConcatDataModule'
return instantiate_from_config(data_config)
@@ -0,0 +1 @@
This code is based on https://github.com/EricGuo5513/text-to-motion.git
@@ -0,0 +1,7 @@
from .dataset_t2m import Text2MotionDataset
from .dataset_t2m_eval import Text2MotionDatasetEval
from .dataset_t2m_cb import Text2MotionDatasetCB
from .dataset_t2m_token import Text2MotionDatasetToken
from .dataset_t2m_m2t import Text2MotionDatasetM2T
from .dataset_m import MotionDataset
from .dataset_m_vq import MotionDatasetVQ
@@ -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.float64).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)
@@ -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
@@ -0,0 +1,156 @@
import os
import rich
import random
import pickle
import codecs as cs
import numpy as np
from torch.utils import data
from rich.progress import track
from os.path import join as pjoin
class MotionDataset(data.Dataset):
def __init__(
self,
data_root,
split,
mean,
std,
max_motion_length=196,
min_motion_length=20,
unit_length=4,
fps=20,
tmpFile=True,
tiny=False,
debug=False,
**kwargs,
):
# restrian the length of motion and text
self.max_motion_length = max_motion_length
self.min_motion_length = min_motion_length
self.unit_length = unit_length
# Data mean and std
self.mean = mean
self.std = std
# Data path
split_file = pjoin(data_root, split + '.txt')
motion_dir = pjoin(data_root, 'new_joint_vecs')
text_dir = pjoin(data_root, 'texts')
# Data id list
self.id_list = []
with cs.open(split_file, "r") as f:
for line in f.readlines():
self.id_list.append(line.strip())
# Debug mode
if tiny or debug:
enumerator = enumerate(
track(
self.id_list,
f"Loading HumanML3D {split}",
))
maxdata = 100
subset = '_tiny'
else:
enumerator = enumerate(self.id_list)
maxdata = 1e10
subset = ''
new_name_list = []
motion_dict = {}
# Fast loading
if os.path.exists(pjoin(data_root, f'tmp/{split}{subset}_motion.pkl')):
with rich.progress.open(pjoin(data_root, f'tmp/{split}{subset}_motion.pkl'),
'rb', description=f"Loading HumanML3D {split}") as file:
motion_dict = pickle.load(file)
with open(pjoin(data_root, f'tmp/{split}{subset}_index.pkl'), 'rb') as file:
new_name_list = pickle.load(file)
else:
for idx, name in enumerator:
if len(new_name_list) > maxdata:
break
try:
motion = [np.load(pjoin(motion_dir, name + ".npy"))]
# Read text
with cs.open(pjoin(text_dir, name + '.txt')) as f:
text_data = []
flag = False
lines = f.readlines()
for line in lines:
try:
line_split = line.strip().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
if f_tag == 0.0 and to_tag == 0.0:
flag = True
else:
motion_new = [tokens[int(f_tag*fps/unit_length) : int(to_tag*fps/unit_length)] for tokens in motion if int(f_tag*fps/unit_length) < int(to_tag*fps/unit_length)]
if len(motion_new) == 0:
continue
new_name = '%s_%f_%f'%(name, f_tag, to_tag)
motion_dict[new_name] = {
'motion': motion_new,
"length": [len(m[0]) for m in motion_new]}
new_name_list.append(new_name)
except:
pass
if flag:
motion_dict[name] = {
'motion': motion,
"length": [len(m[0]) for m in motion]}
new_name_list.append(name)
except:
pass
if tmpFile:
os.makedirs(pjoin(data_root, 'tmp'), exist_ok=True)
with open(pjoin(data_root, f'tmp/{split}{subset}_motion.pkl'),'wb') as file:
pickle.dump(motion_dict, file)
with open(pjoin(data_root, f'tmp/{split}{subset}_index.pkl'), 'wb') as file:
pickle.dump(new_name_list, file)
self.motion_dict = motion_dict
self.name_list = new_name_list
self.nfeats = motion_dict[new_name_list[0]]['motion'][0].shape[1]
def __len__(self):
return len(self.name_list)
def __getitem__(self, item):
data = self.motion_dict[self.name_list[item]]
motion_list, m_length = data["motion"], data["length"]
# Randomly select a motion
motion = random.choice(motion_list)
# Crop the motions in to times of 4, and introduce small variations
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
return None, motion, m_length, None, None, None, None,
@@ -0,0 +1,54 @@
import random
import codecs as cs
import numpy as np
from torch.utils import data
from rich.progress import track
from os.path import join as pjoin
from .dataset_m import MotionDataset
from .dataset_t2m import Text2MotionDataset
class MotionDatasetVQ(Text2MotionDataset):
def __init__(
self,
data_root,
split,
mean,
std,
max_motion_length,
min_motion_length,
win_size,
unit_length=4,
fps=20,
tmpFile=True,
tiny=False,
debug=False,
**kwargs,
):
super().__init__(data_root, split, mean, std, max_motion_length,
min_motion_length, unit_length, fps, tmpFile, tiny,
debug, **kwargs)
# Filter out the motions that are too short
self.window_size = win_size
name_list = list(self.name_list)
for name in self.name_list:
motion = self.data_dict[name]["motion"]
if motion.shape[0] < self.window_size:
name_list.remove(name)
self.data_dict.pop(name)
self.name_list = name_list
def __len__(self):
return len(self.name_list)
def __getitem__(self, item):
idx = self.pointer + item
data = self.data_dict[self.name_list[idx]]
motion, length = data["motion"], data["length"]
idx = random.randint(0, motion.shape[0] - self.window_size)
motion = motion[idx:idx + self.window_size]
motion = (motion - self.mean) / self.std
return None, motion, length, None, None, None, None,
@@ -0,0 +1,211 @@
import os
import rich
import random
import pickle
import codecs as cs
import numpy as np
from torch.utils import data
from rich.progress import track
from os.path import join as pjoin
class Text2MotionDataset(data.Dataset):
def __init__(
self,
data_root,
split,
mean,
std,
max_motion_length=196,
min_motion_length=40,
unit_length=4,
fps=20,
tmpFile=True,
tiny=False,
debug=False,
**kwargs,
):
# restrian the length of motion and text
self.max_length = 20
self.max_motion_length = max_motion_length
self.min_motion_length = min_motion_length
self.unit_length = unit_length
# Data mean and std
self.mean = mean
self.std = std
# Data path
split_file = pjoin(data_root, split + '.txt')
motion_dir = pjoin(data_root, 'new_joint_vecs')
text_dir = pjoin(data_root, 'texts')
# Data id list
self.id_list = []
with cs.open(split_file, "r") as f:
for line in f.readlines():
self.id_list.append(line.strip())
# Debug mode
if tiny or debug:
enumerator = enumerate(self.id_list)
maxdata = 100
subset = '_tiny'
else:
enumerator = enumerate(
track(
self.id_list,
f"Loading HumanML3D {split}",
))
maxdata = 1e10
subset = ''
new_name_list = []
length_list = []
data_dict = {}
# Fast loading
if os.path.exists(pjoin(data_root, f'tmp/{split}{subset}_data.pkl')):
if tiny or debug:
with open(pjoin(data_root, f'tmp/{split}{subset}_data.pkl'),
'rb') as file:
data_dict = pickle.load(file)
else:
with rich.progress.open(
pjoin(data_root, f'tmp/{split}{subset}_data.pkl'),
'rb',
description=f"Loading HumanML3D {split}") as file:
data_dict = pickle.load(file)
with open(pjoin(data_root, f'tmp/{split}{subset}_index.pkl'),
'rb') as file:
name_list = pickle.load(file)
for name in new_name_list:
length_list.append(data_dict[name]['length'])
else:
for idx, name in enumerator:
if len(new_name_list) > maxdata:
break
try:
motion = np.load(pjoin(motion_dir, name + ".npy"))
if (len(motion)) < self.min_motion_length or (len(motion)
>= 200):
continue
# Read text
text_data = []
flag = False
with cs.open(pjoin(text_dir, name + '.txt')) as f:
lines = f.readlines()
for line in lines:
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:
motion_new = motion[int(f_tag *
fps):int(to_tag * fps)]
if (len(motion_new)
) < self.min_motion_length or (
len(motion_new) >= 200):
continue
new_name = random.choice(
'ABCDEFGHIJKLMNOPQRSTUVW') + '_' + name
while new_name in new_name_list:
new_name = random.choice(
'ABCDEFGHIJKLMNOPQRSTUVW') + '_' + name
name_count = 1
while new_name in data_dict:
new_name += '_' + name_count
name_count += 1
data_dict[new_name] = {
'motion': motion_new,
"length": len(motion_new),
'text': [text_dict]
}
new_name_list.append(new_name)
length_list.append(len(motion_new))
if flag:
data_dict[name] = {
'motion': motion,
"length": len(motion),
'text': text_data
}
new_name_list.append(name)
length_list.append(len(motion))
except:
pass
name_list, length_list = zip(
*sorted(zip(new_name_list, length_list), key=lambda x: x[1]))
if tmpFile:
os.makedirs(pjoin(data_root, 'tmp'), exist_ok=True)
with open(pjoin(data_root, f'tmp/{split}{subset}_data.pkl'),
'wb') as file:
pickle.dump(data_dict, file)
with open(pjoin(data_root, f'tmp/{split}{subset}_index.pkl'),
'wb') as file:
pickle.dump(name_list, file)
self.length_arr = np.array(length_list)
self.data_dict = data_dict
self.name_list = name_list
self.nfeats = data_dict[name_list[0]]['motion'].shape[1]
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 __len__(self):
return len(self.name_list) - self.pointer
def __getitem__(self, item):
idx = self.pointer + item
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 = text_data["caption"]
all_captions = [
' '.join([token.split('/')[0] for token in text_dic['tokens']])
for text_dic in text_list
]
# Crop the motions in to times of 4, and introduce small variations
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
return caption, motion, m_length, None, None, None, None, all_captions
@@ -0,0 +1,211 @@
import rich
import random
import pickle
import os
import numpy as np
import codecs as cs
from torch.utils import data
from os.path import join as pjoin
from rich.progress import track
import json
import spacy
class Text2MotionDatasetCB(data.Dataset):
def __init__(
self,
data_root,
split,
mean,
std,
max_motion_length=196,
min_motion_length=20,
unit_length=4,
fps=20,
tmpFile=True,
tiny=False,
debug=False,
stage='lm_pretrain',
code_path='VQVAE',
task_path=None,
std_text=False,
**kwargs,
):
self.tiny = tiny
self.unit_length = unit_length
# Data mean and std
self.mean = mean
self.std = std
# Data path
split = 'train'
split_file = pjoin(data_root, split + '.txt')
motion_dir = pjoin(data_root, code_path)
text_dir = pjoin(data_root, 'texts')
if task_path:
instructions = task_path
elif stage == 'lm_pretrain':
instructions = pjoin(data_root, 'template_pretrain.json')
elif stage in ['lm_instruct', "lm_rl"]:
instructions = pjoin(data_root, 'template_instructions.json')
else:
raise NotImplementedError(f"stage {stage} not implemented")
# Data id list
self.id_list = []
with cs.open(split_file, "r") as f:
for line in f.readlines():
self.id_list.append(line.strip())
# Debug mode
if tiny or debug:
enumerator = enumerate(self.id_list)
maxdata = 100
subset = '_tiny'
else:
enumerator = enumerate(
track(
self.id_list,
f"Loading HumanML3D {split}",
))
maxdata = 1e10
subset = ''
new_name_list = []
data_dict = {}
# Fast loading
for i, name in enumerator:
if len(new_name_list) > maxdata:
break
try:
# Load motion tokens
m_token_list = np.load(pjoin(motion_dir, f'{name}.npy'))
# Read text
with cs.open(pjoin(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
if tmpFile:
os.makedirs(pjoin(data_root, 'tmp'), exist_ok=True)
with open(
pjoin(data_root,
f'tmp/{split}{subset}_tokens_data.pkl'),
'wb') as file:
pickle.dump(data_dict, file)
with open(
pjoin(data_root,
f'tmp/{split}{subset}_tokens_index.pkl'),
'wb') as file:
pickle.dump(new_name_list, file)
self.data_dict = data_dict
self.name_list = new_name_list
self.nlp = spacy.load('en_core_web_sm')
self.std_text = std_text
self.instructions = json.load(open(instructions, 'r'))
self.tasks = []
for task in self.instructions.keys():
for subtask in self.instructions[task].keys():
self.tasks.append(self.instructions[task][subtask])
def __len__(self):
return len(self.name_list) * len(self.tasks)
def __getitem__(self, item):
data_idx = item % len(self.name_list)
task_idx = item // len(self.name_list)
data = self.data_dict[self.name_list[data_idx]]
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']
if self.std_text:
doc = self.nlp(caption)
word_list = []
pos_list = []
for token in doc:
word = token.text
if not word.isalpha():
continue
if (token.pos_ == 'NOUN'
or token.pos_ == 'VERB') and (word != 'left'):
word_list.append(token.lemma_)
else:
word_list.append(word)
pos_list.append(token.pos_)
caption = ' '.join(word_list)
all_captions = [
' '.join([token.split('/')[0] for token in text_dic['tokens']])
for text_dic in text_list
]
coin = np.random.choice([False, False, True])
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]
tasks = self.tasks[task_idx]
return caption, m_tokens, m_tokens_len, None, None, None, None, all_captions, tasks
@@ -0,0 +1,92 @@
import random
import numpy as np
from .dataset_t2m import Text2MotionDataset
class Text2MotionDatasetEval(Text2MotionDataset):
def __init__(
self,
data_root,
split,
mean,
std,
w_vectorizer,
max_motion_length=196,
min_motion_length=40,
unit_length=4,
fps=20,
tmpFile=True,
tiny=False,
debug=False,
**kwargs,
):
super().__init__(data_root, split, mean, std, max_motion_length,
min_motion_length, unit_length, fps, tmpFile, tiny,
debug, **kwargs)
self.w_vectorizer = w_vectorizer
def __getitem__(self, item):
# Get text data
idx = self.pointer + item
data = self.data_dict[self.name_list[idx]]
motion, m_length, text_list = data["motion"], data["length"], data["text"]
all_captions = [
' '.join([token.split('/')[0] for token in text_dic['tokens']])
for text_dic in text_list
]
if len(all_captions) > 3:
all_captions = all_captions[:3]
elif len(all_captions) == 2:
all_captions = all_captions + all_captions[0:1]
elif len(all_captions) == 1:
all_captions = all_captions * 3
# Randomly select a caption
text_data = random.choice(text_list)
caption, tokens = text_data["caption"], text_data["tokens"]
# Text
max_text_len = 20
if len(tokens) < max_text_len:
# pad with "unk"
tokens = ["sos/OTHER"] + tokens + ["eos/OTHER"]
sent_len = len(tokens)
tokens = tokens + ["unk/OTHER"] * (max_text_len + 2 - sent_len)
else:
# crop
tokens = tokens[: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)
# Random crop
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
return caption, motion, m_length, word_embeddings, pos_one_hots, sent_len, "_".join(
tokens), all_captions
@@ -0,0 +1,119 @@
import random
import numpy as np
from torch.utils import data
from .dataset_t2m import Text2MotionDataset
import codecs as cs
from os.path import join as pjoin
class Text2MotionDatasetM2T(data.Dataset):
def __init__(
self,
data_root,
split,
mean,
std,
max_motion_length=196,
min_motion_length=40,
unit_length=4,
fps=20,
tmpFile=True,
tiny=False,
debug=False,
**kwargs,
):
self.max_motion_length = max_motion_length
self.min_motion_length = min_motion_length
self.unit_length = unit_length
# Data mean and std
self.mean = mean
self.std = std
# Data path
split_file = pjoin(data_root, split + '.txt')
motion_dir = pjoin(data_root, 'new_joint_vecs')
text_dir = pjoin(data_root, 'texts')
# Data id list
self.id_list = []
with cs.open(split_file, "r") as f:
for line in f.readlines():
self.id_list.append(line.strip())
new_name_list = []
length_list = []
data_dict = {}
for name in self.id_list:
# try:
motion = np.load(pjoin(motion_dir, name + '.npy'))
if (len(motion)) < self.min_motion_length or (len(motion) >= 200):
continue
text_data = []
flag = False
with cs.open(pjoin(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*20) : int(to_tag*20)]
if (len(n_motion)) < min_motion_length or (len(n_motion) >= 200):
continue
new_name = "%s_%f_%f"%(name, f_tag, to_tag)
data_dict[new_name] = {'motion': n_motion,
'length': len(n_motion),
'text':[text_dict]}
new_name_list.append(new_name)
except:
print(line_split)
print(line_split[2], line_split[3], f_tag, to_tag, name)
if flag:
data_dict[name] = {'motion': motion,
'length': len(motion),
'name': name,
'text': text_data}
new_name_list.append(name)
length_list.append(len(motion))
# except:
# # Some motion may not exist in KIT dataset
# pass
self.length_arr = np.array(length_list)
self.data_dict = data_dict
self.name_list = new_name_list
self.nfeats = motion.shape[-1]
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']
"Z Normalization"
motion = (motion - self.mean) / self.std
return name, motion, m_length, True, True, True, True, True, True
@@ -0,0 +1,86 @@
import random
import numpy as np
from torch.utils import data
from .dataset_t2m import Text2MotionDataset
import codecs as cs
from os.path import join as pjoin
class Text2MotionDatasetToken(data.Dataset):
def __init__(
self,
data_root,
split,
mean,
std,
max_motion_length=196,
min_motion_length=40,
unit_length=4,
fps=20,
tmpFile=True,
tiny=False,
debug=False,
**kwargs,
):
self.max_motion_length = max_motion_length
self.min_motion_length = min_motion_length
self.unit_length = unit_length
# Data mean and std
self.mean = mean
self.std = std
# Data path
split_file = pjoin(data_root, split + '.txt')
motion_dir = pjoin(data_root, 'new_joint_vecs')
text_dir = pjoin(data_root, 'texts')
# Data id list
self.id_list = []
with cs.open(split_file, "r") as f:
for line in f.readlines():
self.id_list.append(line.strip())
new_name_list = []
length_list = []
data_dict = {}
for name in self.id_list:
try:
motion = np.load(pjoin(motion_dir, name + '.npy'))
if (len(motion)) < self.min_motion_length 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.length_arr = np.array(length_list)
self.data_dict = data_dict
self.name_list = new_name_list
self.nfeats = motion.shape[-1]
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 name, motion, m_length, True, True, True, True, True, True
@@ -0,0 +1,529 @@
from os.path import join as pjoin
from ..common.skeleton import Skeleton
import numpy as np
import os
from ..common.quaternion import *
from ..utils.paramUtil import *
import torch
from tqdm import tqdm
# positions (batch, joint_num, 3)
def uniform_skeleton(positions, target_offset):
src_skel = Skeleton(n_raw_offsets, kinematic_chain, 'cpu')
src_offset = src_skel.get_offsets_joints(torch.from_numpy(positions[0]))
src_offset = src_offset.numpy()
tgt_offset = target_offset.numpy()
# print(src_offset)
# print(tgt_offset)
'''Calculate Scale Ratio as the ratio of legs'''
src_leg_len = np.abs(src_offset[l_idx1]).max() + np.abs(src_offset[l_idx2]).max()
tgt_leg_len = np.abs(tgt_offset[l_idx1]).max() + np.abs(tgt_offset[l_idx2]).max()
scale_rt = tgt_leg_len / src_leg_len
# print(scale_rt)
src_root_pos = positions[:, 0]
tgt_root_pos = src_root_pos * scale_rt
'''Inverse Kinematics'''
quat_params = src_skel.inverse_kinematics_np(positions, face_joint_indx)
# print(quat_params.shape)
'''Forward Kinematics'''
src_skel.set_offset(target_offset)
new_joints = src_skel.forward_kinematics_np(quat_params, tgt_root_pos)
return new_joints
def extract_features(positions, feet_thre, n_raw_offsets, kinematic_chain, face_joint_indx, fid_r, fid_l):
global_positions = positions.copy()
""" Get Foot Contacts """
def foot_detect(positions, thres):
velfactor, heightfactor = np.array([thres, thres]), np.array([3.0, 2.0])
feet_l_x = (positions[1:, fid_l, 0] - positions[:-1, fid_l, 0]) ** 2
feet_l_y = (positions[1:, fid_l, 1] - positions[:-1, fid_l, 1]) ** 2
feet_l_z = (positions[1:, fid_l, 2] - positions[:-1, fid_l, 2]) ** 2
# feet_l_h = positions[:-1,fid_l,1]
# feet_l = (((feet_l_x + feet_l_y + feet_l_z) < velfactor) & (feet_l_h < heightfactor)).astype(np.float64)
feet_l = ((feet_l_x + feet_l_y + feet_l_z) < velfactor).astype(np.float64)
feet_r_x = (positions[1:, fid_r, 0] - positions[:-1, fid_r, 0]) ** 2
feet_r_y = (positions[1:, fid_r, 1] - positions[:-1, fid_r, 1]) ** 2
feet_r_z = (positions[1:, fid_r, 2] - positions[:-1, fid_r, 2]) ** 2
# feet_r_h = positions[:-1,fid_r,1]
# feet_r = (((feet_r_x + feet_r_y + feet_r_z) < velfactor) & (feet_r_h < heightfactor)).astype(np.float64)
feet_r = (((feet_r_x + feet_r_y + feet_r_z) < velfactor)).astype(np.float64)
return feet_l, feet_r
#
feet_l, feet_r = foot_detect(positions, feet_thre)
# feet_l, feet_r = foot_detect(positions, 0.002)
'''Quaternion and Cartesian representation'''
r_rot = None
def get_rifke(positions):
'''Local pose'''
positions[..., 0] -= positions[:, 0:1, 0]
positions[..., 2] -= positions[:, 0:1, 2]
'''All pose face Z+'''
positions = qrot_np(np.repeat(r_rot[:, None], positions.shape[1], axis=1), positions)
return positions
def get_quaternion(positions):
skel = Skeleton(n_raw_offsets, kinematic_chain, "cpu")
# (seq_len, joints_num, 4)
quat_params = skel.inverse_kinematics_np(positions, face_joint_indx, smooth_forward=False)
'''Fix Quaternion Discontinuity'''
quat_params = qfix(quat_params)
# (seq_len, 4)
r_rot = quat_params[:, 0].copy()
# print(r_rot[0])
'''Root Linear Velocity'''
# (seq_len - 1, 3)
velocity = (positions[1:, 0] - positions[:-1, 0]).copy()
# print(r_rot.shape, velocity.shape)
velocity = qrot_np(r_rot[1:], velocity)
'''Root Angular Velocity'''
# (seq_len - 1, 4)
r_velocity = qmul_np(r_rot[1:], qinv_np(r_rot[:-1]))
quat_params[1:, 0] = r_velocity
# (seq_len, joints_num, 4)
return quat_params, r_velocity, velocity, r_rot
def get_cont6d_params(positions):
skel = Skeleton(n_raw_offsets, kinematic_chain, "cpu")
# (seq_len, joints_num, 4)
quat_params = skel.inverse_kinematics_np(positions, face_joint_indx, smooth_forward=True)
'''Quaternion to continuous 6D'''
cont_6d_params = quaternion_to_cont6d_np(quat_params)
# (seq_len, 4)
r_rot = quat_params[:, 0].copy()
# print(r_rot[0])
'''Root Linear Velocity'''
# (seq_len - 1, 3)
velocity = (positions[1:, 0] - positions[:-1, 0]).copy()
# print(r_rot.shape, velocity.shape)
velocity = qrot_np(r_rot[1:], velocity)
'''Root Angular Velocity'''
# (seq_len - 1, 4)
r_velocity = qmul_np(r_rot[1:], qinv_np(r_rot[:-1]))
# (seq_len, joints_num, 4)
return cont_6d_params, r_velocity, velocity, r_rot
cont_6d_params, r_velocity, velocity, r_rot = get_cont6d_params(positions)
positions = get_rifke(positions)
# trejec = np.cumsum(np.concatenate([np.array([[0, 0, 0]]), velocity], axis=0), axis=0)
# r_rotations, r_pos = recover_ric_glo_np(r_velocity, velocity[:, [0, 2]])
# plt.plot(positions_b[:, 0, 0], positions_b[:, 0, 2], marker='*')
# plt.plot(ground_positions[:, 0, 0], ground_positions[:, 0, 2], marker='o', color='r')
# plt.plot(trejec[:, 0], trejec[:, 2], marker='^', color='g')
# plt.plot(r_pos[:, 0], r_pos[:, 2], marker='s', color='y')
# plt.xlabel('x')
# plt.ylabel('z')
# plt.axis('equal')
# plt.show()
'''Root height'''
root_y = positions[:, 0, 1:2]
'''Root rotation and linear velocity'''
# (seq_len-1, 1) rotation velocity along y-axis
# (seq_len-1, 2) linear velovity on xz plane
r_velocity = np.arcsin(r_velocity[:, 2:3])
l_velocity = velocity[:, [0, 2]]
# print(r_velocity.shape, l_velocity.shape, root_y.shape)
root_data = np.concatenate([r_velocity, l_velocity, root_y[:-1]], axis=-1)
'''Get Joint Rotation Representation'''
# (seq_len, (joints_num-1) *6) quaternion for skeleton joints
rot_data = cont_6d_params[:, 1:].reshape(len(cont_6d_params), -1)
'''Get Joint Rotation Invariant Position Represention'''
# (seq_len, (joints_num-1)*3) local joint position
ric_data = positions[:, 1:].reshape(len(positions), -1)
'''Get Joint Velocity Representation'''
# (seq_len-1, joints_num*3)
local_vel = qrot_np(np.repeat(r_rot[:-1, None], global_positions.shape[1], axis=1),
global_positions[1:] - global_positions[:-1])
local_vel = local_vel.reshape(len(local_vel), -1)
data = root_data
data = np.concatenate([data, ric_data[:-1]], axis=-1)
data = np.concatenate([data, rot_data[:-1]], axis=-1)
# print(dataset.shape, local_vel.shape)
data = np.concatenate([data, local_vel], axis=-1)
data = np.concatenate([data, feet_l, feet_r], axis=-1)
return data
def process_file(positions, feet_thre):
# (seq_len, joints_num, 3)
# '''Down Sample'''
# positions = positions[::ds_num]
'''Uniform Skeleton'''
positions = uniform_skeleton(positions, tgt_offsets)
'''Put on Floor'''
floor_height = positions.min(axis=0).min(axis=0)[1]
positions[:, :, 1] -= floor_height
# print(floor_height)
# plot_3d_motion("./positions_1.mp4", kinematic_chain, positions, 'title', fps=20)
'''XZ at origin'''
root_pos_init = positions[0]
root_pose_init_xz = root_pos_init[0] * np.array([1, 0, 1])
positions = positions - root_pose_init_xz
# '''Move the first pose to origin '''
# root_pos_init = positions[0]
# positions = positions - root_pos_init[0]
'''All initially face Z+'''
r_hip, l_hip, sdr_r, sdr_l = face_joint_indx
across1 = root_pos_init[r_hip] - root_pos_init[l_hip]
across2 = root_pos_init[sdr_r] - root_pos_init[sdr_l]
across = across1 + across2
across = across / np.sqrt((across ** 2).sum(axis=-1))[..., np.newaxis]
# forward (3,), rotate around y-axis
forward_init = np.cross(np.array([[0, 1, 0]]), across, axis=-1)
# forward (3,)
forward_init = forward_init / np.sqrt((forward_init ** 2).sum(axis=-1))[..., np.newaxis]
# print(forward_init)
target = np.array([[0, 0, 1]])
root_quat_init = qbetween_np(forward_init, target)
root_quat_init = np.ones(positions.shape[:-1] + (4,)) * root_quat_init
positions_b = positions.copy()
positions = qrot_np(root_quat_init, positions)
# plot_3d_motion("./positions_2.mp4", kinematic_chain, positions, 'title', fps=20)
'''New ground truth positions'''
global_positions = positions.copy()
# plt.plot(positions_b[:, 0, 0], positions_b[:, 0, 2], marker='*')
# plt.plot(positions[:, 0, 0], positions[:, 0, 2], marker='o', color='r')
# plt.xlabel('x')
# plt.ylabel('z')
# plt.axis('equal')
# plt.show()
""" Get Foot Contacts """
def foot_detect(positions, thres):
velfactor, heightfactor = np.array([thres, thres]), np.array([3.0, 2.0])
feet_l_x = (positions[1:, fid_l, 0] - positions[:-1, fid_l, 0]) ** 2
feet_l_y = (positions[1:, fid_l, 1] - positions[:-1, fid_l, 1]) ** 2
feet_l_z = (positions[1:, fid_l, 2] - positions[:-1, fid_l, 2]) ** 2
# feet_l_h = positions[:-1,fid_l,1]
# feet_l = (((feet_l_x + feet_l_y + feet_l_z) < velfactor) & (feet_l_h < heightfactor)).astype(np.float64)
feet_l = ((feet_l_x + feet_l_y + feet_l_z) < velfactor).astype(np.float64)
feet_r_x = (positions[1:, fid_r, 0] - positions[:-1, fid_r, 0]) ** 2
feet_r_y = (positions[1:, fid_r, 1] - positions[:-1, fid_r, 1]) ** 2
feet_r_z = (positions[1:, fid_r, 2] - positions[:-1, fid_r, 2]) ** 2
# feet_r_h = positions[:-1,fid_r,1]
# feet_r = (((feet_r_x + feet_r_y + feet_r_z) < velfactor) & (feet_r_h < heightfactor)).astype(np.float64)
feet_r = (((feet_r_x + feet_r_y + feet_r_z) < velfactor)).astype(np.float64)
return feet_l, feet_r
#
feet_l, feet_r = foot_detect(positions, feet_thre)
# feet_l, feet_r = foot_detect(positions, 0.002)
'''Quaternion and Cartesian representation'''
r_rot = None
def get_rifke(positions):
'''Local pose'''
positions[..., 0] -= positions[:, 0:1, 0]
positions[..., 2] -= positions[:, 0:1, 2]
'''All pose face Z+'''
positions = qrot_np(np.repeat(r_rot[:, None], positions.shape[1], axis=1), positions)
return positions
def get_quaternion(positions):
skel = Skeleton(n_raw_offsets, kinematic_chain, "cpu")
# (seq_len, joints_num, 4)
quat_params = skel.inverse_kinematics_np(positions, face_joint_indx, smooth_forward=False)
'''Fix Quaternion Discontinuity'''
quat_params = qfix(quat_params)
# (seq_len, 4)
r_rot = quat_params[:, 0].copy()
# print(r_rot[0])
'''Root Linear Velocity'''
# (seq_len - 1, 3)
velocity = (positions[1:, 0] - positions[:-1, 0]).copy()
# print(r_rot.shape, velocity.shape)
velocity = qrot_np(r_rot[1:], velocity)
'''Root Angular Velocity'''
# (seq_len - 1, 4)
r_velocity = qmul_np(r_rot[1:], qinv_np(r_rot[:-1]))
quat_params[1:, 0] = r_velocity
# (seq_len, joints_num, 4)
return quat_params, r_velocity, velocity, r_rot
def get_cont6d_params(positions):
skel = Skeleton(n_raw_offsets, kinematic_chain, "cpu")
# (seq_len, joints_num, 4)
quat_params = skel.inverse_kinematics_np(positions, face_joint_indx, smooth_forward=True)
'''Quaternion to continuous 6D'''
cont_6d_params = quaternion_to_cont6d_np(quat_params)
# (seq_len, 4)
r_rot = quat_params[:, 0].copy()
# print(r_rot[0])
'''Root Linear Velocity'''
# (seq_len - 1, 3)
velocity = (positions[1:, 0] - positions[:-1, 0]).copy()
# print(r_rot.shape, velocity.shape)
velocity = qrot_np(r_rot[1:], velocity)
'''Root Angular Velocity'''
# (seq_len - 1, 4)
r_velocity = qmul_np(r_rot[1:], qinv_np(r_rot[:-1]))
# (seq_len, joints_num, 4)
return cont_6d_params, r_velocity, velocity, r_rot
cont_6d_params, r_velocity, velocity, r_rot = get_cont6d_params(positions)
positions = get_rifke(positions)
# trejec = np.cumsum(np.concatenate([np.array([[0, 0, 0]]), velocity], axis=0), axis=0)
# r_rotations, r_pos = recover_ric_glo_np(r_velocity, velocity[:, [0, 2]])
# plt.plot(positions_b[:, 0, 0], positions_b[:, 0, 2], marker='*')
# plt.plot(ground_positions[:, 0, 0], ground_positions[:, 0, 2], marker='o', color='r')
# plt.plot(trejec[:, 0], trejec[:, 2], marker='^', color='g')
# plt.plot(r_pos[:, 0], r_pos[:, 2], marker='s', color='y')
# plt.xlabel('x')
# plt.ylabel('z')
# plt.axis('equal')
# plt.show()
'''Root height'''
root_y = positions[:, 0, 1:2]
'''Root rotation and linear velocity'''
# (seq_len-1, 1) rotation velocity along y-axis
# (seq_len-1, 2) linear velovity on xz plane
r_velocity = np.arcsin(r_velocity[:, 2:3])
l_velocity = velocity[:, [0, 2]]
# print(r_velocity.shape, l_velocity.shape, root_y.shape)
root_data = np.concatenate([r_velocity, l_velocity, root_y[:-1]], axis=-1)
'''Get Joint Rotation Representation'''
# (seq_len, (joints_num-1) *6) quaternion for skeleton joints
rot_data = cont_6d_params[:, 1:].reshape(len(cont_6d_params), -1)
'''Get Joint Rotation Invariant Position Represention'''
# (seq_len, (joints_num-1)*3) local joint position
ric_data = positions[:, 1:].reshape(len(positions), -1)
'''Get Joint Velocity Representation'''
# (seq_len-1, joints_num*3)
local_vel = qrot_np(np.repeat(r_rot[:-1, None], global_positions.shape[1], axis=1),
global_positions[1:] - global_positions[:-1])
local_vel = local_vel.reshape(len(local_vel), -1)
data = root_data
data = np.concatenate([data, ric_data[:-1]], axis=-1)
data = np.concatenate([data, rot_data[:-1]], axis=-1)
# print(dataset.shape, local_vel.shape)
data = np.concatenate([data, local_vel], axis=-1)
data = np.concatenate([data, feet_l, feet_r], axis=-1)
return data, global_positions, positions, l_velocity
# Recover global angle and positions for rotation dataset
# root_rot_velocity (B, seq_len, 1)
# root_linear_velocity (B, seq_len, 2)
# root_y (B, seq_len, 1)
# ric_data (B, seq_len, (joint_num - 1)*3)
# rot_data (B, seq_len, (joint_num - 1)*6)
# local_velocity (B, seq_len, joint_num*3)
# foot contact (B, seq_len, 4)
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_rot(data):
# dataset [bs, seqlen, 263/251] HumanML/KIT
joints_num = 22 if data.shape[-1] == 263 else 21
r_rot_quat, r_pos = recover_root_rot_pos(data)
r_pos_pad = torch.cat([r_pos, torch.zeros_like(r_pos)], dim=-1).unsqueeze(-2)
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]
cont6d_params = torch.cat([r_rot_cont6d, cont6d_params], dim=-1)
cont6d_params = cont6d_params.view(-1, joints_num, 6)
cont6d_params = torch.cat([cont6d_params, r_pos_pad], dim=-2)
return cont6d_params
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
'''
For Text2Motion Dataset
'''
'''
if __name__ == "__main__":
example_id = "000021"
# Lower legs
l_idx1, l_idx2 = 5, 8
# Right/Left foot
fid_r, fid_l = [8, 11], [7, 10]
# Face direction, r_hip, l_hip, sdr_r, sdr_l
face_joint_indx = [2, 1, 17, 16]
# l_hip, r_hip
r_hip, l_hip = 2, 1
joints_num = 22
# ds_num = 8
data_dir = '../dataset/pose_data_raw/joints/'
save_dir1 = '../dataset/pose_data_raw/new_joints/'
save_dir2 = '../dataset/pose_data_raw/new_joint_vecs/'
n_raw_offsets = torch.from_numpy(t2m_raw_offsets)
kinematic_chain = t2m_kinematic_chain
# Get offsets of target skeleton
example_data = np.load(os.path.join(data_dir, example_id + '.npy'))
example_data = example_data.reshape(len(example_data), -1, 3)
example_data = torch.from_numpy(example_data)
tgt_skel = Skeleton(n_raw_offsets, kinematic_chain, 'cpu')
# (joints_num, 3)
tgt_offsets = tgt_skel.get_offsets_joints(example_data[0])
# print(tgt_offsets)
source_list = os.listdir(data_dir)
frame_num = 0
for source_file in tqdm(source_list):
source_data = np.load(os.path.join(data_dir, source_file))[:, :joints_num]
try:
dataset, ground_positions, positions, l_velocity = process_file(source_data, 0.002)
rec_ric_data = recover_from_ric(torch.from_numpy(dataset).unsqueeze(0).float(), joints_num)
np.save(pjoin(save_dir1, source_file), rec_ric_data.squeeze().numpy())
np.save(pjoin(save_dir2, source_file), dataset)
frame_num += dataset.shape[0]
except Exception as e:
print(source_file)
print(e)
print('Total clips: %d, Frames: %d, Duration: %fm' %
(len(source_list), frame_num, frame_num / 20 / 60))
'''
if __name__ == "__main__":
example_id = "03950_gt"
# Lower legs
l_idx1, l_idx2 = 17, 18
# Right/Left foot
fid_r, fid_l = [14, 15], [19, 20]
# Face direction, r_hip, l_hip, sdr_r, sdr_l
face_joint_indx = [11, 16, 5, 8]
# l_hip, r_hip
r_hip, l_hip = 11, 16
joints_num = 21
# ds_num = 8
data_dir = '../dataset/kit_mocap_dataset/joints/'
save_dir1 = '../dataset/kit_mocap_dataset/new_joints/'
save_dir2 = '../dataset/kit_mocap_dataset/new_joint_vecs/'
n_raw_offsets = torch.from_numpy(kit_raw_offsets)
kinematic_chain = kit_kinematic_chain
'''Get offsets of target skeleton'''
example_data = np.load(os.path.join(data_dir, example_id + '.npy'))
example_data = example_data.reshape(len(example_data), -1, 3)
example_data = torch.from_numpy(example_data)
tgt_skel = Skeleton(n_raw_offsets, kinematic_chain, 'cpu')
# (joints_num, 3)
tgt_offsets = tgt_skel.get_offsets_joints(example_data[0])
# print(tgt_offsets)
source_list = os.listdir(data_dir)
frame_num = 0
'''Read source dataset'''
for source_file in tqdm(source_list):
source_data = np.load(os.path.join(data_dir, source_file))[:, :joints_num]
try:
name = ''.join(source_file[:-7].split('_')) + '.npy'
data, ground_positions, positions, l_velocity = process_file(source_data, 0.05)
rec_ric_data = recover_from_ric(torch.from_numpy(data).unsqueeze(0).float(), joints_num)
if np.isnan(rec_ric_data.numpy()).any():
print(source_file)
continue
np.save(pjoin(save_dir1, name), rec_ric_data.squeeze().numpy())
np.save(pjoin(save_dir2, name), data)
frame_num += data.shape[0]
except Exception as e:
print(source_file)
print(e)
print('Total clips: %d, Frames: %d, Duration: %fm' %
(len(source_list), frame_num, frame_num / 12.5 / 60))
@@ -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'
@@ -0,0 +1,79 @@
import numpy as np
import pickle
from os.path import join as pjoin
POS_enumerator = {
'VERB': 0,
'NOUN': 1,
'DET': 2,
'ADP': 3,
'NUM': 4,
'AUX': 5,
'PRON': 6,
'ADJ': 7,
'ADV': 8,
'Loc_VIP': 9,
'Body_VIP': 10,
'Obj_VIP': 11,
'Act_VIP': 12,
'Desc_VIP': 13,
'OTHER': 14,
}
Loc_list = ('left', 'right', 'clockwise', 'counterclockwise', 'anticlockwise', 'forward', 'back', 'backward',
'up', 'down', 'straight', 'curve')
Body_list = ('arm', 'chin', 'foot', 'feet', 'face', 'hand', 'mouth', 'leg', 'waist', 'eye', 'knee', 'shoulder', 'thigh')
Obj_List = ('stair', 'dumbbell', 'chair', 'window', 'floor', 'car', 'ball', 'handrail', 'baseball', 'basketball')
Act_list = ('walk', 'run', 'swing', 'pick', 'bring', 'kick', 'put', 'squat', 'throw', 'hop', 'dance', 'jump', 'turn',
'stumble', 'dance', 'stop', 'sit', 'lift', 'lower', 'raise', 'wash', 'stand', 'kneel', 'stroll',
'rub', 'bend', 'balance', 'flap', 'jog', 'shuffle', 'lean', 'rotate', 'spin', 'spread', 'climb')
Desc_list = ('slowly', 'carefully', 'fast', 'careful', 'slow', 'quickly', 'happy', 'angry', 'sad', 'happily', 'angrily', 'sadly')
VIP_dict = {
'Loc_VIP': Loc_list,
'Body_VIP': Body_list,
'Obj_VIP': Obj_List,
'Act_VIP': Act_list,
'Desc_VIP': Desc_list,
}
class WordVectorizer(object):
def __init__(self, meta_root, prefix):
vectors = np.load(pjoin(meta_root, '%s_data.npy'%prefix))
words = pickle.load(open(pjoin(meta_root, '%s_words.pkl'%prefix), 'rb'))
word2idx = pickle.load(open(pjoin(meta_root, '%s_idx.pkl'%prefix), 'rb'))
self.word2vec = {w: vectors[word2idx[w]] for w in words}
def _get_pos_ohot(self, pos):
pos_vec = np.zeros(len(POS_enumerator))
if pos in POS_enumerator:
pos_vec[POS_enumerator[pos]] = 1
else:
pos_vec[POS_enumerator['OTHER']] = 1
return pos_vec
def __len__(self):
return len(self.word2vec)
def __getitem__(self, item):
word, pos = item.split('/')
if word in self.word2vec:
word_vec = self.word2vec[word]
vip_pos = None
for key, values in VIP_dict.items():
if word in values:
vip_pos = key
break
if vip_pos is not None:
pos_vec = self._get_pos_ohot(vip_pos)
else:
pos_vec = self._get_pos_ohot(pos)
else:
word_vec = self.word2vec['unk']
pos_vec = self._get_pos_ohot('OTHER')
return word_vec, pos_vec
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,2 @@
from .tensors import lengths_to_mask
from .collate import collate_text_and_length, collate_pairs_and_text, collate_datastruct_and_text, collate_tensor_with_padding
@@ -0,0 +1,99 @@
# -*- 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©2020 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 typing import List, Dict
from torch import Tensor
def collate_tensor_with_padding(batch: List[Tensor]) -> Tensor:
dims = batch[0].dim()
max_size = [max([b.size(i) for b in batch]) for i in range(dims)]
size = (len(batch),) + tuple(max_size)
canvas = batch[0].new_zeros(size=size)
for i, b in enumerate(batch):
sub_tensor = canvas[i]
for d in range(dims):
sub_tensor = sub_tensor.narrow(d, 0, b.size(d))
sub_tensor.add_(b)
return canvas
def collate_datastruct_and_text(lst_elements: List) -> Dict:
collate_datastruct = lst_elements[0]["datastruct"].transforms.collate
batch = {
# Collate with padding for the datastruct
"datastruct": collate_datastruct([x["datastruct"] for x in lst_elements]),
# Collate normally for the length
"length": [x["length"] for x in lst_elements],
# Collate the text
"text": [x["text"] for x in lst_elements]}
# add keyid for example
otherkeys = [x for x in lst_elements[0].keys() if x not in batch]
for key in otherkeys:
batch[key] = [x[key] for x in lst_elements]
return batch
def collate_length_and_text(lst_elements: List) -> Dict:
batch = {
"length_0": [x["length_0"] for x in lst_elements],
"length_1": [x["length_1"] for x in lst_elements],
"length_transition": [x["length_transition"] for x in lst_elements],
"length_1_with_transition": [x["length_1_with_transition"] for x in lst_elements],
"text_0": [x["text_0"] for x in lst_elements],
"text_1": [x["text_1"] for x in lst_elements]
}
return batch
def collate_pairs_and_text(lst_elements: List, ) -> Dict:
if 'features_0' not in lst_elements[0]: # test set
collate_datastruct = lst_elements[0]["datastruct"].transforms.collate
batch = {"datastruct": collate_datastruct([x["datastruct"] for x in lst_elements]),
"length_0": [x["length_0"] for x in lst_elements],
"length_1": [x["length_1"] for x in lst_elements],
"length_transition": [x["length_transition"] for x in lst_elements],
"length_1_with_transition": [x["length_1_with_transition"] for x in lst_elements],
"text_0": [x["text_0"] for x in lst_elements],
"text_1": [x["text_1"] for x in lst_elements]
}
else:
batch = {"motion_feats_0": collate_tensor_with_padding([el["features_0"] for el in lst_elements]),
"motion_feats_1": collate_tensor_with_padding([el["features_1"] for el in lst_elements]),
"motion_feats_1_with_transition": collate_tensor_with_padding([el["features_1_with_transition"] for el in lst_elements]),
"length_0": [x["length_0"] for x in lst_elements],
"length_1": [x["length_1"] for x in lst_elements],
"length_transition": [x["length_transition"] for x in lst_elements],
"length_1_with_transition": [x["length_1_with_transition"] for x in lst_elements],
"text_0": [x["text_0"] for x in lst_elements],
"text_1": [x["text_1"] for x in lst_elements]
}
return batch
def collate_text_and_length(lst_elements: Dict) -> Dict:
batch = {"length": [x["length"] for x in lst_elements],
"text": [x["text"] for x in lst_elements]}
# add keyid for example
otherkeys = [x for x in lst_elements[0].keys() if x not in batch and x != "datastruct"]
for key in otherkeys:
batch[key] = [x[key] for x in lst_elements]
return batch
@@ -0,0 +1,72 @@
from .geometry import *
def nfeats_of(rottype):
if rottype in ["rotvec", "axisangle"]:
return 3
elif rottype in ["rotquat", "quaternion"]:
return 4
elif rottype in ["rot6d", "6drot", "rotation6d"]:
return 6
elif rottype in ["rotmat"]:
return 9
else:
return TypeError("This rotation type doesn't have features.")
def axis_angle_to(newtype, rotations):
if newtype in ["matrix"]:
rotations = axis_angle_to_matrix(rotations)
return rotations
elif newtype in ["rotmat"]:
rotations = axis_angle_to_matrix(rotations)
rotations = matrix_to("rotmat", rotations)
return rotations
elif newtype in ["rot6d", "6drot", "rotation6d"]:
rotations = axis_angle_to_matrix(rotations)
rotations = matrix_to("rot6d", rotations)
return rotations
elif newtype in ["rotquat", "quaternion"]:
rotations = axis_angle_to_quaternion(rotations)
return rotations
elif newtype in ["rotvec", "axisangle"]:
return rotations
else:
raise NotImplementedError
def matrix_to(newtype, rotations):
if newtype in ["matrix"]:
return rotations
if newtype in ["rotmat"]:
rotations = rotations.reshape((*rotations.shape[:-2], 9))
return rotations
elif newtype in ["rot6d", "6drot", "rotation6d"]:
rotations = matrix_to_rotation_6d(rotations)
return rotations
elif newtype in ["rotquat", "quaternion"]:
rotations = matrix_to_quaternion(rotations)
return rotations
elif newtype in ["rotvec", "axisangle"]:
rotations = matrix_to_axis_angle(rotations)
return rotations
else:
raise NotImplementedError
def to_matrix(oldtype, rotations):
if oldtype in ["matrix"]:
return rotations
if oldtype in ["rotmat"]:
rotations = rotations.reshape((*rotations.shape[:-2], 3, 3))
return rotations
elif oldtype in ["rot6d", "6drot", "rotation6d"]:
rotations = rotation_6d_to_matrix(rotations)
return rotations
elif oldtype in ["rotquat", "quaternion"]:
rotations = quaternion_to_matrix(rotations)
return rotations
elif oldtype in ["rotvec", "axisangle"]:
rotations = axis_angle_to_matrix(rotations)
return rotations
else:
raise NotImplementedError
@@ -0,0 +1,566 @@
# 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)
"""
# Added
def matrix_of_angles(cos, sin, inv=False, dim=2):
assert dim in [2, 3]
sin = -sin if inv else sin
if dim == 2:
row1 = torch.stack((cos, -sin), axis=-1)
row2 = torch.stack((sin, cos), axis=-1)
return torch.stack((row1, row2), axis=-2)
elif dim == 3:
row1 = torch.stack((cos, -sin, 0*cos), axis=-1)
row2 = torch.stack((sin, cos, 0*cos), axis=-1)
row3 = torch.stack((0*sin, 0*cos, 1+0*cos), axis=-1)
return torch.stack((row1, row2, row3),axis=-2)
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)
@@ -0,0 +1,26 @@
# -*- 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©2020 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 typing import List, Dict
import torch
from torch import Tensor
def lengths_to_mask(lengths: List[int], device: torch.device) -> Tensor:
lengths = torch.tensor(lengths, device=device)
max_len = max(lengths)
mask = torch.arange(max_len, device=device).expand(len(lengths), max_len) < lengths.unsqueeze(1)
return mask
@@ -0,0 +1,15 @@
from .base import Transform
from .smpl import SMPLTransform
from .xyz import XYZTransform
# rots2rfeats
from .rots2rfeats import Rots2Rfeats
from .rots2rfeats import Globalvelandy
# rots2joints
from .rots2joints import Rots2Joints
from .rots2joints import SMPLH, SMPLX
# joints2jfeats
from .joints2jfeats import Joints2Jfeats
from .joints2jfeats import Rifke

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