Add 4DHuman
This commit is contained in:
@@ -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>
|
||||
|
||||
[](https://arxiv.org/pdf/2305.20091.pdf) [](https://shubham-goel.github.io/4dhumans/) [](https://colab.research.google.com/drive/1Ex4gE5v1bPR3evfhtG7sDHxQGsWwNwby?usp=sharing) [](https://huggingface.co/spaces/brjathu/HMR2.0)
|
||||
|
||||
|
||||

|
||||
|
||||
## 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}
|
||||
}
|
||||
```
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
Reference in New Issue
Block a user