Add SpectreFaceRecon and move vert normalization of Human4D to its internal

This commit is contained in:
Hacker 17082006
2024-04-16 23:00:20 +07:00
parent f342f59dbf
commit 99cb0ca807
69 changed files with 58533 additions and 80 deletions
+95
View File
@@ -0,0 +1,95 @@
from pathlib import Path
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
CKPT_DIR_PATH = (Path(__file__).parent.parent / "ckpts").resolve()
def download_models(filename_links={}):
"""Download checkpoints and files for running inference.
"""
folder = CKPT_DIR_PATH
import os
os.makedirs(folder, exist_ok=True)
download_files = {
**{filename: [link, folder] for filename, link in 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])
+2 -1
View File
@@ -2,8 +2,9 @@ import os
from typing import Dict
from yacs.config import CfgNode as CN
from pathlib import Path
from motiondiff_modules import CKPT_DIR_PATH
CACHE_DIR_4DHUMANS = os.environ.get("4DHUMAN_CACHE", str(Path(__file__).parent.parent.parent.parent / "ckpts"))
CACHE_DIR_4DHUMANS = os.environ.get("4DHUMAN_CACHE", str(CKPT_DIR_PATH))
def to_lower(x: Dict) -> Dict:
"""
+4 -20
View File
@@ -298,22 +298,7 @@ def render_from_smpl(thetas, yfov, move_x, move_y, move_z, x_rot, y_rot, z_rot,
return np.stack(vid, axis=0), np.stack(vid_depth, axis=0)
# verts_frames: list of [num_subjects, num_verts, 3]
# cam_t_frames: list of [num_subjects, 3]
def render_from_smpl_multiple_subjects(verts_frames, cam_t_frames, focal_length, fx_offset, fy_offset, move_x, move_y, move_z, x_rot, y_rot, z_rot, frame_width, frame_height, draw_platform=True, depth_only=False, normals=False, smpl_model_path=None):
def vertices_to_trimesh(vertices, camera_translation, faces, rot_axis=[1,0,0], rot_angle=0,):
mesh = trimesh.Trimesh(vertices + camera_translation, faces.copy())
rot = trimesh.transformations.rotation_matrix(
np.radians(rot_angle), rot_axis)
mesh.apply_transform(rot)
rot = trimesh.transformations.rotation_matrix(
np.radians(180), [1, 0, 0])
mesh.apply_transform(rot)
return mesh
rot2xyz = Rotation2xyz(device="cpu", smpl_model_path=smpl_model_path)
faces = rot2xyz.smpl_model.faces
def render_from_smpl_multiple_subjects(verts_frames, faces, focal_length, fx_offset, fy_offset, move_x, move_y, move_z, x_rot, y_rot, z_rot, frame_width, frame_height, draw_platform=True, depth_only=False, normals=False, vertical_flip=True, cx=0, cy=0):
MINS = torch.stack([verts_frame.min(0).values.min(0).values for verts_frame in verts_frames if verts_frame is not None]).min(0).values
MAXS = torch.stack([verts_frame.max(0).values.max(0).values for verts_frame in verts_frames if verts_frame is not None]).max(0).values
minx = MINS[0] - 0.5
@@ -400,7 +385,7 @@ def render_from_smpl_multiple_subjects(verts_frames, cam_t_frames, focal_length,
#Build the scene
camera = pyrender.IntrinsicsCamera(fx=focal_length + fx_offset, fy=focal_length + fy_offset,
cx=frame_width / 2, cy=frame_height / 2, zfar=1e12)
cx=cx, cy=cy)
bg_color = [1, 1, 1, 0.8]
scene = pyrender.Scene(bg_color=bg_color, ambient_light=(0.4, 0.4, 0.4))
scene.add(camera, pose=camera_pose)
@@ -416,14 +401,13 @@ def render_from_smpl_multiple_subjects(verts_frames, cam_t_frames, focal_length,
pbar = comfy.utils.ProgressBar(len(verts_frames))
for i in tqdm(range(len(verts_frames))):
subjects = verts_frames[i]
cam_t_subjects = cam_t_frames[i]
mesh_nodes = []
if subjects is None:
vid.append(np.ones([frame_height, frame_width, 4], dtype=np.uint8))
vid_depth.append(np.zeros([frame_height, frame_width], dtype=np.float32))
continue
for subject_vertices, cam_t in zip(subjects, cam_t_subjects):
mesh = vertices_to_trimesh(subject_vertices, cam_t, faces)
for subject_vertices in subjects:
mesh = trimesh.Trimesh(subject_vertices, faces=faces)
mesh = pyrender.Mesh.from_trimesh(mesh, material=material)
mesh_node = pyrender.Node(mesh=mesh)
scene.add_node(mesh_node)
+437
View File
@@ -0,0 +1,437 @@
Attribution-NonCommercial-ShareAlike 4.0 International
=======================================================================
Creative Commons Corporation ("Creative Commons") is not a law firm and
does not provide legal services or legal advice. Distribution of
Creative Commons public licenses does not create a lawyer-client or
other relationship. Creative Commons makes its licenses and related
information available on an "as-is" basis. Creative Commons gives no
warranties regarding its licenses, any material licensed under their
terms and conditions, or any related information. Creative Commons
disclaims all liability for damages resulting from their use to the
fullest extent possible.
Using Creative Commons Public Licenses
Creative Commons public licenses provide a standard set of terms and
conditions that creators and other rights holders may use to share
original works of authorship and other material subject to copyright
and certain other rights specified in the public license below. The
following considerations are for informational purposes only, are not
exhaustive, and do not form part of our licenses.
Considerations for licensors: Our public licenses are
intended for use by those authorized to give the public
permission to use material in ways otherwise restricted by
copyright and certain other rights. Our licenses are
irrevocable. Licensors should read and understand the terms
and conditions of the license they choose before applying it.
Licensors should also secure all rights necessary before
applying our licenses so that the public can reuse the
material as expected. Licensors should clearly mark any
material not subject to the license. This includes other CC-
licensed material, or material used under an exception or
limitation to copyright. More considerations for licensors:
wiki.creativecommons.org/Considerations_for_licensors
Considerations for the public: By using one of our public
licenses, a licensor grants the public permission to use the
licensed material under specified terms and conditions. If
the licensor's permission is not necessary for any reason--for
example, because of any applicable exception or limitation to
copyright--then that use is not regulated by the license. Our
licenses grant only permissions under copyright and certain
other rights that a licensor has authority to grant. Use of
the licensed material may still be restricted for other
reasons, including because others have copyright or other
rights in the material. A licensor may make special requests,
such as asking that all changes be marked or described.
Although not required by our licenses, you are encouraged to
respect those requests where reasonable. More considerations
for the public:
wiki.creativecommons.org/Considerations_for_licensees
=======================================================================
Creative Commons Attribution-NonCommercial-ShareAlike 4.0 International
Public License
By exercising the Licensed Rights (defined below), You accept and agree
to be bound by the terms and conditions of this Creative Commons
Attribution-NonCommercial-ShareAlike 4.0 International Public License
("Public License"). To the extent this Public License may be
interpreted as a contract, You are granted the Licensed Rights in
consideration of Your acceptance of these terms and conditions, and the
Licensor grants You such rights in consideration of benefits the
Licensor receives from making the Licensed Material available under
these terms and conditions.
Section 1 -- Definitions.
a. Adapted Material means material subject to Copyright and Similar
Rights that is derived from or based upon the Licensed Material
and in which the Licensed Material is translated, altered,
arranged, transformed, or otherwise modified in a manner requiring
permission under the Copyright and Similar Rights held by the
Licensor. For purposes of this Public License, where the Licensed
Material is a musical work, performance, or sound recording,
Adapted Material is always produced where the Licensed Material is
synched in timed relation with a moving image.
b. Adapter's License means the license You apply to Your Copyright
and Similar Rights in Your contributions to Adapted Material in
accordance with the terms and conditions of this Public License.
c. BY-NC-SA Compatible License means a license listed at
creativecommons.org/compatiblelicenses, approved by Creative
Commons as essentially the equivalent of this Public License.
d. Copyright and Similar Rights means copyright and/or similar rights
closely related to copyright including, without limitation,
performance, broadcast, sound recording, and Sui Generis Database
Rights, without regard to how the rights are labeled or
categorized. For purposes of this Public License, the rights
specified in Section 2(b)(1)-(2) are not Copyright and Similar
Rights.
e. Effective Technological Measures means those measures that, in the
absence of proper authority, may not be circumvented under laws
fulfilling obligations under Article 11 of the WIPO Copyright
Treaty adopted on December 20, 1996, and/or similar international
agreements.
f. Exceptions and Limitations means fair use, fair dealing, and/or
any other exception or limitation to Copyright and Similar Rights
that applies to Your use of the Licensed Material.
g. License Elements means the license attributes listed in the name
of a Creative Commons Public License. The License Elements of this
Public License are Attribution, NonCommercial, and ShareAlike.
h. Licensed Material means the artistic or literary work, database,
or other material to which the Licensor applied this Public
License.
i. Licensed Rights means the rights granted to You subject to the
terms and conditions of this Public License, which are limited to
all Copyright and Similar Rights that apply to Your use of the
Licensed Material and that the Licensor has authority to license.
j. Licensor means the individual(s) or entity(ies) granting rights
under this Public License.
k. NonCommercial means not primarily intended for or directed towards
commercial advantage or monetary compensation. For purposes of
this Public License, the exchange of the Licensed Material for
other material subject to Copyright and Similar Rights by digital
file-sharing or similar means is NonCommercial provided there is
no payment of monetary compensation in connection with the
exchange.
l. Share means to provide material to the public by any means or
process that requires permission under the Licensed Rights, such
as reproduction, public display, public performance, distribution,
dissemination, communication, or importation, and to make material
available to the public including in ways that members of the
public may access the material from a place and at a time
individually chosen by them.
m. Sui Generis Database Rights means rights other than copyright
resulting from Directive 96/9/EC of the European Parliament and of
the Council of 11 March 1996 on the legal protection of databases,
as amended and/or succeeded, as well as other essentially
equivalent rights anywhere in the world.
n. You means the individual or entity exercising the Licensed Rights
under this Public License. Your has a corresponding meaning.
Section 2 -- Scope.
a. License grant.
1. Subject to the terms and conditions of this Public License,
the Licensor hereby grants You a worldwide, royalty-free,
non-sublicensable, non-exclusive, irrevocable license to
exercise the Licensed Rights in the Licensed Material to:
a. reproduce and Share the Licensed Material, in whole or
in part, for NonCommercial purposes only; and
b. produce, reproduce, and Share Adapted Material for
NonCommercial purposes only.
2. Exceptions and Limitations. For the avoidance of doubt, where
Exceptions and Limitations apply to Your use, this Public
License does not apply, and You do not need to comply with
its terms and conditions.
3. Term. The term of this Public License is specified in Section
6(a).
4. Media and formats; technical modifications allowed. The
Licensor authorizes You to exercise the Licensed Rights in
all media and formats whether now known or hereafter created,
and to make technical modifications necessary to do so. The
Licensor waives and/or agrees not to assert any right or
authority to forbid You from making technical modifications
necessary to exercise the Licensed Rights, including
technical modifications necessary to circumvent Effective
Technological Measures. For purposes of this Public License,
simply making modifications authorized by this Section 2(a)
(4) never produces Adapted Material.
5. Downstream recipients.
a. Offer from the Licensor -- Licensed Material. Every
recipient of the Licensed Material automatically
receives an offer from the Licensor to exercise the
Licensed Rights under the terms and conditions of this
Public License.
b. Additional offer from the Licensor -- Adapted Material.
Every recipient of Adapted Material from You
automatically receives an offer from the Licensor to
exercise the Licensed Rights in the Adapted Material
under the conditions of the Adapter's License You apply.
c. No downstream restrictions. You may not offer or impose
any additional or different terms or conditions on, or
apply any Effective Technological Measures to, the
Licensed Material if doing so restricts exercise of the
Licensed Rights by any recipient of the Licensed
Material.
6. No endorsement. Nothing in this Public License constitutes or
may be construed as permission to assert or imply that You
are, or that Your use of the Licensed Material is, connected
with, or sponsored, endorsed, or granted official status by,
the Licensor or others designated to receive attribution as
provided in Section 3(a)(1)(A)(i).
b. Other rights.
1. Moral rights, such as the right of integrity, are not
licensed under this Public License, nor are publicity,
privacy, and/or other similar personality rights; however, to
the extent possible, the Licensor waives and/or agrees not to
assert any such rights held by the Licensor to the limited
extent necessary to allow You to exercise the Licensed
Rights, but not otherwise.
2. Patent and trademark rights are not licensed under this
Public License.
3. To the extent possible, the Licensor waives any right to
collect royalties from You for the exercise of the Licensed
Rights, whether directly or through a collecting society
under any voluntary or waivable statutory or compulsory
licensing scheme. In all other cases the Licensor expressly
reserves any right to collect such royalties, including when
the Licensed Material is used other than for NonCommercial
purposes.
Section 3 -- License Conditions.
Your exercise of the Licensed Rights is expressly made subject to the
following conditions.
a. Attribution.
1. If You Share the Licensed Material (including in modified
form), You must:
a. retain the following if it is supplied by the Licensor
with the Licensed Material:
i. identification of the creator(s) of the Licensed
Material and any others designated to receive
attribution, in any reasonable manner requested by
the Licensor (including by pseudonym if
designated);
ii. a copyright notice;
iii. a notice that refers to this Public License;
iv. a notice that refers to the disclaimer of
warranties;
v. a URI or hyperlink to the Licensed Material to the
extent reasonably practicable;
b. indicate if You modified the Licensed Material and
retain an indication of any previous modifications; and
c. indicate the Licensed Material is licensed under this
Public License, and include the text of, or the URI or
hyperlink to, this Public License.
2. You may satisfy the conditions in Section 3(a)(1) in any
reasonable manner based on the medium, means, and context in
which You Share the Licensed Material. For example, it may be
reasonable to satisfy the conditions by providing a URI or
hyperlink to a resource that includes the required
information.
3. If requested by the Licensor, You must remove any of the
information required by Section 3(a)(1)(A) to the extent
reasonably practicable.
b. ShareAlike.
In addition to the conditions in Section 3(a), if You Share
Adapted Material You produce, the following conditions also apply.
1. The Adapter's License You apply must be a Creative Commons
license with the same License Elements, this version or
later, or a BY-NC-SA Compatible License.
2. You must include the text of, or the URI or hyperlink to, the
Adapter's License You apply. You may satisfy this condition
in any reasonable manner based on the medium, means, and
context in which You Share Adapted Material.
3. You may not offer or impose any additional or different terms
or conditions on, or apply any Effective Technological
Measures to, Adapted Material that restrict exercise of the
rights granted under the Adapter's License You apply.
Section 4 -- Sui Generis Database Rights.
Where the Licensed Rights include Sui Generis Database Rights that
apply to Your use of the Licensed Material:
a. for the avoidance of doubt, Section 2(a)(1) grants You the right
to extract, reuse, reproduce, and Share all or a substantial
portion of the contents of the database for NonCommercial purposes
only;
b. if You include all or a substantial portion of the database
contents in a database in which You have Sui Generis Database
Rights, then the database in which You have Sui Generis Database
Rights (but not its individual contents) is Adapted Material,
including for purposes of Section 3(b); and
c. You must comply with the conditions in Section 3(a) if You Share
all or a substantial portion of the contents of the database.
For the avoidance of doubt, this Section 4 supplements and does not
replace Your obligations under this Public License where the Licensed
Rights include other Copyright and Similar Rights.
Section 5 -- Disclaimer of Warranties and Limitation of Liability.
a. UNLESS OTHERWISE SEPARATELY UNDERTAKEN BY THE LICENSOR, TO THE
EXTENT POSSIBLE, THE LICENSOR OFFERS THE LICENSED MATERIAL AS-IS
AND AS-AVAILABLE, AND MAKES NO REPRESENTATIONS OR WARRANTIES OF
ANY KIND CONCERNING THE LICENSED MATERIAL, WHETHER EXPRESS,
IMPLIED, STATUTORY, OR OTHER. THIS INCLUDES, WITHOUT LIMITATION,
WARRANTIES OF TITLE, MERCHANTABILITY, FITNESS FOR A PARTICULAR
PURPOSE, NON-INFRINGEMENT, ABSENCE OF LATENT OR OTHER DEFECTS,
ACCURACY, OR THE PRESENCE OR ABSENCE OF ERRORS, WHETHER OR NOT
KNOWN OR DISCOVERABLE. WHERE DISCLAIMERS OF WARRANTIES ARE NOT
ALLOWED IN FULL OR IN PART, THIS DISCLAIMER MAY NOT APPLY TO YOU.
b. TO THE EXTENT POSSIBLE, IN NO EVENT WILL THE LICENSOR BE LIABLE
TO YOU ON ANY LEGAL THEORY (INCLUDING, WITHOUT LIMITATION,
NEGLIGENCE) OR OTHERWISE FOR ANY DIRECT, SPECIAL, INDIRECT,
INCIDENTAL, CONSEQUENTIAL, PUNITIVE, EXEMPLARY, OR OTHER LOSSES,
COSTS, EXPENSES, OR DAMAGES ARISING OUT OF THIS PUBLIC LICENSE OR
USE OF THE LICENSED MATERIAL, EVEN IF THE LICENSOR HAS BEEN
ADVISED OF THE POSSIBILITY OF SUCH LOSSES, COSTS, EXPENSES, OR
DAMAGES. WHERE A LIMITATION OF LIABILITY IS NOT ALLOWED IN FULL OR
IN PART, THIS LIMITATION MAY NOT APPLY TO YOU.
c. The disclaimer of warranties and limitation of liability provided
above shall be interpreted in a manner that, to the extent
possible, most closely approximates an absolute disclaimer and
waiver of all liability.
Section 6 -- Term and Termination.
a. This Public License applies for the term of the Copyright and
Similar Rights licensed here. However, if You fail to comply with
this Public License, then Your rights under this Public License
terminate automatically.
b. Where Your right to use the Licensed Material has terminated under
Section 6(a), it reinstates:
1. automatically as of the date the violation is cured, provided
it is cured within 30 days of Your discovery of the
violation; or
2. upon express reinstatement by the Licensor.
For the avoidance of doubt, this Section 6(b) does not affect any
right the Licensor may have to seek remedies for Your violations
of this Public License.
c. For the avoidance of doubt, the Licensor may also offer the
Licensed Material under separate terms or conditions or stop
distributing the Licensed Material at any time; however, doing so
will not terminate this Public License.
d. Sections 1, 5, 6, 7, and 8 survive termination of this Public
License.
Section 7 -- Other Terms and Conditions.
a. The Licensor shall not be bound by any additional or different
terms or conditions communicated by You unless expressly agreed.
b. Any arrangements, understandings, or agreements regarding the
Licensed Material not stated herein are separate from and
independent of the terms and conditions of this Public License.
Section 8 -- Interpretation.
a. For the avoidance of doubt, this Public License does not, and
shall not be interpreted to, reduce, limit, restrict, or impose
conditions on any use of the Licensed Material that could lawfully
be made without permission under this Public License.
b. To the extent possible, if any provision of this Public License is
deemed unenforceable, it shall be automatically reformed to the
minimum extent necessary to make it enforceable. If the provision
cannot be reformed, it shall be severed from this Public License
without affecting the enforceability of the remaining terms and
conditions.
c. No term or condition of this Public License will be waived and no
failure to comply consented to unless expressly agreed to by the
Licensor.
d. Nothing in this Public License constitutes or may be interpreted
as a limitation upon, or waiver of, any privileges and immunities
that apply to the Licensor or You, including from the legal
processes of any jurisdiction or authority.
=======================================================================
Creative Commons is not a party to its public
licenses. Notwithstanding, Creative Commons may elect to apply one of
its public licenses to material it publishes and in those instances
will be considered the “Licensor.” The text of the Creative Commons
public licenses is dedicated to the public domain under the CC0 Public
Domain Dedication. Except for the limited purpose of indicating that
material is shared under a Creative Commons public license or as
otherwise permitted by the Creative Commons policies published at
creativecommons.org/policies, Creative Commons does not authorize the
use of the trademark "Creative Commons" or any other trademark or logo
of Creative Commons without its prior written consent including,
without limitation, in connection with any unauthorized modifications
to any of its public licenses or any other arrangements,
understandings, or agreements concerning use of licensed material. For
the avoidance of doubt, this paragraph does not form part of the
public licenses.
Creative Commons may be contacted at creativecommons.org.
+144
View File
@@ -0,0 +1,144 @@
<div align="center">
# SPECTRE: Visual Speech-Aware Perceptual 3D Facial Expression Reconstruction from Videos
[![Paper](https://img.shields.io/badge/arXiv-2207.11094-brightgreen)](https://arxiv.org/abs/2207.11094)
&nbsp; [![Project WebPage](https://img.shields.io/badge/Project-webpage-blue)](https://filby89.github.io/spectre/)
&nbsp; <a href='https://youtu.be/P1kqrxWNizI'>
<img src='https://img.shields.io/badge/Youtube-Video-red?style=flat&logo=youtube&logoColor=red' alt='Youtube Video'>
</a>
</div>
<p align="center">
<img src="samples/visualizations/M003_level_1_angry_014_grid.gif">
<img src="samples/visualizations/test_BImnT7lcLDE_00003_grid.gif">
</p>
<p align="center">
<img src="cover.png">
</p>
<p align="center"> Our method performs visual-speech aware 3D reconstruction so that speech perception from the original footage is preserved in the reconstructed talking head. On the left we include the word/phrase being said for each example. <p align="center">
This is the official Pytorch implementation of the paper:
```
Visual Speech-Aware Perceptual 3D Facial Expression Reconstruction from Videos
Panagiotis P. Filntisis, George Retsinas, Foivos Paraperas-Papantoniou, Athanasios Katsamanis, Anastasios Roussos, and Petros Maragos
arXiv 2022
```
## Installation
Clone the repo and its submodules:
```bash
git clone --recurse-submodules -j4 https://github.com/filby89/spectre
cd spectre
```
You need to have installed a working version of Pytorch with Python 3.6 or higher and Pytorch 3D. You can use the following commands to create a working installation:
```bash
conda create -n "spectre" python=3.8
conda install -c pytorch pytorch=1.11.0 torchvision torchaudio # you might need to select cudatoolkit version here by adding e.g. cudatoolkit=11.3
conda install -c conda-forge -c fvcore fvcore iopath
conda install pytorch3d -c pytorch3d
pip install -r requirements.txt # install the rest of the requirements
```
Installing a working setup of Pytorch3d with Pytorch can be a bit tricky. For development we used Pytorch3d 0.6.1 with Pytorch 1.10.0.
PyTorch3d 0.6.2 with pytorch 1.11.0 are also compatible.
Install the face_alignment and face_detection packages:
```bash
cd external/face_alignment
pip install -e .
cd ../face_detection
git lfs pull
pip install -e .
cd ../..
```
You may need to install git-lfs to run the above commands. [More details](https://stackoverflow.com/questions/48734119/git-lfs-is-not-a-git-command-unclear)
```bash
curl -s https://packagecloud.io/install/repositories/github/git-lfs/script.deb.sh | sudo bash
sudo apt-get install git-lfs
```
Download the FLAME model and the pretrained SPECTRE model:
```bash
pip install gdown
bash quick_install.sh
```
## Demo
Samples are included in ``samples`` folder. You can run the demo by running
```bash
python demo.py --input samples/LRS3/0Fi83BHQsMA_00002.mp4 --audio
```
The audio flag extracts audio from the input video and puts it in the output shape video for visualization purposes (ffmpeg is required for video creation).
## Training and Testing
In order to train the model you need to download the `trainval` and `test` sets of the [LRS3 dataset](https://www.robots.ox.ac.uk/~vgg/data/lip_reading/lrs3.html). After downloading
the dataset, run the following command to extract frames and audio from the videos (audio is not needed for training but it is nice for visualizing the result):
```bash
python utils/extract_frames_and_audio.py --dataset_path ./data/LRS3
```
After downloading and preprocessing the dataset, download the rest needed assets:
```bash
bash get_training_data.sh
```
This command downloads the original [DECA](https://github.com/YadiraF/DECA/) pretrained model,
the ResNet50 emotion recognition model provided by [EMOCA](https://github.com/radekd91/emoca),
the pretrained lipreading model and detected landmarks for the videos of the LRS3 dataset provided by [Visual_Speech_Recognition_for_Multiple_Languages](https://github.com/mpc001/Visual_Speech_Recognition_for_Multiple_Languages).
Finally, you need to create a texture model using the repository [BFM_to_FLAME](https://github.com/TimoBolkart/BFM_to_FLAME#create-texture-model). Due
to licencing reasons we are not allowed to share it to you.
Now, you can run the following command to train the model:
```bash
python main.py --output_dir logs --landmark 50 --relative_landmark 25 --lipread 2 --expression 0.5 --epochs 6 --LRS3_path data/LRS3 --LRS3_landmarks_path data/LRS3_landmarks
```
and then test it on the LRS3 dataset test set:
```bash
python main.py --test --output_dir logs --model_path logs/model.tar --LRS3_path data/LRS3 --LRS3_landmarks_path data/LRS3_landmarks
```
and run lipreading with AV-hubert:
```bash
# and run lipreading with our script
python utils/run_av_hubert.py --videos "logs/test_videos_000000/*_mouth.avi --LRS3_path data/LRS3"
```
## Acknowledgements
This repo is has been heavily based on the original implementation of [DECA](https://github.com/YadiraF/DECA/). We also acknowledge the following
repositories which we have benefited greatly from as well:
- [EMOCA](https://github.com/radekd91/emoca)
- [face_alignment](https://github.com/hhj1897/face_alignment)
- [face_detection](https://github.com/hhj1897/face_detection)
- [Visual_Speech_Recognition_for_Multiple_Languages](https://github.com/mpc001/Visual_Speech_Recognition_for_Multiple_Languages)
## Citation
If your research benefits from this repository, consider citing the following:
```
@misc{filntisis2022visual,
title = {Visual Speech-Aware Perceptual 3D Facial Expression Reconstruction from Videos},
author = {Filntisis, Panagiotis P. and Retsinas, George and Paraperas-Papantoniou, Foivos and Katsamanis, Athanasios and Roussos, Anastasios and Maragos, Petros},
publisher = {arXiv},
year = {2022},
}
```
+182
View File
@@ -0,0 +1,182 @@
'''
Default config for SPECTRE - adapted from DECA
'''
from yacs.config import CfgNode as CN
import argparse
import yaml
import os
cfg = CN()
cfg.project_dir = os.path.abspath(os.path.join(os.path.dirname(__file__), 'src', '..'))
cfg.device = 'cuda'
cfg.device_ids = '0'
cfg.pretrained_modelpath = os.path.join(cfg.project_dir, 'data', 'deca_model.tar')
cfg.output_dir = ''
cfg.rasterizer_type = 'pytorch3d'
# ---------------------------------------------------------------------------- #
# Options for FLAME and from original DECA
# ---------------------------------------------------------------------------- #
cfg.model = CN()
cfg.model.topology_path = os.path.join(cfg.project_dir, 'data' , 'head_template.obj')
# texture data original from http://files.is.tue.mpg.de/tbolkart/FLAME/FLAME_texture_data.zip
cfg.model.dense_template_path = os.path.join(cfg.project_dir, 'data', 'texture_data_256.npy')
cfg.model.fixed_displacement_path = os.path.join(cfg.project_dir, 'data', 'fixed_displacement_256.npy')
cfg.model.flame_model_path = os.path.join(cfg.project_dir, 'data', 'FLAME2020', 'generic_model.pkl')
cfg.model.flame_lmk_embedding_path = os.path.join(cfg.project_dir, 'data', 'landmark_embedding.npy')
cfg.model.face_mask_path = os.path.join(cfg.project_dir, 'data', 'uv_face_mask.png')
cfg.model.face_eye_mask_path = os.path.join(cfg.project_dir, 'data', 'uv_face_eye_mask.png')
cfg.model.mean_tex_path = os.path.join(cfg.project_dir, 'data', 'mean_texture.jpg')
cfg.model.tex_path = os.path.join(cfg.project_dir, 'data', 'FLAME_albedo_from_BFM.npz')
cfg.model.tex_type = 'BFM' # BFM, FLAME, albedoMM
cfg.model.uv_size = 256
cfg.model.param_list = ['shape', 'tex', 'exp', 'pose', 'cam', 'light']
cfg.model.n_shape = 100
cfg.model.n_tex = 50
cfg.model.n_exp = 50
cfg.model.n_cam = 3
cfg.model.n_pose = 6
cfg.model.n_light = 27
cfg.model.jaw_type = 'aa' # default use axis angle, another option: euler. Note that: aa is not stable in the beginning
cfg.model.model_type = "SPECTRE"
cfg.model.temporal = True
# ---------------------------------------------------------------------------- #
# Options for Dataset
# ---------------------------------------------------------------------------- #
cfg.dataset = CN()
cfg.dataset.LRS3_path = "/gpu-data3/filby/LRS3"
cfg.dataset.LRS3_landmarks_path = "../Visual_Speech_Recognition_for_Multiple_Languages/landmarks/LRS3/LRS3_landmarks"
cfg.dataset.LRS3_path = "/gpu-data3/filby/LRS3"
cfg.dataset.LRS3_landmarks_path = "../Visual_Speech_Recognition_for_Multiple_Languages/landmarks/LRS3/LRS3_landmarks"
cfg.dataset.LRS3_path = "/gpu-data3/filby/LRS3"
cfg.dataset.LRS3_landmarks_path = "../Visual_Speech_Recognition_for_Multiple_Languages/landmarks/LRS3/LRS3_landmarks"
cfg.dataset.batch_size = 1
cfg.dataset.K = 20
cfg.dataset.num_workers = 8
cfg.dataset.image_size = 224
cfg.dataset.scale_min = 1.4
cfg.dataset.scale_max = 1.8
cfg.dataset.trans_scale = 0.
cfg.dataset.fps = 25
cfg.dataset.test_datasets = ['LRS3']
# ---------------------------------------------------------------------------- #
# Options for training
# ---------------------------------------------------------------------------- #
cfg.train = CN()
cfg.train.max_epochs = 6
cfg.train.log_dir = 'logs'
cfg.train.log_steps = 10
cfg.train.vis_dir = 'train_images'
cfg.train.vis_steps = 500
cfg.train.write_summary = True
cfg.train.checkpoint_steps = 10000
cfg.train.val_vis_dir = 'val_images'
cfg.train.evaluation_steps = 10000
# ---------------------------------------------------------------------------- #
# Options for Losses
# ---------------------------------------------------------------------------- #
cfg.loss = CN()
cfg.loss.train = CN()
cfg.model.use_tex = True
cfg.model.regularization_type = 'nonlinear'
cfg.model.backbone = 'mobilenetv2' # perceptual encoder backbone
cfg.loss.train.landmark = 50
cfg.loss.train.lip_landmarks = 0
cfg.loss.train.relative_landmark = 50# 50
cfg.loss.train.photometric_texture = 0
cfg.loss.train.lipread = 2
cfg.loss.train.jaw_reg = 200
cfg.train.lr = 5e-5
cfg.loss.train.expression = 0.5
cfg.test_mode = False
def get_cfg_defaults():
"""Get a yacs CfgNode object with default values for my_project."""
# Return a clone so that the defaults will not be altered
# This is for the "local variable" use pattern
return cfg.clone()
def update_cfg(cfg, cfg_file):
cfg.merge_from_file(cfg_file)
return cfg.clone()
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument('--output_dir', type=str, help='output path')
parser.add_argument('--LRS3_path', default=None, type=str, help='path to LRS3 dataset')
parser.add_argument('--LRS3_landmarks_path', default=None, type=str, help='path to LRS3 landmarks')
parser.add_argument('--model_path', default=None, help='path to pretrained model')
parser.add_argument('--batch-size', type=int, default=1, help='the batch size')
parser.add_argument('--epochs', type=int, default=6, help='number of epochs to train for')
parser.add_argument('--K', type=int, default=20, help='length of sampled frame sequence')
parser.add_argument('--lipread', type=float, default=None, help='lipread loss weight')
parser.add_argument('--expression', type=float, default=None, help='expression loss weight')
parser.add_argument('--lr', type=float, default=None, help='learning rate')
parser.add_argument('--landmark', type=float, default=None, help='landmark loss weight')
parser.add_argument('--relative_landmark', type=float, default=None, help='relative landmark loss weight')
parser.add_argument('--backbone', type=str, default='mobilenetv2', choices=['mobilenetv2', 'resnet50'])
parser.add_argument('--test', action='store_true', help='test mode')
parser.add_argument('--test_datasets', type=str, nargs='+', default=['LRS3'], help='test datasets')
args = parser.parse_args()
cfg = get_cfg_defaults()
cfg.output_dir = args.output_dir
if args.model_path is not None:
cfg.pretrained_modelpath = args.model_path
if args.batch_size is not None:
cfg.dataset.batch_size = args.batch_size
cfg.dataset.K = args.K
if args.landmark is not None:
cfg.loss.train.landmark = args.landmark
if args.relative_landmark is not None:
cfg.loss.train.relative_landmark = args.relative_landmark
if args.lipread is not None:
cfg.loss.train.lipread = args.lipread
if args.expression is not None:
cfg.loss.train.expression = args.expression
if args.lr is not None:
cfg.train.lr = args.lr
if args.epochs is not None:
cfg.train.max_epochs = args.epochs
if args.LRS3_path is not None:
cfg.dataset.LRS3_path = args.LRS3_path
if args.LRS3_landmarks_path is not None:
cfg.dataset.LRS3_landmarks_path = args.LRS3_landmarks_path
cfg.model.backbone = args.backbone
cfg.test_mode = args.test
cfg.test_datasets = args.test_datasets
return cfg
@@ -0,0 +1,18 @@
[input]
modality=video
v_fps=25
[model]
v_fps=25
model_path=data/LRS3_V_WER32.3/model.pth
model_conf=data/LRS3_V_WER32.3/model.json
rnnlm=
rnnlm_conf=
[decode]
beam_size=1
penalty=0.5
maxlenratio=0.0
minlenratio=0.0
ctc_weight=0.1
lm_weight=0.6
Binary file not shown.
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
Binary file not shown.

After

Width:  |  Height:  |  Size: 52 KiB

@@ -0,0 +1,67 @@
b,b,voiced bilabial plosive,bed,p
d,d,voiced alveolar plosive,dig,t
d͡ʒ,dZ,voiced postalveolar affricate,jump,S
dʒ,dZ,voiced postalveolar affricate,,S
ð,D,voiced dental fricative,then,T
f,f,voiceless labiodental fricative,five,f
ɡ,g,voiced velar plosive,game,k
h,h,voiceless glottal fricative,house,k
j,j,palatal approximant,yes,i
k,k,voiceless velar plosive,cat,k
l,l,alveolar lateral approximant,lay,t
ɾ,,alveolar flap,,t
m,m,bilabial nasal,mouse,p
n,n,alveolar nasal,nap,t
ŋ,N,velar nasal,thing,k
p,p,voiceless bilabial plosive,speak,p
ɹ,r\,alveolar approximant,red,r
ɹ̩,r\,alveolar approximant,red,r
s,s,voiceless alveolar fricative,seem,s
ʃ,S,voiceless postalveolar fricative,ship,S
t,t,voiceless alveolar plosive,trap,t
t͡ʃ,tS,voiceless postalveolar affricate,chart,S
tʃ,,,,S
θ,T,voiceless dental fricative,thin,T
v,v,voiced labiodental fricative,vest,f
w,w,labial-velar approximant,west,u
z,z,voiced alveolar fricative,zero,s
ʒ,Z,voiced postalveolar fricative,vision,S
ə,@,mid-central vowel,arena,@
ɚ,@`,mid-central r-colored vowel,reader,@
æ,{,near open-front unrounded vowel,trap,a
aɪ,aI,diphthong,price,a
aʊ,aU,diphthong,mouth,a
ɑ,A,long open-back unrounded vowel,father,a
ɑː,A,long open-back unrounded vowel,father,a
ɐ,,near-open central vowel,,a
eɪ,eI,diphthong,face,e
ɝ,3`,open mid-central unrounded r-colored vowel,nurse,E
ɜː,,long open mid-central unrounded vowel,,E
ɛ,E,open mid-front unrounded vowel,dress,E
i,i,long close front unrounded vowel,fleece,i
iː,,long close front unrounded vowel,,i
ɪ,I,near-close near-front unrounded vowel,kit,i
iə,,,,i
ᵻ,,,,i
oʊ,oU,diphthong,goat,o
ɔ,O,long open mid-back rounded vowel,thought,O
ɔː,,long open mid-back rounded vowel,,O
ɔɪ,OI,diphthong,choice,O
u,u,long close-back rounded vowel,goose,u
uː,,long close-back rounded vowel,goose,u
ʊ,U,near-close near-back rounded vowel,foot,u
ʌ,V,open-mid-back unrounded vowel,strut,E
ɛɹ,,,,er
ʊɹ,,,,er
ɔːɹ,,,,Or
ɑːɹ,,,,ar
əl,,,,@t
oːɹ,,,,Or
ɪɹ,,,,ir
oː,,,,O
o,,,,O
e,,,,E
a,,,,a
n̩,,,,t
ʔ,,,,
aɪə,,,,a
1 b b voiced bilabial plosive bed p
2 d d voiced alveolar plosive dig t
3 d͡ʒ dZ voiced postalveolar affricate jump S
4 dʒ dZ voiced postalveolar affricate S
5 ð D voiced dental fricative then T
6 f f voiceless labiodental fricative five f
7 ɡ g voiced velar plosive game k
8 h h voiceless glottal fricative house k
9 j j palatal approximant yes i
10 k k voiceless velar plosive cat k
11 l l alveolar lateral approximant lay t
12 ɾ alveolar flap t
13 m m bilabial nasal mouse p
14 n n alveolar nasal nap t
15 ŋ N velar nasal thing k
16 p p voiceless bilabial plosive speak p
17 ɹ r\ alveolar approximant red r
18 ɹ̩ r\ alveolar approximant red r
19 s s voiceless alveolar fricative seem s
20 ʃ S voiceless postalveolar fricative ship S
21 t t voiceless alveolar plosive trap t
22 t͡ʃ tS voiceless postalveolar affricate chart S
23 tʃ S
24 θ T voiceless dental fricative thin T
25 v v voiced labiodental fricative vest f
26 w w labial-velar approximant west u
27 z z voiced alveolar fricative zero s
28 ʒ Z voiced postalveolar fricative vision S
29 ə @ mid-central vowel arena @
30 ɚ @` mid-central r-colored vowel reader @
31 æ { near open-front unrounded vowel trap a
32 aɪ aI diphthong price a
33 aʊ aU diphthong mouth a
34 ɑ A long open-back unrounded vowel father a
35 ɑː A long open-back unrounded vowel father a
36 ɐ near-open central vowel a
37 eɪ eI diphthong face e
38 ɝ 3` open mid-central unrounded r-colored vowel nurse E
39 ɜː long open mid-central unrounded vowel E
40 ɛ E open mid-front unrounded vowel dress E
41 i i long close front unrounded vowel fleece i
42 iː long close front unrounded vowel i
43 ɪ I near-close near-front unrounded vowel kit i
44 iə i
45 ᵻ i
46 oʊ oU diphthong goat o
47 ɔ O long open mid-back rounded vowel thought O
48 ɔː long open mid-back rounded vowel O
49 ɔɪ OI diphthong choice O
50 u u long close-back rounded vowel goose u
51 uː long close-back rounded vowel goose u
52 ʊ U near-close near-back rounded vowel foot u
53 ʌ V open-mid-back unrounded vowel strut E
54 ɛɹ er
55 ʊɹ er
56 ɔːɹ Or
57 ɑːɹ ar
58 əl @t
59 oːɹ Or
60 ɪɹ ir
61 oː O
62 o O
63 e E
64 a a
65 n̩ t
66 ʔ
67 aɪə a
Binary file not shown.
Binary file not shown.

After

Width:  |  Height:  |  Size: 11 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 12 KiB

+231
View File
@@ -0,0 +1,231 @@
# -*- coding: utf-8 -*-
import os, sys
import argparse
# sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
import os, sys
import torch
import numpy as np
import cv2
from skimage.transform import estimate_transform, warp, resize, rescale
import scipy.io
import collections
from tqdm import tqdm
from datasets.data_utils import landmarks_interpolate
from src.spectre import SPECTRE
from config import cfg as spectre_cfg
from src.utils.util import tensor2video
import torchvision
def extract_frames(video_path, detect_landmarks=True):
videofolder = os.path.splitext(video_path)[0]
os.makedirs(videofolder, exist_ok=True)
vidcap = cv2.VideoCapture(video_path)
if detect_landmarks is True:
from .tracker.face_tracker import FaceTracker
from .tracker.utils import get_landmarks
face_tracker = FaceTracker()
imagepath_list = []
count = 0
face_info = collections.defaultdict(list)
fps = vidcap.get(cv2.CAP_PROP_FPS)
with tqdm(total=int(vidcap.get(cv2.CAP_PROP_FRAME_COUNT))) as pbar:
while True:
success, image = vidcap.read()
if not success:
break
if detect_landmarks is True:
detected_faces = face_tracker.face_detector(image, rgb=False)
# -- face alignment
landmarks, scores = face_tracker.landmark_detector(image, detected_faces, rgb=False)
face_info['bbox'].append(detected_faces)
face_info['landmarks'].append(landmarks)
face_info['landmarks_scores'].append(scores)
imagepath = os.path.join(videofolder, f'{count:06d}.jpg')
cv2.imwrite(imagepath, image) # save frame as JPEG file
count += 1
imagepath_list.append(imagepath)
pbar.update(1)
pbar.set_description("Preprocessing frame %d" % count)
landmarks = get_landmarks(face_info)
print('video frames are stored in {}'.format(videofolder))
return imagepath_list, landmarks, videofolder, fps
def crop_face(frame, landmarks, scale=1.0):
image_size = 224
left = np.min(landmarks[:, 0])
right = np.max(landmarks[:, 0])
top = np.min(landmarks[:, 1])
bottom = np.max(landmarks[:, 1])
h, w, _ = frame.shape
old_size = (right - left + bottom - top) / 2
center = np.array([right - (right - left) / 2.0, bottom - (bottom - top) / 2.0])
size = int(old_size * scale)
src_pts = np.array([[center[0] - size / 2, center[1] - size / 2], [center[0] - size / 2, center[1] + size / 2],
[center[0] + size / 2, center[1] - size / 2]])
DST_PTS = np.array([[0, 0], [0, image_size - 1], [image_size - 1, 0]])
tform = estimate_transform('similarity', src_pts, DST_PTS)
return tform
def main(args):
args.crop_face = True
spectre_cfg.pretrained_modelpath = "pretrained/spectre_model.tar"
spectre_cfg.model.use_tex = False
spectre = SPECTRE(spectre_cfg, args.device)
spectre.eval()
image_paths, landmarks, videofolder, fps = extract_frames(args.input, detect_landmarks=args.crop_face)
if args.crop_face:
landmarks = landmarks_interpolate(landmarks)
if landmarks is None:
print('No faces detected in input {}'.format(args.input))
original_video_length = len(image_paths)
""" SPECTRE uses a temporal convolution of size 5.
Thus, in order to predict the parameters for a contiguous video with need to
process the video in chunks of overlap 2, dropping values which were computed from the
temporal kernel which uses pad 'same'. For the start and end of the video we
pad using the first and last frame of the video.
e.g., consider a video of size 48 frames and we want to predict it in chunks of 20 frames
(due to memory limitations). We first pad the video two frames at the start and end using
the first and last frames correspondingly, making the video 52 frames length.
Then we process independently the following chunks:
[[ 0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19]
[16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35]
[32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51]]
In the first chunk, after computing the 3DMM params we drop 0,1 and 18,19, since they were computed
from the temporal kernel with padding (we followed the same procedure in training and computed loss
only from valid outputs of the temporal kernel) In the second chunk, we drop 16,17 and 34,35, and in
the last chunk we drop 32,33 and 50,51. As a result we get:
[2..17], [18..33], [34..49] (end included) which correspond to all frames of the original video
(removing the initial padding).
"""
# pad
image_paths.insert(0,image_paths[0])
image_paths.insert(0,image_paths[0])
image_paths.append(image_paths[-1])
image_paths.append(image_paths[-1])
landmarks.insert(0,landmarks[0])
landmarks.insert(0,landmarks[0])
landmarks.append(landmarks[-1])
landmarks.append(landmarks[-1])
landmarks = np.array(landmarks)
L = 50 # chunk size
# create lists of overlapping indices
indices = list(range(len(image_paths)))
overlapping_indices = [indices[i: i + L] for i in range(0, len(indices), L-4)]
if len(overlapping_indices[-1]) < 5:
# if the last chunk has less than 5 frames, pad it with the semilast frame
overlapping_indices[-2] = overlapping_indices[-2] + overlapping_indices[-1]
overlapping_indices[-2] = np.unique(overlapping_indices[-2]).tolist()
overlapping_indices = overlapping_indices[:-1]
overlapping_indices = np.array(overlapping_indices)
image_paths = np.array(image_paths) # do this to index with multiple indices
all_shape_images = []
all_images = []
with torch.no_grad():
for chunk_id in range(len(overlapping_indices)):
print('Processing frames {} to {}'.format(overlapping_indices[chunk_id][0], overlapping_indices[chunk_id][-1]))
image_paths_chunk = image_paths[overlapping_indices[chunk_id]]
landmarks_chunk = landmarks[overlapping_indices[chunk_id]] if args.crop_face else None
images_list = []
""" load each image and crop it around the face if necessary """
for j in range(len(image_paths_chunk)):
frame = cv2.imread(image_paths_chunk[j])
frame = cv2.cvtColor(frame,cv2.COLOR_BGR2RGB)
kpt = landmarks_chunk[j]
tform = crop_face(frame,kpt,scale=1.6)
cropped_image = warp(frame, tform.inverse, output_shape=(224, 224))
images_list.append(cropped_image.transpose(2,0,1))
images_array = torch.from_numpy(np.array(images_list)).type(dtype = torch.float32).to(args.device) #K,224,224,3
codedict, initial_deca_exp, initial_deca_jaw = spectre.encode(images_array)
codedict['exp'] = codedict['exp'] + initial_deca_exp
codedict['pose'][..., 3:] = codedict['pose'][..., 3:] + initial_deca_jaw
for key in codedict.keys():
""" filter out invalid indices - see explanation at the top of the function """
if chunk_id == 0 and chunk_id == len(overlapping_indices) - 1:
pass
elif chunk_id == 0:
codedict[key] = codedict[key][:-2]
elif chunk_id == len(overlapping_indices) - 1:
codedict[key] = codedict[key][2:]
else:
codedict[key] = codedict[key][2:-2]
opdict, visdict = spectre.decode(codedict, rendering=True, vis_lmk=False, return_vis=True)
all_shape_images.append(visdict['shape_images'].detach().cpu())
all_images.append(codedict['images'].detach().cpu())
vid_shape = tensor2video(torch.cat(all_shape_images, dim=0))[2:-2] # remove padding
vid_orig = tensor2video(torch.cat(all_images, dim=0))[2:-2] # remove padding
grid_vid = np.concatenate((vid_shape, vid_orig), axis=2)
assert original_video_length == len(vid_shape)
if args.audio:
import librosa
wav, sr = librosa.load(args.input)
wav = torch.FloatTensor(wav)
if len(wav.shape) == 1:
wav = wav.unsqueeze(0)
torchvision.io.write_video(videofolder+"_shape.mp4", vid_shape, fps=fps, audio_codec='aac', audio_array=wav, audio_fps=sr)
torchvision.io.write_video(videofolder+"_grid.mp4", grid_vid, fps=fps,
audio_codec='aac', audio_array=wav, audio_fps=sr)
else:
torchvision.io.write_video(videofolder+"_shape.mp4", vid_shape, fps=fps)
torchvision.io.write_video(videofolder+"_grid.mp4", grid_vid, fps=fps)
if __name__ == '__main__':
parser = argparse.ArgumentParser(description='DECA: Detailed Expression Capture and Animation')
parser.add_argument('-i', '--input', default='examples', type=str,
help='path to the test data, can be image folder, image path, image list, video')
# parser.add_argument('-o', '--outpath', default='examples/results', type=str,
# help='path to the output directory, where results(obj, txt files) will be stored.')
parser.add_argument('--device', default='cuda', type=str,
help='set device, cpu for using cpu')
parser.add_argument('--audio', action='store_true',
help='extract audio from the original video and add it to the output video')
main(parser.parse_args())
@@ -0,0 +1,4 @@
from .fan import FANPredictor
__version__ = '0.1.0'
@@ -0,0 +1 @@
from .fan_predictor import FANPredictor
@@ -0,0 +1,179 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
def conv3x3(in_planes, out_planes, strd=1, padding=1, bias=False):
"3x3 convolution with padding"
return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=strd, padding=padding, bias=bias)
class ConvBlock(nn.Module):
def __init__(self, in_planes, out_planes, use_instance_norm):
super(ConvBlock, self).__init__()
self.bn1 = nn.InstanceNorm2d(in_planes) if use_instance_norm else nn.BatchNorm2d(in_planes)
self.conv1 = conv3x3(in_planes, int(out_planes / 2))
self.bn2 = (nn.InstanceNorm2d(int(out_planes / 2)) if use_instance_norm
else nn.BatchNorm2d(int(out_planes / 2)))
self.conv2 = conv3x3(int(out_planes / 2), int(out_planes / 4))
self.bn3 = (nn.InstanceNorm2d(int(out_planes / 4)) if use_instance_norm
else nn.BatchNorm2d(int(out_planes / 4)))
self.conv3 = conv3x3(int(out_planes / 4), int(out_planes / 4))
if in_planes != out_planes:
self.downsample = nn.Sequential(nn.InstanceNorm2d(in_planes) if use_instance_norm
else nn.BatchNorm2d(in_planes),
nn.ReLU(True),
nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=1, bias=False))
else:
self.downsample = None
def forward(self, x):
residual = x
out1 = self.bn1(x)
out1 = F.relu(out1, True)
out1 = self.conv1(out1)
out2 = self.bn2(out1)
out2 = F.relu(out2, True)
out2 = self.conv2(out2)
out3 = self.bn3(out2)
out3 = F.relu(out3, True)
out3 = self.conv3(out3)
out3 = torch.cat((out1, out2, out3), 1)
if self.downsample is not None:
residual = self.downsample(residual)
out3 += residual
return out3
class HourGlass(nn.Module):
def __init__(self, config):
super(HourGlass, self).__init__()
self.config = config
self._generate_network(self.config.hg_depth)
def _generate_network(self, level):
self.add_module('b1_' + str(level), ConvBlock(self.config.hg_num_features,
self.config.hg_num_features,
self.config.use_instance_norm))
self.add_module('b2_' + str(level), ConvBlock(self.config.hg_num_features,
self.config.hg_num_features,
self.config.use_instance_norm))
if level > 1:
self._generate_network(level - 1)
else:
self.add_module('b2_plus_' + str(level),ConvBlock(self.config.hg_num_features,
self.config.hg_num_features,
self.config.use_instance_norm))
self.add_module('b3_' + str(level), ConvBlock(self.config.hg_num_features,
self.config.hg_num_features,
self.config.use_instance_norm))
def _forward(self, level, inp):
up1 = inp
up1 = self._modules['b1_' + str(level)](up1)
if self.config.use_avg_pool:
low1 = F.avg_pool2d(inp, 2)
else:
low1 = F.max_pool2d(inp, 2)
low1 = self._modules['b2_' + str(level)](low1)
if level > 1:
low2 = self._forward(level - 1, low1)
else:
low2 = low1
low2 = self._modules['b2_plus_' + str(level)](low2)
low3 = low2
low3 = self._modules['b3_' + str(level)](low3)
up2 = F.interpolate(low3, scale_factor=2, mode='nearest')
return up1 + up2
def forward(self, x):
return self._forward(self.config.hg_depth, x)
class FAN(nn.Module):
def __init__(self, config):
super(FAN, self).__init__()
self.config = config
# Stem
self.conv1 = nn.Conv2d(3, 64, kernel_size=self.config.stem_conv_kernel_size,
stride=self.config.stem_conv_stride,
padding=self.config.stem_conv_kernel_size // 2)
self.bn1 = nn.InstanceNorm2d(64) if self.config.use_instance_norm else nn.BatchNorm2d(64)
self.conv2 = ConvBlock(64, 128, self.config.use_instance_norm)
self.conv3 = ConvBlock(128, 128, self.config.use_instance_norm)
self.conv4 = ConvBlock(128, self.config.hg_num_features, self.config.use_instance_norm)
# Hourglasses
for hg_module in range(self.config.num_modules):
self.add_module('m' + str(hg_module), HourGlass(self.config))
self.add_module('top_m_' + str(hg_module), ConvBlock(self.config.hg_num_features,
self.config.hg_num_features,
self.config.use_instance_norm))
self.add_module('conv_last' + str(hg_module), nn.Conv2d(self.config.hg_num_features,
self.config.hg_num_features,
kernel_size=1, stride=1, padding=0))
self.add_module('bn_end' + str(hg_module),
nn.InstanceNorm2d(self.config.hg_num_features) if self.config.use_instance_norm
else nn.BatchNorm2d(self.config.hg_num_features))
self.add_module('l' + str(hg_module), nn.Conv2d(self.config.hg_num_features,
self.config.num_landmarks,
kernel_size=1, stride=1, padding=0))
if hg_module < self.config.num_modules - 1:
self.add_module('bl' + str(hg_module), nn.Conv2d(self.config.hg_num_features,
self.config.hg_num_features,
kernel_size=1, stride=1, padding=0))
self.add_module('al' + str(hg_module), nn.Conv2d(self.config.num_landmarks,
self.config.hg_num_features,
kernel_size=1, stride=1, padding=0))
def forward(self, x):
x = self.conv2(F.relu(self.bn1(self.conv1(x)), True))
if self.config.stem_pool_kernel_size > 1:
if self.config.use_avg_pool:
x = F.avg_pool2d(x, self.config.stem_pool_kernel_size)
else:
x = F.max_pool2d(x, self.config.stem_pool_kernel_size)
x = self.conv3(x)
x = self.conv4(x)
previous = x
hg_feats = []
tmp_out = None
for i in range(self.config.num_modules):
hg = self._modules['m' + str(i)](previous)
ll = hg
ll = self._modules['top_m_' + str(i)](ll)
ll = F.relu(self._modules['bn_end' + str(i)](self._modules['conv_last' + str(i)](ll)), True)
# Predict heatmaps
tmp_out = self._modules['l' + str(i)](ll)
if i < self.config.num_modules - 1:
ll = self._modules['bl' + str(i)](ll)
tmp_out_ = self._modules['al' + str(i)](tmp_out)
previous = previous + ll + tmp_out_
hg_feats.append(ll)
return tmp_out, x, tuple(hg_feats)
@@ -0,0 +1,168 @@
import os
import cv2
import torch
import numpy as np
from types import SimpleNamespace
from typing import Union, Optional, Tuple
from .fan import FAN
__all__ = ['FANPredictor']
class FANPredictor(object):
def __init__(self, device: Union[str, torch.device] = 'cuda:0', model: Optional[SimpleNamespace] = None,
config: Optional[SimpleNamespace] = None) -> None:
self.device = device
if model is None:
model = FANPredictor.get_model()
if config is None:
config = FANPredictor.create_config()
self.config = SimpleNamespace(**model.config.__dict__, **config.__dict__)
self.net = FAN(config=self.config).to(self.device)
self.net.load_state_dict(torch.load(model.weights, map_location=self.device))
self.net.eval()
if self.config.use_jit:
self.net = torch.jit.trace(self.net, torch.rand(1, 3, self.config.input_size,
self.config.input_size).to(self.device))
@staticmethod
def get_model(name: str = '2dfan2') -> SimpleNamespace:
from motiondiff_modules import CKPT_DIR_PATH, download_models
name = name.lower()
FACE_ALIGNMENT_PREFIX = "https://github.com/hhj1897/face_alignment/raw/9cf5494e443f26d567972f3f50f6212d65b76c01/ibug/face_alignment/fan/weights/"
download_models({f'{name}.pth': FACE_ALIGNMENT_PREFIX + f'{name}.pth'})
if name == '2dfan2':
return SimpleNamespace(weights=os.path.join(str(CKPT_DIR_PATH), '2dfan2.pth'),
config=SimpleNamespace(crop_ratio=0.55, input_size=256, num_modules=2,
hg_num_features=256, hg_depth=4, use_avg_pool=False,
use_instance_norm=False, stem_conv_kernel_size=7,
stem_conv_stride=2, stem_pool_kernel_size=2,
num_landmarks=68))
elif name == '2dfan4':
return SimpleNamespace(weights=os.path.join(str(CKPT_DIR_PATH), '2dfan4.pth'),
config=SimpleNamespace(crop_ratio=0.55, input_size=256, num_modules=4,
hg_num_features=256, hg_depth=4, use_avg_pool=True,
use_instance_norm=False, stem_conv_kernel_size=7,
stem_conv_stride=2, stem_pool_kernel_size=2,
num_landmarks=68))
elif name == '2dfan2_alt':
return SimpleNamespace(weights=os.path.join(str(CKPT_DIR_PATH), '2dfan2_alt.pth'),
config=SimpleNamespace(crop_ratio=0.55, input_size=256, num_modules=2,
hg_num_features=256, hg_depth=4, use_avg_pool=False,
use_instance_norm=False, stem_conv_kernel_size=7,
stem_conv_stride=2, stem_pool_kernel_size=2,
num_landmarks=68))
else:
raise ValueError('name must be set to either 2dfan2, 2dfan4, or 2dfan2_alt')
@staticmethod
def create_config(gamma: float = 1.0, radius: float = 0.1, use_jit: bool = True) -> SimpleNamespace:
return SimpleNamespace(gamma=gamma, radius=radius, use_jit=use_jit)
@torch.no_grad()
def __call__(self, image: np.ndarray, face_boxes: np.ndarray, rgb: bool = True,
return_features: bool = False) -> Union[Tuple[np.ndarray, np.ndarray],
Tuple[np.ndarray, np.ndarray, torch.Tensor]]:
if face_boxes.size > 0:
if not rgb:
image = image[..., ::-1]
if face_boxes.ndim == 1:
face_boxes = face_boxes[np.newaxis, ...]
# Crop the faces
face_patches = []
centres = (face_boxes[:, [0, 1]] + face_boxes[:, [2, 3]]) / 2.0
face_sizes = (face_boxes[:, [3, 2]] - face_boxes[:, [1, 0]]).mean(axis=1)
enlarged_face_box_sizes = (face_sizes / self.config.crop_ratio)[:, np.newaxis].repeat(2, axis=1)
enlarged_face_boxes = np.zeros_like(face_boxes[:, :4])
enlarged_face_boxes[:, :2] = np.round(centres - enlarged_face_box_sizes / 2.0)
enlarged_face_boxes[:, 2:] = np.round(enlarged_face_boxes[:, :2] + enlarged_face_box_sizes) + 1
enlarged_face_boxes = enlarged_face_boxes.astype(int)
outer_bounding_box = np.hstack((enlarged_face_boxes[:, :2].min(axis=0),
enlarged_face_boxes[:, 2:].max(axis=0)))
pad_widths = np.zeros(shape=(3, 2), dtype=int)
if outer_bounding_box[0] < 0:
pad_widths[1][0] = -outer_bounding_box[0]
if outer_bounding_box[1] < 0:
pad_widths[0][0] = -outer_bounding_box[1]
if outer_bounding_box[2] > image.shape[1]:
pad_widths[1][1] = outer_bounding_box[2] - image.shape[1]
if outer_bounding_box[3] > image.shape[0]:
pad_widths[0][1] = outer_bounding_box[3] - image.shape[0]
if np.any(pad_widths > 0):
image = np.pad(image, pad_widths)
for left, top, right, bottom in enlarged_face_boxes:
left += pad_widths[1][0]
top += pad_widths[0][0]
right += pad_widths[1][0]
bottom += pad_widths[0][0]
face_patches.append(cv2.resize(image[top: bottom, left: right, :],
(self.config.input_size, self.config.input_size)))
face_patches = torch.from_numpy(np.array(face_patches).transpose(
(0, 3, 1, 2)).astype(np.float32)).to(self.device) / 255.0
# Get heatmaps
heatmaps, stem_feats, hg_feats = self.net(face_patches)
# Get landmark coordinates and scores
landmarks, landmark_scores = self._decode(heatmaps)
# Rectify landmark coordinates
hh, hw = heatmaps.size(2), heatmaps.size(3)
for landmark, (left, top, right, bottom) in zip(landmarks, enlarged_face_boxes):
landmark[:, 0] = landmark[:, 0] * (right - left) / hw + left
landmark[:, 1] = landmark[:, 1] * (bottom - top) / hh + top
if return_features:
return landmarks, landmark_scores, torch.cat((stem_feats, torch.cat(hg_feats, dim=1) *
torch.sum(heatmaps, dim=1, keepdim=True)), dim=1)
else:
return landmarks, landmark_scores
else:
landmarks = np.empty(shape=(0, 68, 2), dtype=np.float32)
landmark_scores = np.empty(shape=(0, 68), dtype=np.float32)
if return_features:
return landmarks, landmark_scores, torch.Tensor([])
else:
return landmarks, landmark_scores
def _decode(self, heatmaps: torch.Tensor) -> Tuple[np.ndarray, np.ndarray]:
heatmaps = heatmaps.contiguous()
scores = heatmaps.max(dim=3)[0].max(dim=2)[0]
if (self.config.radius ** 2 * heatmaps.shape[2] * heatmaps.shape[3] <
heatmaps.shape[2] ** 2 + heatmaps.shape[3] ** 2):
# Find peaks in all heatmaps
m = heatmaps.view(heatmaps.shape[0] * heatmaps.shape[1], -1).argmax(1)
all_peaks = torch.cat(
[(m / heatmaps.shape[3]).trunc().view(-1, 1), (m % heatmaps.shape[3]).view(-1, 1)], dim=1
).reshape((heatmaps.shape[0], heatmaps.shape[1], 1, 1, 2)).repeat(
1, 1, heatmaps.shape[2], heatmaps.shape[3], 1).float()
# Apply masks created from the peaks
all_indices = torch.zeros_like(all_peaks) + torch.stack(
[torch.arange(0.0, all_peaks.shape[2],
device=all_peaks.device).unsqueeze(-1).repeat(1, all_peaks.shape[3]),
torch.arange(0.0, all_peaks.shape[3],
device=all_peaks.device).unsqueeze(0).repeat(all_peaks.shape[2], 1)], dim=-1)
heatmaps = heatmaps * ((all_indices - all_peaks).norm(dim=-1) <= self.config.radius *
(heatmaps.shape[2] * heatmaps.shape[3]) ** 0.5).float()
# Prepare the indices for calculating centroids
x_indices = (torch.zeros((*heatmaps.shape[:2], heatmaps.shape[3]), device=heatmaps.device) +
torch.arange(0.5, heatmaps.shape[3], device=heatmaps.device))
y_indices = (torch.zeros(heatmaps.shape[:3], device=heatmaps.device) +
torch.arange(0.5, heatmaps.shape[2], device=heatmaps.device))
# Finally, find centroids as landmark locations
heatmaps = heatmaps.clamp_min(0.0)
if self.config.gamma != 1.0:
heatmaps = heatmaps.pow(self.config.gamma)
m00s = heatmaps.sum(dim=(2, 3)).clamp_min(torch.finfo(heatmaps.dtype).eps)
xs = heatmaps.sum(dim=2).mul(x_indices).sum(dim=2).div(m00s)
ys = heatmaps.sum(dim=3).mul(y_indices).sum(dim=2).div(m00s)
lm_info = torch.stack((xs, ys, scores), dim=-1).cpu().numpy()
return lm_info[..., :-1], lm_info[..., -1]
@@ -0,0 +1,51 @@
import cv2
import numpy as np
from typing import Optional, Sequence, Tuple
__all__ = ['get_landmark_connectivity', 'plot_landmarks']
def get_landmark_connectivity(num_landmarks: int) -> Optional[Sequence[Tuple[int, int]]]:
if num_landmarks == 68:
return ((0, 1), (1, 2), (2, 3), (3, 4), (4, 5), (5, 6), (6, 7), (7, 8), (8, 9), (9, 10), (10, 11), (11, 12),
(12, 13), (13, 14), (14, 15), (15, 16), (17, 18), (18, 19), (19, 20), (20, 21), (22, 23), (23, 24),
(24, 25), (25, 26), (27, 28), (28, 29), (29, 30), (30, 33), (31, 32), (32, 33), (33, 34), (34, 35),
(36, 37), (37, 38), (38, 39), (40, 41), (41, 36), (42, 43), (43, 44), (44, 45), (45, 46), (46, 47),
(47, 42), (48, 49), (49, 50), (50, 51), (51, 52), (52, 53), (53, 54), (54, 55), (55, 56), (56, 57),
(57, 58), (58, 59), (59, 48), (60, 61), (61, 62), (62, 63), (63, 64), (64, 65), (65, 66), (66, 67),
(67, 60), (39, 40))
elif num_landmarks == 100:
return ((0, 1), (1, 2), (2, 3), (3, 4), (4, 5), (5, 6), (6, 7), (7, 8), (8, 9), (9, 10), (10, 11), (11, 12),
(12, 13), (13, 14), (14, 15), (15, 16), (17, 18), (18, 19), (19, 20), (20, 21), (22, 23), (23, 24),
(24, 25), (25, 26), (68, 69), (69, 70), (70, 71), (72, 73), (73, 74), (74, 75), (36, 76), (76, 37),
(37, 77), (77, 38), (38, 78), (78, 39), (39, 40), (40, 79), (79, 41), (41, 36), (42, 80), (80, 43),
(43, 81), (81, 44), (44, 82), (82, 45), (45, 46), (46, 83), (83, 47), (47, 42), (27, 28), (28, 29),
(29, 30), (30, 33), (31, 32), (32, 33), (33, 34), (34, 35), (84, 85), (86, 87), (48, 49), (49, 88),
(88, 50), (50, 51), (51, 52), (52, 89), (89, 53), (53, 54), (54, 55), (55, 90), (90, 56), (56, 57),
(57, 58), (58, 91), (91, 59), (59, 48), (60, 92), (92, 93), (93, 61), (61, 62), (62, 63), (63, 94),
(94, 95), (95, 64), (64, 96), (96, 97), (97, 65), (65, 66), (66, 67), (67, 98), (98, 99), (99, 60),
(17, 68), (21, 71), (22, 72), (26, 75))
else:
return None
def plot_landmarks(image: np.ndarray, landmarks: np.ndarray, landmark_scores: Optional[Sequence[float]] = None,
threshold: float = 0.2, line_colour: Tuple[int, int, int] = (0, 255, 0),
pts_colour: Tuple[int, int, int] = (0, 0, 255), line_thickness: int = 1, pts_radius: int = 1,
landmark_connectivity: Optional[Sequence[Tuple[int, int]]] = None) -> None:
num_landmarks = len(landmarks)
if landmark_scores is None:
landmark_scores = np.full((num_landmarks,), threshold + 1.0, dtype=float)
if landmark_connectivity is None:
landmark_connectivity = get_landmark_connectivity(len(landmarks))
if landmark_connectivity is not None:
for (idx1, idx2) in landmark_connectivity:
if (idx1 < num_landmarks and idx2 < num_landmarks and
landmark_scores[idx1] >= threshold and landmark_scores[idx2] >= threshold):
cv2.line(image, tuple(landmarks[idx1].astype(int).tolist()),
tuple(landmarks[idx2].astype(int).tolist()),
color=line_colour, thickness=line_thickness, lineType=cv2.LINE_AA)
for landmark, score in zip(landmarks, landmark_scores):
if score >= threshold:
cv2.circle(image, tuple(landmark.astype(int).tolist()), pts_radius, pts_colour, -1)
@@ -0,0 +1,5 @@
from .s3fd import S3FDPredictor
from .retina_face import RetinaFacePredictor
__version__ = '0.1.0'
@@ -0,0 +1 @@
from .retina_face_predictor import RetinaFacePredictor
@@ -0,0 +1,332 @@
import torch
import numpy as np
def point_form(boxes):
""" Convert prior_boxes to (xmin, ymin, xmax, ymax)
representation for comparison to point form ground truth data.
Args:
boxes: (tensor) center-size default boxes from priorbox layers.
Return:
boxes: (tensor) Converted xmin, ymin, xmax, ymax form of boxes.
"""
return torch.cat((boxes[:, :2] - boxes[:, 2:]/2, # xmin, ymin
boxes[:, :2] + boxes[:, 2:]/2), 1) # xmax, ymax
def center_size(boxes):
""" Convert prior_boxes to (cx, cy, w, h)
representation for comparison to center-size form ground truth data.
Args:
boxes: (tensor) point_form boxes
Return:
boxes: (tensor) Converted xmin, ymin, xmax, ymax form of boxes.
"""
return torch.cat((boxes[:, 2:] + boxes[:, :2])/2, # cx, cy
boxes[:, 2:] - boxes[:, :2], 1) # w, h
def intersect(box_a, box_b):
""" We resize both tensors to [A,B,2] without new malloc:
[A,2] -> [A,1,2] -> [A,B,2]
[B,2] -> [1,B,2] -> [A,B,2]
Then we compute the area of intersect between box_a and box_b.
Args:
box_a: (tensor) bounding boxes, Shape: [A,4].
box_b: (tensor) bounding boxes, Shape: [B,4].
Return:
(tensor) intersection area, Shape: [A,B].
"""
A = box_a.size(0)
B = box_b.size(0)
max_xy = torch.min(box_a[:, 2:].unsqueeze(1).expand(A, B, 2),
box_b[:, 2:].unsqueeze(0).expand(A, B, 2))
min_xy = torch.max(box_a[:, :2].unsqueeze(1).expand(A, B, 2),
box_b[:, :2].unsqueeze(0).expand(A, B, 2))
inter = torch.clamp((max_xy - min_xy), min=0)
return inter[:, :, 0] * inter[:, :, 1]
def jaccard(box_a, box_b):
"""Compute the jaccard overlap of two sets of boxes. The jaccard overlap
is simply the intersection over union of two boxes. Here we operate on
ground truth boxes and default boxes.
E.g.:
A ∩ B / A ∪ B = A ∩ B / (area(A) + area(B) - A ∩ B)
Args:
box_a: (tensor) Ground truth bounding boxes, Shape: [num_objects,4]
box_b: (tensor) Prior boxes from priorbox layers, Shape: [num_priors,4]
Return:
jaccard overlap: (tensor) Shape: [box_a.size(0), box_b.size(0)]
"""
inter = intersect(box_a, box_b)
area_a = ((box_a[:, 2]-box_a[:, 0]) *
(box_a[:, 3]-box_a[:, 1])).unsqueeze(1).expand_as(inter) # [A,B]
area_b = ((box_b[:, 2]-box_b[:, 0]) *
(box_b[:, 3]-box_b[:, 1])).unsqueeze(0).expand_as(inter) # [A,B]
union = area_a + area_b - inter
return inter / union # [A,B]
def matrix_iou(a, b):
"""
return iou of a and b, numpy version for data augenmentation
"""
lt = np.maximum(a[:, np.newaxis, :2], b[:, :2])
rb = np.minimum(a[:, np.newaxis, 2:], b[:, 2:])
area_i = np.prod(rb - lt, axis=2) * (lt < rb).all(axis=2)
area_a = np.prod(a[:, 2:] - a[:, :2], axis=1)
area_b = np.prod(b[:, 2:] - b[:, :2], axis=1)
return area_i / (area_a[:, np.newaxis] + area_b - area_i)
def matrix_iof(a, b):
"""
return iof of a and b, numpy version for data augenmentation
"""
lt = np.maximum(a[:, np.newaxis, :2], b[:, :2])
rb = np.minimum(a[:, np.newaxis, 2:], b[:, 2:])
area_i = np.prod(rb - lt, axis=2) * (lt < rb).all(axis=2)
area_a = np.prod(a[:, 2:] - a[:, :2], axis=1)
return area_i / np.maximum(area_a[:, np.newaxis], 1)
def match(threshold, truths, priors, variances, labels, landms, loc_t, conf_t, landm_t, idx):
"""Match each prior box with the ground truth box of the highest jaccard
overlap, encode the bounding boxes, then return the matched indices
corresponding to both confidence and location preds.
Args:
threshold: (float) The overlap threshold used when mathing boxes.
truths: (tensor) Ground truth boxes, Shape: [num_obj, 4].
priors: (tensor) Prior boxes from priorbox layers, Shape: [n_priors,4].
variances: (tensor) Variances corresponding to each prior coord,
Shape: [num_priors, 4].
labels: (tensor) All the class labels for the image, Shape: [num_obj].
landms: (tensor) Ground truth landms, Shape [num_obj, 10].
loc_t: (tensor) Tensor to be filled w/ endcoded location targets.
conf_t: (tensor) Tensor to be filled w/ matched indices for conf preds.
landm_t: (tensor) Tensor to be filled w/ endcoded landm targets.
idx: (int) current batch index
Return:
The matched indices corresponding to 1)location 2)confidence 3)landm preds.
"""
# jaccard index
overlaps = jaccard(
truths,
point_form(priors)
)
# (Bipartite Matching)
# [1,num_objects] best prior for each ground truth
best_prior_overlap, best_prior_idx = overlaps.max(1, keepdim=True)
# ignore hard gt
valid_gt_idx = best_prior_overlap[:, 0] >= 0.2
best_prior_idx_filter = best_prior_idx[valid_gt_idx, :]
if best_prior_idx_filter.shape[0] <= 0:
loc_t[idx] = 0
conf_t[idx] = 0
return
# [1,num_priors] best ground truth for each prior
best_truth_overlap, best_truth_idx = overlaps.max(0, keepdim=True)
best_truth_idx.squeeze_(0)
best_truth_overlap.squeeze_(0)
best_prior_idx.squeeze_(1)
best_prior_idx_filter.squeeze_(1)
best_prior_overlap.squeeze_(1)
best_truth_overlap.index_fill_(0, best_prior_idx_filter, 2) # ensure best prior
# TODO refactor: index best_prior_idx with long tensor
# ensure every gt matches with its prior of max overlap
for j in range(best_prior_idx.size(0)): # 判别此anchor是预测哪一个boxes
best_truth_idx[best_prior_idx[j]] = j
matches = truths[best_truth_idx] # Shape: [num_priors,4] 此处为每一个anchor对应的bbox取出来
conf = labels[best_truth_idx] # Shape: [num_priors] 此处为每一个anchor对应的label取出来
conf[best_truth_overlap < threshold] = 0 # label as background overlap<0.35的全部作为负样本
loc = encode(matches, priors, variances)
matches_landm = landms[best_truth_idx]
landm = encode_landm(matches_landm, priors, variances)
loc_t[idx] = loc # [num_priors,4] encoded offsets to learn
conf_t[idx] = conf # [num_priors] top class label for each prior
landm_t[idx] = landm
def encode(matched, priors, variances):
"""Encode the variances from the priorbox layers into the ground truth boxes
we have matched (based on jaccard overlap) with the prior boxes.
Args:
matched: (tensor) Coords of ground truth for each prior in point-form
Shape: [num_priors, 4].
priors: (tensor) Prior boxes in center-offset form
Shape: [num_priors,4].
variances: (list[float]) Variances of priorboxes
Return:
encoded boxes (tensor), Shape: [num_priors, 4]
"""
# dist b/t match center and prior's center
g_cxcy = (matched[:, :2] + matched[:, 2:])/2 - priors[:, :2]
# encode variance
g_cxcy /= (variances[0] * priors[:, 2:])
# match wh / prior wh
g_wh = (matched[:, 2:] - matched[:, :2]) / priors[:, 2:]
g_wh = torch.log(g_wh) / variances[1]
# return target for smooth_l1_loss
return torch.cat([g_cxcy, g_wh], 1) # [num_priors,4]
def encode_landm(matched, priors, variances):
"""Encode the variances from the priorbox layers into the ground truth boxes
we have matched (based on jaccard overlap) with the prior boxes.
Args:
matched: (tensor) Coords of ground truth for each prior in point-form
Shape: [num_priors, 10].
priors: (tensor) Prior boxes in center-offset form
Shape: [num_priors,4].
variances: (list[float]) Variances of priorboxes
Return:
encoded landm (tensor), Shape: [num_priors, 10]
"""
# dist b/t match center and prior's center
matched = torch.reshape(matched, (matched.size(0), 5, 2))
priors_cx = priors[:, 0].unsqueeze(1).expand(matched.size(0), 5).unsqueeze(2)
priors_cy = priors[:, 1].unsqueeze(1).expand(matched.size(0), 5).unsqueeze(2)
priors_w = priors[:, 2].unsqueeze(1).expand(matched.size(0), 5).unsqueeze(2)
priors_h = priors[:, 3].unsqueeze(1).expand(matched.size(0), 5).unsqueeze(2)
priors = torch.cat([priors_cx, priors_cy, priors_w, priors_h], dim=2)
g_cxcy = matched[:, :, :2] - priors[:, :, :2]
# encode variance
g_cxcy /= (variances[0] * priors[:, :, 2:])
# g_cxcy /= priors[:, :, 2:]
g_cxcy = g_cxcy.reshape(g_cxcy.size(0), -1)
# return target for smooth_l1_loss
return g_cxcy
# Adapted from https://github.com/Hakuyume/chainer-ssd
def decode(loc, priors, variances):
"""Decode locations from predictions using priors to undo
the encoding we did for offset regression at train time.
Args:
loc (tensor): location predictions for loc layers,
Shape: [num_priors,4]
priors (tensor): Prior boxes in center-offset form.
Shape: [num_priors,4].
variances: (list[float]) Variances of priorboxes
Return:
decoded bounding box predictions
"""
boxes = torch.cat((
priors[:, :2] + loc[:, :2] * variances[0] * priors[:, 2:],
priors[:, 2:] * torch.exp(loc[:, 2:] * variances[1])), 1)
boxes[:, :2] -= boxes[:, 2:] / 2
boxes[:, 2:] += boxes[:, :2]
return boxes
def decode_landm(pre, priors, variances):
"""Decode landm from predictions using priors to undo
the encoding we did for offset regression at train time.
Args:
pre (tensor): landm predictions for loc layers,
Shape: [num_priors,10]
priors (tensor): Prior boxes in center-offset form.
Shape: [num_priors,4].
variances: (list[float]) Variances of priorboxes
Return:
decoded landm predictions
"""
landms = torch.cat((priors[:, :2] + pre[:, :2] * variances[0] * priors[:, 2:],
priors[:, :2] + pre[:, 2:4] * variances[0] * priors[:, 2:],
priors[:, :2] + pre[:, 4:6] * variances[0] * priors[:, 2:],
priors[:, :2] + pre[:, 6:8] * variances[0] * priors[:, 2:],
priors[:, :2] + pre[:, 8:10] * variances[0] * priors[:, 2:],
), dim=1)
return landms
def log_sum_exp(x):
"""Utility function for computing log_sum_exp while determining
This will be used to determine unaveraged confidence loss across
all examples in a batch.
Args:
x (Variable(tensor)): conf_preds from conf layers
"""
x_max = x.data.max()
return torch.log(torch.sum(torch.exp(x-x_max), 1, keepdim=True)) + x_max
# Original author: Francisco Massa:
# https://github.com/fmassa/object-detection.torch
# Ported to PyTorch by Max deGroot (02/01/2017)
def nms(boxes, scores, overlap=0.5, top_k=200):
"""Apply non-maximum suppression at test time to avoid detecting too many
overlapping bounding boxes for a given object.
Args:
boxes: (tensor) The location preds for the img, Shape: [num_priors,4].
scores: (tensor) The class predscores for the img, Shape:[num_priors].
overlap: (float) The overlap thresh for suppressing unnecessary boxes.
top_k: (int) The Maximum number of box preds to consider.
Return:
The indices of the kept boxes with respect to num_priors.
"""
keep = torch.Tensor(scores.size(0)).fill_(0).long()
if boxes.numel() == 0:
return keep
x1 = boxes[:, 0]
y1 = boxes[:, 1]
x2 = boxes[:, 2]
y2 = boxes[:, 3]
area = torch.mul(x2 - x1, y2 - y1)
v, idx = scores.sort(0) # sort in ascending order
# I = I[v >= 0.01]
idx = idx[-top_k:] # indices of the top-k largest vals
xx1 = boxes.new()
yy1 = boxes.new()
xx2 = boxes.new()
yy2 = boxes.new()
w = boxes.new()
h = boxes.new()
# keep = torch.Tensor()
count = 0
while idx.numel() > 0:
i = idx[-1] # index of current largest val
# keep.append(i)
keep[count] = i
count += 1
if idx.size(0) == 1:
break
idx = idx[:-1] # remove kept element from view
# load bboxes of next highest vals
torch.index_select(x1, 0, idx, out=xx1)
torch.index_select(y1, 0, idx, out=yy1)
torch.index_select(x2, 0, idx, out=xx2)
torch.index_select(y2, 0, idx, out=yy2)
# store element-wise max with next highest score
xx1 = torch.clamp(xx1, min=x1[i])
yy1 = torch.clamp(yy1, min=y1[i])
xx2 = torch.clamp(xx2, max=x2[i])
yy2 = torch.clamp(yy2, max=y2[i])
w.resize_as_(xx2)
h.resize_as_(yy2)
w = xx2 - xx1
h = yy2 - yy1
# check sizes of xx1 and xx2.. after each iteration
w = torch.clamp(w, min=0.0)
h = torch.clamp(h, min=0.0)
inter = w*h
# IoU = i / (area(a) + area(b) - i)
rem_areas = torch.index_select(area, 0, idx) # load remaining areas)
union = (rem_areas - inter) + area[i]
IoU = inter/union # store result in iou
# keep only elements with an IoU <= overlap
idx = idx[IoU.le(overlap)]
return keep, count
@@ -0,0 +1,41 @@
# config.py
cfg_mnet = {
'name': 'mobilenet0.25',
'min_sizes': [[16, 32], [64, 128], [256, 512]],
'steps': [8, 16, 32],
'variance': [0.1, 0.2],
'clip': False,
'loc_weight': 2.0,
'gpu_train': True,
'batch_size': 32,
'ngpu': 1,
'epoch': 250,
'decay1': 190,
'decay2': 220,
'image_size': 640,
'pretrain': False,
'return_layers': {'stage1': 1, 'stage2': 2, 'stage3': 3},
'in_channel': 32,
'out_channel': 64
}
cfg_re50 = {
'name': 'Resnet50',
'min_sizes': [[16, 32], [64, 128], [256, 512]],
'steps': [8, 16, 32],
'variance': [0.1, 0.2],
'clip': False,
'loc_weight': 2.0,
'gpu_train': True,
'batch_size': 24,
'ngpu': 4,
'epoch': 100,
'decay1': 70,
'decay2': 90,
'image_size': 840,
'pretrain': False,
'return_layers': {'layer2': 1, 'layer3': 2, 'layer4': 3},
'in_channel': 256,
'out_channel': 256
}
@@ -0,0 +1,33 @@
import torch
from itertools import product as product
from math import ceil
class PriorBox(object):
def __init__(self, cfg, image_size=None):
super(PriorBox, self).__init__()
self.min_sizes = cfg['min_sizes']
self.steps = cfg['steps']
self.clip = cfg['clip']
self.image_size = image_size
self.feature_maps = [[ceil(self.image_size[0]/step), ceil(self.image_size[1]/step)] for step in self.steps]
self.name = "s"
def forward(self):
anchors = []
for k, f in enumerate(self.feature_maps):
min_sizes = self.min_sizes[k]
for i, j in product(range(f[0]), range(f[1])):
for min_size in min_sizes:
s_kx = min_size / self.image_size[1]
s_ky = min_size / self.image_size[0]
dense_cx = [x * self.steps[k] / self.image_size[1] for x in [j + 0.5]]
dense_cy = [y * self.steps[k] / self.image_size[0] for y in [i + 0.5]]
for cy, cx in product(dense_cy, dense_cx):
anchors += [cx, cy, s_kx, s_ky]
# back to torch land
output = torch.Tensor(anchors).view(-1, 4)
if self.clip:
output.clamp_(max=1, min=0)
return output
@@ -0,0 +1,39 @@
# --------------------------------------------------------
# Fast R-CNN
# Copyright (c) 2015 Microsoft
# Licensed under The MIT License [see LICENSE for details]
# Written by Ross Girshick
# --------------------------------------------------------
import numpy as np
def py_cpu_nms(dets, thresh, top_k):
"""Pure Python NMS baseline."""
x1 = dets[:, 0]
y1 = dets[:, 1]
x2 = dets[:, 2]
y2 = dets[:, 3]
scores = dets[:, 4]
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
order = scores.argsort()[: -top_k - 1: -1]
keep = []
while order.size > 0:
i = order[0]
keep.append(i)
xx1 = np.maximum(x1[i], x1[order[1:]])
yy1 = np.maximum(y1[i], y1[order[1:]])
xx2 = np.minimum(x2[i], x2[order[1:]])
yy2 = np.minimum(y2[i], y2[order[1:]])
w = np.maximum(0.0, xx2 - xx1 + 1)
h = np.maximum(0.0, yy2 - yy1 + 1)
inter = w * h
ovr = inter / (areas[i] + areas[order[1:]] - inter)
inds = np.where(ovr <= thresh)[0]
order = order[inds + 1]
return keep
@@ -0,0 +1,117 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.models as models
import torchvision.models._utils as _utils
from .retina_face_net import MobileNetV1, FPN, SSH
class ClassHead(nn.Module):
def __init__(self, inchannels=512, num_anchors=3):
super(ClassHead, self).__init__()
self.num_anchors = num_anchors
self.conv1x1 = nn.Conv2d(inchannels, self.num_anchors*2, kernel_size=(1, 1), stride=1, padding=0)
def forward(self, x):
out = self.conv1x1(x)
out = out.permute(0, 2, 3, 1).contiguous()
return out.view(out.shape[0], -1, 2)
class BboxHead(nn.Module):
def __init__(self, inchannels=512, num_anchors=3):
super(BboxHead, self).__init__()
self.conv1x1 = nn.Conv2d(inchannels, num_anchors*4, kernel_size=(1, 1), stride=1,padding=0)
def forward(self, x):
out = self.conv1x1(x)
out = out.permute(0, 2, 3, 1).contiguous()
return out.view(out.shape[0], -1, 4)
class LandmarkHead(nn.Module):
def __init__(self, inchannels=512, num_anchors=3):
super(LandmarkHead, self).__init__()
self.conv1x1 = nn.Conv2d(inchannels,num_anchors*10, kernel_size=(1, 1), stride=1, padding=0)
def forward(self, x):
out = self.conv1x1(x)
out = out.permute(0, 2, 3, 1).contiguous()
return out.view(out.shape[0], -1, 10)
class RetinaFace(nn.Module):
def __init__(self, cfg=None, phase='train'):
"""
:param cfg: Network related settings.
:param phase: train or test.
"""
super(RetinaFace, self).__init__()
self.phase = phase
backbone = None
if cfg['name'] == 'mobilenet0.25':
backbone = MobileNetV1()
if cfg['pretrain']:
raise ValueError('cfg[\'pretrain\'] cannot be set to True for mobilenet0.25')
elif cfg['name'] == 'Resnet50':
backbone = models.resnet50(pretrained=cfg['pretrain'])
self.body = _utils.IntermediateLayerGetter(backbone, cfg['return_layers'])
in_channels_stage2 = cfg['in_channel']
in_channels_list = [
in_channels_stage2 * 2,
in_channels_stage2 * 4,
in_channels_stage2 * 8,
]
out_channels = cfg['out_channel']
self.fpn = FPN(in_channels_list,out_channels)
self.ssh1 = SSH(out_channels, out_channels)
self.ssh2 = SSH(out_channels, out_channels)
self.ssh3 = SSH(out_channels, out_channels)
self.ClassHead = self._make_class_head(fpn_num=3, inchannels=cfg['out_channel'])
self.BboxHead = self._make_bbox_head(fpn_num=3, inchannels=cfg['out_channel'])
self.LandmarkHead = self._make_landmark_head(fpn_num=3, inchannels=cfg['out_channel'])
def _make_class_head(self, fpn_num=3, inchannels=64, anchor_num=2):
classhead = nn.ModuleList()
for i in range(fpn_num):
classhead.append(ClassHead(inchannels, anchor_num))
return classhead
def _make_bbox_head(self, fpn_num=3, inchannels=64, anchor_num=2):
bboxhead = nn.ModuleList()
for i in range(fpn_num):
bboxhead.append(BboxHead(inchannels, anchor_num))
return bboxhead
def _make_landmark_head(self, fpn_num=3, inchannels=64, anchor_num=2):
landmarkhead = nn.ModuleList()
for i in range(fpn_num):
landmarkhead.append(LandmarkHead(inchannels, anchor_num))
return landmarkhead
def forward(self, inputs):
out = self.body(inputs)
# FPN
fpn = self.fpn(out)
# SSH
feature1 = self.ssh1(fpn[0])
feature2 = self.ssh2(fpn[1])
feature3 = self.ssh3(fpn[2])
features = [feature1, feature2, feature3]
bbox_regressions = torch.cat([self.BboxHead[i](feature) for i, feature in enumerate(features)], dim=1)
classifications = torch.cat([self.ClassHead[i](feature) for i, feature in enumerate(features)], dim=1)
ldm_regressions = torch.cat([self.LandmarkHead[i](feature) for i, feature in enumerate(features)], dim=1)
if self.phase == 'train':
output = (bbox_regressions, classifications, ldm_regressions)
else:
output = (bbox_regressions, F.softmax(classifications, dim=-1), ldm_regressions)
return output
@@ -0,0 +1,137 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
def conv_bn(inp, oup, stride = 1, leaky = 0):
return nn.Sequential(
nn.Conv2d(inp, oup, 3, stride, 1, bias=False),
nn.BatchNorm2d(oup),
nn.LeakyReLU(negative_slope=leaky, inplace=True)
)
def conv_bn_no_relu(inp, oup, stride):
return nn.Sequential(
nn.Conv2d(inp, oup, 3, stride, 1, bias=False),
nn.BatchNorm2d(oup),
)
def conv_bn1X1(inp, oup, stride, leaky=0):
return nn.Sequential(
nn.Conv2d(inp, oup, 1, stride, padding=0, bias=False),
nn.BatchNorm2d(oup),
nn.LeakyReLU(negative_slope=leaky, inplace=True)
)
def conv_dw(inp, oup, stride, leaky=0.1):
return nn.Sequential(
nn.Conv2d(inp, inp, 3, stride, 1, groups=inp, bias=False),
nn.BatchNorm2d(inp),
nn.LeakyReLU(negative_slope=leaky, inplace=True),
nn.Conv2d(inp, oup, 1, 1, 0, bias=False),
nn.BatchNorm2d(oup),
nn.LeakyReLU(negative_slope=leaky, inplace=True),
)
class SSH(nn.Module):
def __init__(self, in_channel, out_channel):
super(SSH, self).__init__()
assert out_channel % 4 == 0
leaky = 0
if out_channel <= 64:
leaky = 0.1
self.conv3X3 = conv_bn_no_relu(in_channel, out_channel//2, stride=1)
self.conv5X5_1 = conv_bn(in_channel, out_channel//4, stride=1, leaky = leaky)
self.conv5X5_2 = conv_bn_no_relu(out_channel//4, out_channel//4, stride=1)
self.conv7X7_2 = conv_bn(out_channel//4, out_channel//4, stride=1, leaky = leaky)
self.conv7x7_3 = conv_bn_no_relu(out_channel//4, out_channel//4, stride=1)
def forward(self, input):
conv3X3 = self.conv3X3(input)
conv5X5_1 = self.conv5X5_1(input)
conv5X5 = self.conv5X5_2(conv5X5_1)
conv7X7_2 = self.conv7X7_2(conv5X5_1)
conv7X7 = self.conv7x7_3(conv7X7_2)
out = torch.cat([conv3X3, conv5X5, conv7X7], dim=1)
out = F.relu(out)
return out
class FPN(nn.Module):
def __init__(self,in_channels_list,out_channels):
super(FPN,self).__init__()
leaky = 0
if out_channels <= 64:
leaky = 0.1
self.output1 = conv_bn1X1(in_channels_list[0], out_channels, stride=1, leaky=leaky)
self.output2 = conv_bn1X1(in_channels_list[1], out_channels, stride=1, leaky=leaky)
self.output3 = conv_bn1X1(in_channels_list[2], out_channels, stride=1, leaky=leaky)
self.merge1 = conv_bn(out_channels, out_channels, leaky=leaky)
self.merge2 = conv_bn(out_channels, out_channels, leaky=leaky)
def forward(self, input):
# names = list(input.keys())
input = list(input.values())
output1 = self.output1(input[0])
output2 = self.output2(input[1])
output3 = self.output3(input[2])
up3 = F.interpolate(output3, size=[output2.size(2), output2.size(3)], mode="nearest")
output2 = output2 + up3
output2 = self.merge2(output2)
up2 = F.interpolate(output2, size=[output1.size(2), output1.size(3)], mode="nearest")
output1 = output1 + up2
output1 = self.merge1(output1)
out = [output1, output2, output3]
return out
class MobileNetV1(nn.Module):
def __init__(self):
super(MobileNetV1, self).__init__()
self.stage1 = nn.Sequential(
conv_bn(3, 8, 2, leaky=0.1), # 3
conv_dw(8, 16, 1), # 7
conv_dw(16, 32, 2), # 11
conv_dw(32, 32, 1), # 19
conv_dw(32, 64, 2), # 27
conv_dw(64, 64, 1), # 43
)
self.stage2 = nn.Sequential(
conv_dw(64, 128, 2), # 43 + 16 = 59
conv_dw(128, 128, 1), # 59 + 32 = 91
conv_dw(128, 128, 1), # 91 + 32 = 123
conv_dw(128, 128, 1), # 123 + 32 = 155
conv_dw(128, 128, 1), # 155 + 32 = 187
conv_dw(128, 128, 1), # 187 + 32 = 219
)
self.stage3 = nn.Sequential(
conv_dw(128, 256, 2), # 219 +3 2 = 241
conv_dw(256, 256, 1), # 241 + 64 = 301
)
self.avg = nn.AdaptiveAvgPool2d((1,1))
self.fc = nn.Linear(256, 1000)
def forward(self, x):
x = self.stage1(x)
x = self.stage2(x)
x = self.stage3(x)
x = self.avg(x)
# x = self.model(x)
x = x.view(-1, 256)
x = self.fc(x)
return x
@@ -0,0 +1,110 @@
import os
import torch
import numpy as np
from copy import deepcopy
from types import SimpleNamespace
from typing import Union, Optional
from .prior_box import PriorBox
from .py_cpu_nms import py_cpu_nms
from .retina_face import RetinaFace
from .config import cfg_mnet, cfg_re50
from .box_utils import decode, decode_landm
__all__ = ['RetinaFacePredictor']
class RetinaFacePredictor(object):
def __init__(self, threshold: float = 0.8, device: Union[str, torch.device] = 'cuda:0',
model: Optional[SimpleNamespace] = None, config: Optional[SimpleNamespace] = None) -> None:
self.threshold = threshold
self.device = device
if model is None:
model = RetinaFacePredictor.get_model()
if config is None:
config = RetinaFacePredictor.create_config()
self.config = SimpleNamespace(**model.config.__dict__, **config.__dict__)
self.net = RetinaFace(cfg=self.config.__dict__, phase='test').to(self.device)
pretrained_dict = torch.load(model.weights, map_location=self.device)
if 'state_dict' in pretrained_dict.keys():
pretrained_dict = {key.split('module.', 1)[-1] if key.startswith('module.') else key: value
for key, value in pretrained_dict['state_dict'].items()}
else:
pretrained_dict = {key.split('module.', 1)[-1] if key.startswith('module.') else key: value
for key, value in pretrained_dict.items()}
self.net.load_state_dict(pretrained_dict, strict=False)
self.net.eval()
self.priors = None
self.previous_size = None
@staticmethod
def get_model(name: str = 'resnet50') -> SimpleNamespace:
from motiondiff_modules import CKPT_DIR_PATH, download_models
RETINA_FACE_PREFIX = "https://github.com/hhj1897/face_detection/raw/71852f00b815f568f3b51f045a418ae84cbe162a/ibug/face_detection/retina_face/weights/"
name = name.lower().strip()
if name == 'resnet50':
download_models({'Resnet50_Final.pth': RETINA_FACE_PREFIX + 'Resnet50_Final.pth'})
return SimpleNamespace(weights=os.path.join(str(CKPT_DIR_PATH), 'Resnet50_Final.pth'),
config=SimpleNamespace(**deepcopy(cfg_re50)))
elif name == 'mobilenet0.25':
download_models({'Resnet50_Final.pth': RETINA_FACE_PREFIX + 'mobilenet0.25_Final.pth'})
return SimpleNamespace(weights=os.path.join(str(CKPT_DIR_PATH), 'mobilenet0.25_Final.pth'),
config=SimpleNamespace(**deepcopy(cfg_mnet)))
else:
raise ValueError('name must be set to either resnet50 or mobilenet0.25')
@staticmethod
def create_config(top_k: int = 750, conf_thresh: float = 0.02,
nms_thresh: float = 0.4, nms_top_k: int = 5000) -> SimpleNamespace:
return SimpleNamespace(top_k=top_k, conf_thresh=conf_thresh, nms_thresh=nms_thresh, nms_top_k=nms_top_k)
@torch.no_grad()
def __call__(self, image: np.ndarray, rgb: bool = True) -> np.ndarray:
im_height, im_width, _ = image.shape
if rgb:
image = image[..., ::-1]
image = image.astype(int) - np.array([104, 117, 123])
image = image.transpose(2, 0, 1)
image = torch.from_numpy(image).unsqueeze(0).float().to(self.device)
scale = torch.Tensor([im_width, im_height, im_width, im_height]).to(self.device)
loc, conf, landms = self.net(image)
image_size = (im_height, im_width)
if self.priors is None or self.previous_size != image_size:
self.priors = PriorBox(self.config.__dict__, image_size=image_size).forward().to(self.device)
self.previous_size = image_size
prior_data = self.priors.data
boxes = decode(loc.data.squeeze(0), prior_data, self.config.variance)
boxes = boxes * scale
boxes = boxes.cpu().numpy()
scores = conf.squeeze(0).data.cpu().numpy()[:, 1]
landms = decode_landm(landms.data.squeeze(0), prior_data, self.config.variance)
scale1 = torch.Tensor([image.shape[3], image.shape[2], image.shape[3], image.shape[2],
image.shape[3], image.shape[2], image.shape[3], image.shape[2],
image.shape[3], image.shape[2]]).to(self.device)
landms = landms * scale1
landms = landms.cpu().numpy()
# ignore low scores
inds = np.where(scores > self.config.conf_thresh)[0]
if len(inds) == 0:
return np.empty(shape=(0, 15), dtype=np.float32)
boxes = boxes[inds]
landms = landms[inds]
scores = scores[inds]
# do NMS
dets = np.hstack((boxes, scores[:, np.newaxis])).astype(np.float32, copy=False)
keep = py_cpu_nms(dets, self.config.nms_thresh, self.config.nms_top_k)
dets = dets[keep, :]
landms = landms[keep]
# keep top-K
dets = dets[:self.config.top_k, :]
landms = landms[:self.config.top_k, :]
dets = np.concatenate((dets, landms), axis=1)
# further filter by confidence
inds = np.where(dets[:, 4] >= self.threshold)[0]
if len(inds) == 0:
return np.empty(shape=(0, 15), dtype=np.float32)
else:
return dets[inds]
@@ -0,0 +1 @@
from .s3fd_predictor import S3FDPredictor
@@ -0,0 +1,175 @@
import torch
import torch.nn as nn
import torch.nn.init as init
import torch.nn.functional as F
from .utils import Detect, PriorBox
class L2Norm(nn.Module):
def __init__(self, n_channels, scale):
super(L2Norm, self).__init__()
self.n_channels = n_channels
self.gamma = scale or None
self.eps = 1e-10
self.weight = nn.Parameter(torch.Tensor(self.n_channels))
self.reset_parameters()
def reset_parameters(self):
init.constant_(self.weight, self.gamma)
def forward(self, x):
norm = x.pow(2).sum(dim=1, keepdim=True).sqrt() + self.eps
x = torch.div(x, norm)
out = self.weight.unsqueeze(0).unsqueeze(2).unsqueeze(3).expand_as(x) * x
return out
class S3FDNet(nn.Module):
def __init__(self, config, device='cuda'):
super(S3FDNet, self).__init__()
self.config = config
self.device = device
self.vgg = nn.ModuleList([
nn.Conv2d(3, 64, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(64, 64, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.MaxPool2d(2, 2),
nn.Conv2d(64, 128, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(128, 128, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.MaxPool2d(2, 2),
nn.Conv2d(128, 256, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 256, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 256, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.MaxPool2d(2, 2, ceil_mode=True),
nn.Conv2d(256, 512, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(512, 512, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(512, 512, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.MaxPool2d(2, 2),
nn.Conv2d(512, 512, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(512, 512, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(512, 512, 3, 1, padding=1),
nn.ReLU(inplace=True),
nn.MaxPool2d(2, 2),
nn.Conv2d(512, 1024, 3, 1, padding=6, dilation=6),
nn.ReLU(inplace=True),
nn.Conv2d(1024, 1024, 1, 1),
nn.ReLU(inplace=True),
])
self.L2Norm3_3 = L2Norm(256, 10)
self.L2Norm4_3 = L2Norm(512, 8)
self.L2Norm5_3 = L2Norm(512, 5)
self.extras = nn.ModuleList([
nn.Conv2d(1024, 256, 1, 1),
nn.Conv2d(256, 512, 3, 2, padding=1),
nn.Conv2d(512, 128, 1, 1),
nn.Conv2d(128, 256, 3, 2, padding=1),
])
self.loc = nn.ModuleList([
nn.Conv2d(256, 4, 3, 1, padding=1),
nn.Conv2d(512, 4, 3, 1, padding=1),
nn.Conv2d(512, 4, 3, 1, padding=1),
nn.Conv2d(1024, 4, 3, 1, padding=1),
nn.Conv2d(512, 4, 3, 1, padding=1),
nn.Conv2d(256, 4, 3, 1, padding=1),
])
self.conf = nn.ModuleList([
nn.Conv2d(256, 4, 3, 1, padding=1),
nn.Conv2d(512, 2, 3, 1, padding=1),
nn.Conv2d(512, 2, 3, 1, padding=1),
nn.Conv2d(1024, 2, 3, 1, padding=1),
nn.Conv2d(512, 2, 3, 1, padding=1),
nn.Conv2d(256, 2, 3, 1, padding=1),
])
self.priors = None
self.previous_size = None
self.softmax = nn.Softmax(dim=-1)
self.detect = Detect(self.config)
def forward(self, x):
size = x.size()[2:]
sources = list()
loc = list()
conf = list()
for k in range(16):
x = self.vgg[k](x)
s = self.L2Norm3_3(x)
sources.append(s)
for k in range(16, 23):
x = self.vgg[k](x)
s = self.L2Norm4_3(x)
sources.append(s)
for k in range(23, 30):
x = self.vgg[k](x)
s = self.L2Norm5_3(x)
sources.append(s)
for k in range(30, len(self.vgg)):
x = self.vgg[k](x)
sources.append(x)
# apply extra layers and cache source layer outputs
for k, v in enumerate(self.extras):
x = F.relu(v(x), inplace=True)
if k % 2 == 1:
sources.append(x)
# apply multibox head to source layers
loc_x = self.loc[0](sources[0])
conf_x = self.conf[0](sources[0])
max_conf, _ = torch.max(conf_x[:, 0:3, :, :], dim=1, keepdim=True)
conf_x = torch.cat((max_conf, conf_x[:, 3:, :, :]), dim=1)
loc.append(loc_x.permute(0, 2, 3, 1).contiguous())
conf.append(conf_x.permute(0, 2, 3, 1).contiguous())
for i in range(1, len(sources)):
x = sources[i]
conf.append(self.conf[i](x).permute(0, 2, 3, 1).contiguous())
loc.append(self.loc[i](x).permute(0, 2, 3, 1).contiguous())
if self.priors is None or self.previous_size != size:
with torch.no_grad():
features_maps = []
for i in range(len(loc)):
feat = []
feat += [loc[i].size(1), loc[i].size(2)]
features_maps += [feat]
self.priors = PriorBox(size, features_maps, self.config).forward().to(self.device)
self.previous_size = size
loc = torch.cat([o.view(o.size(0), -1) for o in loc], 1)
conf = torch.cat([o.view(o.size(0), -1) for o in conf], 1)
conf = self.softmax(conf.view(conf.size(0), -1, 2))
output = self.detect(loc.view(loc.size(0), -1, 4), conf, self.priors)
return output
@@ -0,0 +1,70 @@
import os
import torch
import numpy as np
from types import SimpleNamespace
from typing import Union, Optional
from .s3fd_net import S3FDNet
__all__ = ['S3FDPredictor']
class S3FDPredictor(object):
def __init__(self, threshold: float = 0.8, device: Union[str, torch.device] = 'cuda:0',
model: Optional[SimpleNamespace] = None, config: Optional[SimpleNamespace] = None) -> None:
self.threshold = threshold
self.device = device
if model is None:
model = S3FDPredictor.get_model()
if config is None:
config = S3FDPredictor.create_config()
self.config = SimpleNamespace(**model.config.__dict__, **config.__dict__)
self.net = S3FDNet(config=self.config, device=self.device).to(self.device)
self.net.load_state_dict(torch.load(model.weights, map_location=self.device))
self.net.eval()
@staticmethod
def get_model(name: str = 's3fd') -> SimpleNamespace:
from motiondiff_modules import CKPT_DIR_PATH, download_models
S3FD_FACE_PREFIX = "https://github.com/hhj1897/face_detection/raw/71852f00b815f568f3b51f045a418ae84cbe162a/ibug/face_detection/s3fd/weights/"
name = name.lower().strip()
if name == 's3fd':
download_models({'s3fd_weights.pth': S3FD_FACE_PREFIX + 's3fd_weights.pth'})
return SimpleNamespace(weights=os.path.realpath(CKPT_DIR_PATH, 's3fd_weights.pth'),
config=SimpleNamespace(num_classes=2, variance=(0.1, 0.2),
prior_min_sizes=(16, 32, 64, 128, 256, 512),
prior_steps=(4, 8, 16, 32, 64, 128), prior_clip=False))
else:
raise ValueError('name must be set to s3fd')
@staticmethod
def create_config(top_k: int = 750, conf_thresh: float = 0.05,nms_thresh: float = 0.3,
nms_top_k: int = 5000, use_nms_np: bool = True) -> SimpleNamespace:
return SimpleNamespace(top_k=top_k, conf_thresh=conf_thresh, nms_thresh=nms_thresh,
nms_top_k=nms_top_k, use_nms_np=use_nms_np)
@torch.no_grad()
def __call__(self, image: np.ndarray, rgb: bool = True) -> np.ndarray:
w, h = image.shape[1], image.shape[0]
if not rgb:
image = image[..., ::-1]
image = image.astype(int) - np.array([123, 117, 104])
image = image.transpose(2, 0, 1)
image = image.reshape((1,) + image.shape)
image = torch.from_numpy(image).float().to(self.device)
bboxes = []
detections = self.net(image)
scale = torch.Tensor([w, h, w, h]).to(detections.device)
for i in range(detections.size(1)):
j = 0
while detections[0, i, j, 0] >= self.threshold:
score = detections[0, i, j, 0]
pt = (detections[0, i, j, 1:] * scale).cpu().numpy()
bbox = (pt[0], pt[1], pt[2], pt[3], score)
bboxes.append(bbox)
j += 1
if len(bboxes) > 0:
return np.array(bboxes)
else:
return np.empty(shape=(0, 5), dtype=np.float32)
@@ -0,0 +1,206 @@
import torch
import numpy as np
from itertools import product
def decode(loc, priors, variances):
"""Decode locations from predictions using priors to undo
the encoding we did for offset regression at train time.
Args:
loc (tensor): location predictions for loc layers,
Shape: [num_priors,4]
priors (tensor): Prior boxes in center-offset form.
Shape: [num_priors,4].
variances: (list[float]) Variances of priorboxes
Return:
decoded bounding box predictions
"""
boxes = torch.cat((
priors[:, :2] + loc[:, :2] * variances[0] * priors[:, 2:],
priors[:, 2:] * torch.exp(loc[:, 2:] * variances[1])), 1)
boxes[:, :2] -= boxes[:, 2:] / 2
boxes[:, 2:] += boxes[:, :2]
return boxes
def nms(boxes, scores, overlap=0.5, top_k=200):
"""Apply non-maximum suppression at test time to avoid detecting too many
overlapping bounding boxes for a given object.
Args:
boxes: (tensor) The location preds for the img, Shape: [num_priors,4].
scores: (tensor) The class predscores for the img, Shape:[num_priors].
overlap: (float) The overlap thresh for suppressing unnecessary boxes.
top_k: (int) The Maximum number of box preds to consider.
Return:
The indices of the kept boxes with respect to num_priors.
"""
keep = scores.new(scores.size(0)).zero_().long()
if boxes.numel() == 0:
return keep, 0
x1 = boxes[:, 0]
y1 = boxes[:, 1]
x2 = boxes[:, 2]
y2 = boxes[:, 3]
area = torch.mul(x2 - x1, y2 - y1)
v, idx = scores.sort(0) # sort in ascending order
# I = I[v >= 0.01]
idx = idx[-top_k:] # indices of the top-k largest vals
xx1 = boxes.new()
yy1 = boxes.new()
xx2 = boxes.new()
yy2 = boxes.new()
w = boxes.new()
h = boxes.new()
# keep = torch.Tensor()
count = 0
while idx.numel() > 0:
i = idx[-1] # index of current largest val
# keep.append(i)
keep[count] = i
count += 1
if idx.size(0) == 1:
break
idx = idx[:-1] # remove kept element from view
# load bboxes of next highest vals
torch.index_select(x1, 0, idx, out=xx1)
torch.index_select(y1, 0, idx, out=yy1)
torch.index_select(x2, 0, idx, out=xx2)
torch.index_select(y2, 0, idx, out=yy2)
# store element-wise max with next highest score
xx1 = torch.clamp(xx1, min=x1[i])
yy1 = torch.clamp(yy1, min=y1[i])
xx2 = torch.clamp(xx2, max=x2[i])
yy2 = torch.clamp(yy2, max=y2[i])
w.resize_as_(xx2)
h.resize_as_(yy2)
w = xx2 - xx1
h = yy2 - yy1
# check sizes of xx1 and xx2.. after each iteration
w = torch.clamp(w, min=0.0)
h = torch.clamp(h, min=0.0)
inter = w * h
# IoU = i / (area(a) + area(b) - i)
rem_areas = torch.index_select(area, 0, idx) # load remaining areas)
union = (rem_areas - inter) + area[i]
IoU = inter / union # store result in iou
# keep only elements with an IoU <= overlap
idx = idx[IoU.le(overlap)]
return keep, count
def nms_np(boxes, scores, overlap=0.5, top_k=200):
"""Apply non-maximum suppression at test time to avoid detecting too many
overlapping bounding boxes for a given object, using numpy (for speed).
Args:
boxes: (tensor) The location preds for the img, Shape: [num_priors,4].
scores: (tensor) The class predscores for the img, Shape:[num_priors].
overlap: (float) The overlap thresh for suppressing unnecessary boxes.
top_k: (int) The Maximum number of box preds to consider.
Return:
The indices of the kept boxes with respect to num_priors.
"""
if scores.size(0) == 0:
return [], 0
else:
areas = torch.mul(boxes[:, 2] - boxes[:, 0], boxes[:, 3] - boxes[:, 1]).cpu().numpy()
x1, y1 = boxes[:, 0].cpu().numpy(), boxes[:, 1].cpu().numpy()
x2, y2 = boxes[:, 2].cpu().numpy(), boxes[:, 3].cpu().numpy()
scores = scores.cpu().numpy()
order = scores.argsort()[: -top_k - 1: -1]
keep = []
while order.size > 0:
i = order[0]
keep.append(i)
xx1, yy1 = np.maximum(x1[i], x1[order[1:]]), np.maximum(y1[i], y1[order[1:]])
xx2, yy2 = np.minimum(x2[i], x2[order[1:]]), np.minimum(y2[i], y2[order[1:]])
w, h = np.maximum(0.0, xx2 - xx1), np.maximum(0.0, yy2 - yy1)
ovr = w * h / (areas[i] + areas[order[1:]] - w * h)
inds = np.where(ovr <= overlap)[0]
order = order[inds + 1]
return keep, len(keep)
class Detect(object):
def __init__(self, config):
self.config = config
def __call__(self, loc_data, conf_data, prior_data):
num = loc_data.size(0)
num_priors = prior_data.size(0)
conf_preds = conf_data.view(num, num_priors, self.config.num_classes).transpose(2, 1)
batch_priors = prior_data.view(-1, num_priors, 4).expand(num, num_priors, 4)
batch_priors = batch_priors.contiguous().view(-1, 4)
decoded_boxes = decode(loc_data.view(-1, 4), batch_priors, self.config.variance)
decoded_boxes = decoded_boxes.view(num, num_priors, 4)
output = torch.zeros(num, self.config.num_classes, self.config.top_k, 5)
for i in range(num):
boxes = decoded_boxes[i].clone()
conf_scores = conf_preds[i].clone()
for cl in range(1, self.config.num_classes):
c_mask = conf_scores[cl].gt(self.config.conf_thresh)
scores = conf_scores[cl][c_mask]
if scores.dim() == 0:
continue
l_mask = c_mask.unsqueeze(1).expand_as(boxes)
boxes_ = boxes[l_mask].view(-1, 4)
if self.config.use_nms_np:
ids, count = nms_np(boxes_, scores, self.config.nms_thresh, self.config.nms_top_k)
else:
ids, count = nms(boxes_, scores, self.config.nms_thresh, self.config.nms_top_k)
count = count if count < self.config.top_k else self.config.top_k
output[i, cl, :count] = torch.cat((scores[ids[:count]].unsqueeze(1), boxes_[ids[:count]]), 1)
return output
class PriorBox(object):
def __init__(self, input_size, feature_maps, config):
self.imh = input_size[0]
self.imw = input_size[1]
self.feature_maps = feature_maps
self.config = config
def forward(self):
mean = []
for k, fmap in enumerate(self.feature_maps):
feath = fmap[0]
featw = fmap[1]
for i, j in product(range(feath), range(featw)):
f_kw = self.imw / self.config.prior_steps[k]
f_kh = self.imh / self.config.prior_steps[k]
cx = (j + 0.5) / f_kw
cy = (i + 0.5) / f_kh
s_kw = self.config.prior_min_sizes[k] / self.imw
s_kh = self.config.prior_min_sizes[k] / self.imh
mean += [cx, cy, s_kw, s_kh]
output = torch.FloatTensor(mean).view(-1, 4)
if self.config.prior_clip:
output.clamp_(max=1, min=0)
return output
@@ -0,0 +1,2 @@
from .head_pose_estimator import HeadPoseEstimator
from .simple_face_tracker import SimpleFaceTracker
@@ -0,0 +1,78 @@
import os
import cv2
import math
import numpy as np
from typing import Optional, Tuple
__all__ = ['HeadPoseEstimator']
class HeadPoseEstimator(object):
def __init__(self, mean_shape_path: str = os.path.join(os.path.dirname(__file__),
'data', 'bfm_lms.npy')) -> None:
# Load the 68-point mean shape derived from BFM
mean_shape = np.load(mean_shape_path)
# Calculate the 5-points mean shape
left_eye = mean_shape[[37, 38, 40, 41]].mean(axis=0)
right_eye = mean_shape[[43, 44, 46, 47]].mean(axis=0)
self._mean_shape_5pts = np.vstack((left_eye, right_eye, mean_shape[[30, 48, 54]]))
# Flip the y coordinates of the mean shape to match that of the image coordinate system
self._mean_shape_5pts[:, 1] = -self._mean_shape_5pts[:, 1]
def __call__(self, landmarks: np.ndarray, image_width: int = 0, image_height: int = 0,
camera_matrix: Optional[np.ndarray] = None, dist_coeffs: Optional[np.ndarray] = None,
output_preference: int = 0) -> Tuple[float, float, float]:
# Form the camera matrix
if camera_matrix is None:
if image_width <= 0 or image_height <= 0:
raise ValueError(
'image_width and image_height must be specified when camera_matrix is not given directly')
else:
camera_matrix = np.array([[image_width + image_height, 0, image_width / 2.0],
[0, image_width + image_height, image_height / 2.0],
[0, 0, 1]], dtype=float)
# Prepare the landmarks
if landmarks.shape[0] == 68:
landmarks = landmarks[17:]
if landmarks.shape[0] in [49, 51]:
left_eye = landmarks[[20, 21, 23, 24]].mean(axis=0)
right_eye = landmarks[[26, 27, 29, 30]].mean(axis=0)
landmarks = np.vstack((left_eye, right_eye, landmarks[[13, 31, 37]]))
# Use EPnP to estimate pitch, yaw, and roll
_, rvec, _ = cv2.solvePnP(self._mean_shape_5pts, np.expand_dims(landmarks, axis=1),
camera_matrix, dist_coeffs, flags=cv2.SOLVEPNP_EPNP)
rot_mat, _ = cv2.Rodrigues(rvec)
if 1.0 + rot_mat[2, 0] < 1e-9:
pitch = 0.0
yaw = 90.0
roll = -math.atan2(rot_mat[0, 1], rot_mat[0, 2]) / math.pi * 180.0
elif 1.0 - rot_mat[2, 0] < 1e-9:
pitch = 0.0
yaw = -90.0
roll = math.atan2(-rot_mat[0, 1], -rot_mat[0, 2]) / math.pi * 180.0
else:
pitch = math.atan2(rot_mat[2, 1], rot_mat[2, 2]) / math.pi * 180.0
yaw = -math.asin(rot_mat[2, 0]) / math.pi * 180.0
roll = math.atan2(rot_mat[1, 0], rot_mat[0, 0]) / math.pi * 180.0
# Respond to output_preference:
# output_preference == 1: limit pitch to the range of -90.0 ~ 90.0
# output_preference == 2: limit yaw to the range of -90.0 ~ 90.0 (already satisfied)
# output_preference == 3: limit roll to the range of -90.0 ~ 90.0
# otherwise: minimise total rotation, min(abs(pitch) + abs(yaw) + abs(roll))
if output_preference != 2:
alt_pitch = pitch - 180.0 if pitch > 0.0 else pitch + 180.0
alt_yaw = -180.0 - yaw if yaw < 0.0 else 180.0 - yaw
alt_roll = roll - 180.0 if roll > 0.0 else roll + 180.0
if (output_preference == 1 and -90.0 < alt_pitch < 90.0 or
output_preference == 3 and -90.0 < alt_roll < 90.0 or
output_preference not in (1, 2, 3) and
abs(alt_pitch) + abs(alt_yaw) + abs(alt_roll) < abs(pitch) + abs(yaw) + abs(roll)):
pitch, yaw, roll = alt_pitch, alt_yaw, alt_roll
return -pitch, yaw, roll
@@ -0,0 +1,90 @@
import numpy as np
from typing import List, Optional
from scipy.optimize import linear_sum_assignment
__all__ = ['SimpleFaceTracker']
class SimpleFaceTracker(object):
def __init__(self, iou_threshold: float = 0.4, minimum_face_size: float = 0.0) -> None:
self._iou_threshold = iou_threshold
self._minimum_face_size = minimum_face_size
self._tracklets = []
self._tracklet_counter = 0
@property
def iou_threshold(self) -> float:
return self._iou_threshold
@iou_threshold.setter
def iou_threshold(self, threshold: float) -> None:
self._iou_threshold = threshold
@property
def minimum_face_size(self) -> float:
return self._minimum_face_size
@minimum_face_size.setter
def minimum_face_size(self, face_size: float) -> None:
self._minimum_face_size = face_size
def __call__(self, face_boxes: np.ndarray) -> List[Optional[int]]:
if face_boxes.size <= 0:
self._tracklets = []
return []
# Calculate area of the faces
face_areas = np.abs((face_boxes[:, 2] - face_boxes[:, 0]) * (face_boxes[:, 3] - face_boxes[:, 1]))
# Prepare tracklets
for tracklet in self._tracklets:
tracklet['tracked'] = False
# Calculate the distance matrix based on IOU
iou_distance_threshold = np.clip(1.0 - self._iou_threshold, 0.0, 1.0)
min_face_area = max(self._minimum_face_size ** 2, np.finfo(float).eps)
distances = np.full(shape=(face_boxes.shape[0], len(self._tracklets)),
fill_value=2.0 * min(face_boxes.shape[0], len(self._tracklets)), dtype=float)
for row, face_box in enumerate(face_boxes):
if face_areas[row] >= min_face_area:
for col, tracklet in enumerate(self._tracklets):
x_left = max(min(face_box[0], face_box[2]), min(tracklet['bbox'][0], tracklet['bbox'][2]))
y_top = max(min(face_box[1], face_box[3]), min(tracklet['bbox'][1], tracklet['bbox'][3]))
x_right = min(max(face_box[2], face_box[0]), max(tracklet['bbox'][2], tracklet['bbox'][0]))
y_bottom = min(max(face_box[3], face_box[1]), max(tracklet['bbox'][3], tracklet['bbox'][1]))
if x_right <= x_left or y_bottom <= y_top:
distance = 1.0
else:
intersection_area = (x_right - x_left) * (y_bottom - y_top)
distance = 1.0 - intersection_area / float(face_areas[row] + tracklet['area'] -
intersection_area)
if distance <= iou_distance_threshold:
distances[row, col] = distance
# ID assignment
tracked_ids = [None] * face_boxes.shape[0]
for row, col in zip(*linear_sum_assignment(distances)):
if distances[row, col] <= iou_distance_threshold:
tracked_ids[row] = self._tracklets[col]['id']
self._tracklets[col]['bbox'] = face_boxes[row, :4].copy()
self._tracklets[col]['area'] = face_areas[row]
self._tracklets[col]['tracked'] = True
# Remove expired tracklets
self._tracklets = [x for x in self._tracklets if x['tracked']]
# Register new faces
for idx, face_box in enumerate(face_boxes):
if face_areas[idx] >= min_face_area and tracked_ids[idx] is None:
self._tracklet_counter += 1
self._tracklets.append({'bbox': face_box[:4].copy(), 'area': face_areas[idx],
'id': self._tracklet_counter, 'tracked': True})
tracked_ids[idx] = self._tracklets[-1]['id']
return tracked_ids
def reset(self, reset_tracklet_counter: bool = True) -> None:
self._tracklets = []
if reset_tracklet_counter:
self._tracklet_counter = 0
@@ -0,0 +1,272 @@
# -*- coding: utf-8 -*-
#
# Max-Planck-Gesellschaft zur Förderung der Wissenschaften e.V. (MPG) is
# holder of all proprietary rights on this computer program.
# Using this computer program means that you agree to the terms
# in the LICENSE file included with this software distribution.
# Any use not explicitly granted by the LICENSE is prohibited.
#
# Copyright©2019 Max-Planck-Gesellschaft zur Förderung
# der Wissenschaften e.V. (MPG). acting on behalf of its Max Planck Institute
# for Intelligent Systems. All rights reserved.
#
# For comments or questions, please email us at deca@tue.mpg.de
# For commercial licensing contact, please contact ps-license@tuebingen.mpg.de
import torch
import torch.nn as nn
import numpy as np
import pickle
import torch.nn.functional as F
from .lbs import lbs, batch_rodrigues, vertices2landmarks, rot_mat_to_euler
def to_tensor(array, dtype=torch.float32):
if 'torch.tensor' not in str(type(array)):
return torch.tensor(array, dtype=dtype)
def to_np(array, dtype=np.float32):
if 'scipy.sparse' in str(type(array)):
array = array.todense()
return np.array(array, dtype=dtype)
class Struct(object):
def __init__(self, **kwargs):
for key, val in kwargs.items():
setattr(self, key, val)
class FLAME(nn.Module):
"""
borrowed from https://github.com/soubhiksanyal/FLAME_PyTorch/blob/master/FLAME.py
Given flame parameters this class generates a differentiable FLAME function
which outputs the a mesh and 2D/3D facial landmarks
"""
def __init__(self, config):
super(FLAME, self).__init__()
# print("creating the FLAME Decoder")
with open(config.flame_model_path, 'rb') as f:
ss = pickle.load(f, encoding='latin1')
flame_model = Struct(**ss)
self.dtype = torch.float32
self.register_buffer('faces_tensor', to_tensor(to_np(flame_model.f, dtype=np.int64), dtype=torch.long))
# The vertices of the template model
self.register_buffer('v_template', to_tensor(to_np(flame_model.v_template), dtype=self.dtype))
# The shape components and expression
shapedirs = to_tensor(to_np(flame_model.shapedirs), dtype=self.dtype)
shapedirs = torch.cat([shapedirs[:,:,:config.n_shape], shapedirs[:,:,300:300+config.n_exp]], 2)
self.register_buffer('shapedirs', shapedirs)
# The pose components
num_pose_basis = flame_model.posedirs.shape[-1]
posedirs = np.reshape(flame_model.posedirs, [-1, num_pose_basis]).T
self.register_buffer('posedirs', to_tensor(to_np(posedirs), dtype=self.dtype))
#
self.register_buffer('J_regressor', to_tensor(to_np(flame_model.J_regressor), dtype=self.dtype))
parents = to_tensor(to_np(flame_model.kintree_table[0])).long(); parents[0] = -1
self.register_buffer('parents', parents)
self.register_buffer('lbs_weights', to_tensor(to_np(flame_model.weights), dtype=self.dtype))
# Fixing Eyeball and neck rotation
default_eyball_pose = torch.zeros([1, 6], dtype=self.dtype, requires_grad=False)
self.register_parameter('eye_pose', nn.Parameter(default_eyball_pose,
requires_grad=False))
default_neck_pose = torch.zeros([1, 3], dtype=self.dtype, requires_grad=False)
self.register_parameter('neck_pose', nn.Parameter(default_neck_pose,
requires_grad=False))
# Static and Dynamic Landmark embeddings for FLAME
lmk_embeddings = np.load(config.flame_lmk_embedding_path, allow_pickle=True, encoding='latin1')
lmk_embeddings = lmk_embeddings[()]
self.register_buffer('lmk_faces_idx', torch.from_numpy(lmk_embeddings['static_lmk_faces_idx']).long())
self.register_buffer('lmk_bary_coords', torch.from_numpy(lmk_embeddings['static_lmk_bary_coords']).to(self.dtype))
self.register_buffer('dynamic_lmk_faces_idx', lmk_embeddings['dynamic_lmk_faces_idx'].long())
self.register_buffer('dynamic_lmk_bary_coords', lmk_embeddings['dynamic_lmk_bary_coords'].to(self.dtype))
self.register_buffer('full_lmk_faces_idx', torch.from_numpy(lmk_embeddings['full_lmk_faces_idx']).long())
self.register_buffer('full_lmk_bary_coords', torch.from_numpy(lmk_embeddings['full_lmk_bary_coords']).to(self.dtype))
neck_kin_chain = []; NECK_IDX=1
curr_idx = torch.tensor(NECK_IDX, dtype=torch.long)
while curr_idx != -1:
neck_kin_chain.append(curr_idx)
curr_idx = self.parents[curr_idx]
self.register_buffer('neck_kin_chain', torch.stack(neck_kin_chain))
def _find_dynamic_lmk_idx_and_bcoords(self, pose, dynamic_lmk_faces_idx,
dynamic_lmk_b_coords,
neck_kin_chain, dtype=torch.float32):
"""
Selects the face contour depending on the reletive position of the head
Input:
vertices: N X num_of_vertices X 3
pose: N X full pose
dynamic_lmk_faces_idx: The list of contour face indexes
dynamic_lmk_b_coords: The list of contour barycentric weights
neck_kin_chain: The tree to consider for the relative rotation
dtype: Data type
return:
The contour face indexes and the corresponding barycentric weights
"""
batch_size = pose.shape[0]
aa_pose = torch.index_select(pose.view(batch_size, -1, 3), 1,
neck_kin_chain)
rot_mats = batch_rodrigues(
aa_pose.view(-1, 3), dtype=dtype).view(batch_size, -1, 3, 3)
rel_rot_mat = torch.eye(3, device=pose.device,
dtype=dtype).unsqueeze_(dim=0).expand(batch_size, -1, -1)
for idx in range(len(neck_kin_chain)):
rel_rot_mat = torch.bmm(rot_mats[:, idx], rel_rot_mat)
y_rot_angle = torch.round(
torch.clamp(rot_mat_to_euler(rel_rot_mat) * 180.0 / np.pi,
max=39)).to(dtype=torch.long)
neg_mask = y_rot_angle.lt(0).to(dtype=torch.long)
mask = y_rot_angle.lt(-39).to(dtype=torch.long)
neg_vals = mask * 78 + (1 - mask) * (39 - y_rot_angle)
y_rot_angle = (neg_mask * neg_vals +
(1 - neg_mask) * y_rot_angle)
dyn_lmk_faces_idx = torch.index_select(dynamic_lmk_faces_idx,
0, y_rot_angle)
dyn_lmk_b_coords = torch.index_select(dynamic_lmk_b_coords,
0, y_rot_angle)
return dyn_lmk_faces_idx, dyn_lmk_b_coords
def _vertices2landmarks(self, vertices, faces, lmk_faces_idx, lmk_bary_coords):
"""
Calculates landmarks by barycentric interpolation
Input:
vertices: torch.tensor NxVx3, dtype = torch.float32
The tensor of input vertices
faces: torch.tensor (N*F)x3, dtype = torch.long
The faces of the mesh
lmk_faces_idx: torch.tensor N X L, dtype = torch.long
The tensor with the indices of the faces used to calculate the
landmarks.
lmk_bary_coords: torch.tensor N X L X 3, dtype = torch.float32
The tensor of barycentric coordinates that are used to interpolate
the landmarks
Returns:
landmarks: torch.tensor NxLx3, dtype = torch.float32
The coordinates of the landmarks for each mesh in the batch
"""
# Extract the indices of the vertices for each face
# NxLx3
batch_size, num_verts = vertices.shape[:dd2]
lmk_faces = torch.index_select(faces, 0, lmk_faces_idx.view(-1)).view(
1, -1, 3).view(batch_size, lmk_faces_idx.shape[1], -1)
lmk_faces += torch.arange(batch_size, dtype=torch.long).view(-1, 1, 1).to(
device=vertices.device) * num_verts
lmk_vertices = vertices.view(-1, 3)[lmk_faces]
landmarks = torch.einsum('blfi,blf->bli', [lmk_vertices, lmk_bary_coords])
return landmarks
def seletec_3d68(self, vertices):
landmarks3d = vertices2landmarks(vertices, self.faces_tensor,
self.full_lmk_faces_idx.repeat(vertices.shape[0], 1),
self.full_lmk_bary_coords.repeat(vertices.shape[0], 1, 1))
return landmarks3d
def forward(self, shape_params=None, expression_params=None, pose_params=None, eye_pose_params=None):
"""
Input:
shape_params: N X number of shape parameters
expression_params: N X number of expression parameters
pose_params: N X number of pose parameters (6)
return:d
vertices: N X V X 3
landmarks: N X number of landmarks X 3
"""
batch_size = shape_params.shape[0]
if pose_params is None:
pose_params = self.eye_pose.expand(batch_size, -1)
if eye_pose_params is None:
eye_pose_params = self.eye_pose.expand(batch_size, -1)
betas = torch.cat([shape_params, expression_params], dim=1)
full_pose = torch.cat([pose_params[:, :3], self.neck_pose.expand(batch_size, -1), pose_params[:, 3:], eye_pose_params], dim=1)
template_vertices = self.v_template.unsqueeze(0).expand(batch_size, -1, -1)
vertices, _ = lbs(betas, full_pose, template_vertices,
self.shapedirs, self.posedirs,
self.J_regressor, self.parents,
self.lbs_weights, dtype=self.dtype)
lmk_faces_idx = self.lmk_faces_idx.unsqueeze(dim=0).expand(batch_size, -1)
lmk_bary_coords = self.lmk_bary_coords.unsqueeze(dim=0).expand(batch_size, -1, -1)
dyn_lmk_faces_idx, dyn_lmk_bary_coords = self._find_dynamic_lmk_idx_and_bcoords(
full_pose, self.dynamic_lmk_faces_idx,
self.dynamic_lmk_bary_coords,
self.neck_kin_chain, dtype=self.dtype)
lmk_faces_idx = torch.cat([dyn_lmk_faces_idx, lmk_faces_idx], 1)
lmk_bary_coords = torch.cat([dyn_lmk_bary_coords, lmk_bary_coords], 1)
landmarks2d = vertices2landmarks(vertices, self.faces_tensor,
lmk_faces_idx,
lmk_bary_coords)
bz = vertices.shape[0]
landmarks3d = vertices2landmarks(vertices, self.faces_tensor,
self.full_lmk_faces_idx.repeat(bz, 1),
self.full_lmk_bary_coords.repeat(bz, 1, 1))
return vertices, landmarks2d, landmarks3d
class FLAMETex(nn.Module):
"""
FLAME texture:
https://github.com/TimoBolkart/TF_FLAME/blob/ade0ab152300ec5f0e8555d6765411555c5ed43d/sample_texture.py#L64
FLAME texture converted from BFM:
https://github.com/TimoBolkart/BFM_to_FLAME
"""
def __init__(self, config):
super(FLAMETex, self).__init__()
if config.tex_type == 'BFM':
mu_key = 'MU'
pc_key = 'PC'
n_pc = 199
tex_path = config.tex_path
tex_space = np.load(tex_path)
texture_mean = tex_space[mu_key].reshape(1, -1)
texture_basis = tex_space[pc_key].reshape(-1, n_pc)
elif config.tex_type == 'FLAME':
mu_key = 'mean'
pc_key = 'tex_dir'
n_pc = 200
tex_path = config.flame_tex_path
tex_space = np.load(tex_path)
texture_mean = tex_space[mu_key].reshape(1, -1)/255.
texture_basis = tex_space[pc_key].reshape(-1, n_pc)/255.
else:
print('texture type ', config.tex_type, 'not exist!')
raise NotImplementedError
n_tex = config.n_tex
num_components = texture_basis.shape[1]
texture_mean = torch.from_numpy(texture_mean).float()[None,...]
texture_basis = torch.from_numpy(texture_basis[:,:n_tex]).float()[None,...]
self.register_buffer('texture_mean', texture_mean)
self.register_buffer('texture_basis', texture_basis)
def forward(self, texcode):
'''
texcode: [batchsize, n_tex]
texture: [bz, 3, 256, 256], range: 0-1
'''
bs = texcode.shape[0]
texcode = texcode[:1]
# we use the same (first frame) texture for all frames
texture = self.texture_mean + (self.texture_basis*texcode[:,None,:]).sum(-1)
texture = texture.reshape(texcode.shape[0], 512, 512, 3).permute(0,3,1,2)
texture = F.interpolate(texture, [256, 256])
texture = texture[:,[2,1,0], :,:].repeat(bs,1,1,1)
return texture
@@ -0,0 +1,90 @@
# -*- coding: utf-8 -*-
import torch.nn as nn
import torch
import torch.nn.functional as F
from . import resnet
class PerceptualEncoder(nn.Module):
def __init__(self, outsize, cfg):
super(PerceptualEncoder, self).__init__()
if cfg.backbone == "mobilenetv2":
self.encoder = torch.hub.load('pytorch/vision:v0.8.1', 'mobilenet_v2', pretrained=True)
feature_size = 1280
elif cfg.backbone == "resnet50":
self.encoder = resnet.load_ResNet50Model() #out: 2048
feature_size = 2048
### regressor
self.temporal = nn.Sequential(
nn.Conv1d(in_channels=feature_size, out_channels=256, kernel_size=5, stride=1, padding=2),
nn.BatchNorm1d(256),
nn.ReLU()
)
self.layers = nn.Sequential(
nn.Linear(256, 53),
)
self.backbone = cfg.backbone
def forward(self, inputs):
is_video_batch = inputs.ndim == 5
if self.backbone == 'resnet50':
features = self.encoder(inputs).squeeze(-1).squeeze(-1)
else:
inputs_ = inputs
if is_video_batch:
B, T, C, H, W = inputs.shape
inputs_ = inputs.view(B * T, C, H, W)
features = self.encoder.features(inputs_)
features = nn.functional.adaptive_avg_pool2d(features, (1, 1)).squeeze(-1).squeeze(-1)
if is_video_batch:
features = features.view(B, T, -1)
features = features
if is_video_batch:
features = features.permute(0, 2, 1)
else:
features = features.permute(1,0).unsqueeze(0)
features = self.temporal(features)
if is_video_batch:
features = features.permute(0, 2, 1)
else:
features = features.squeeze(0).permute(1,0)
parameters = self.layers(features)
parameters[...,50] = F.relu(parameters[...,50]) # jaw x is highly improbably negative and can introduce artifacts
return parameters[...,:50], parameters[...,50:]
class ResnetEncoder(nn.Module):
def __init__(self, outsize):
super(ResnetEncoder, self).__init__()
feature_size = 2048
self.encoder = resnet.load_ResNet50Model() #out: 2048
### regressor
self.layers = nn.Sequential(
nn.Linear(feature_size, 1024),
nn.ReLU(),
nn.Linear(1024, outsize)
)
def forward(self, inputs):
inputs_ = inputs
if inputs.ndim == 5: # batch of videos
B, T, C, H, W = inputs.shape
inputs_ = inputs.view(B * T, C, H, W)
features = self.encoder(inputs_)
parameters = self.layers(features)
if inputs.ndim == 5: # batch of videos
parameters = parameters.view(B, T, -1)
return parameters
@@ -0,0 +1,39 @@
# -*- coding: utf-8 -*-
#
# Max-Planck-Gesellschaft zur Förderung der Wissenschaften e.V. (MPG) is
# holder of all proprietary rights on this computer program.
# Using this computer program means that you agree to the terms
# in the LICENSE file included with this software distribution.
# Any use not explicitly granted by the LICENSE is prohibited.
#
# Copyright©2019 Max-Planck-Gesellschaft zur Förderung
# der Wissenschaften e.V. (MPG). acting on behalf of its Max Planck Institute
# for Intelligent Systems. All rights reserved.
#
# For comments or questions, please email us at deca@tue.mpg.de
# For commercial licensing contact, please contact ps-license@tuebingen.mpg.de
import torch.nn as nn
from torchvision import models
from . import resnet
class ExpressionLossNet(nn.Module):
""" Code borrowed from EMOCA https://github.com/radekd91/emoca """
def __init__(self):
super(ExpressionLossNet, self).__init__()
self.backbone = resnet.load_ResNet50Model() #out: 2048
self.linear = nn.Sequential(
nn.Linear(2048, 10))
def forward2(self, inputs):
features = self.backbone(inputs)
out = self.linear(features)
return features, out
def forward(self, inputs):
features = self.backbone(inputs)
return features
@@ -0,0 +1,378 @@
# -*- coding: utf-8 -*-
# Max-Planck-Gesellschaft zur Förderung der Wissenschaften e.V. (MPG) is
# holder of all proprietary rights on this computer program.
# You can only use this computer program if you have closed
# a license agreement with MPG or you get the right to use the computer
# program from someone who is authorized to grant you that right.
# Any use of the computer program without a valid license is prohibited and
# liable to prosecution.
#
# Copyright©2019 Max-Planck-Gesellschaft zur Förderung
# der Wissenschaften e.V. (MPG). acting on behalf of its Max Planck Institute
# for Intelligent Systems. All rights reserved.
#
# Contact: ps-license@tuebingen.mpg.de
from __future__ import absolute_import
from __future__ import print_function
from __future__ import division
import numpy as np
import torch
import torch.nn.functional as F
def rot_mat_to_euler(rot_mats):
# Calculates rotation matrix to euler angles
# Careful for extreme cases of eular angles like [0.0, pi, 0.0]
sy = torch.sqrt(rot_mats[:, 0, 0] * rot_mats[:, 0, 0] +
rot_mats[:, 1, 0] * rot_mats[:, 1, 0])
return torch.atan2(-rot_mats[:, 2, 0], sy)
def find_dynamic_lmk_idx_and_bcoords(vertices, pose, dynamic_lmk_faces_idx,
dynamic_lmk_b_coords,
neck_kin_chain, dtype=torch.float32):
''' Compute the faces, barycentric coordinates for the dynamic landmarks
To do so, we first compute the rotation of the neck around the y-axis
and then use a pre-computed look-up table to find the faces and the
barycentric coordinates that will be used.
Special thanks to Soubhik Sanyal (soubhik.sanyal@tuebingen.mpg.de)
for providing the original TensorFlow implementation and for the LUT.
Parameters
----------
vertices: torch.tensor BxVx3, dtype = torch.float32
The tensor of input vertices
pose: torch.tensor Bx(Jx3), dtype = torch.float32
The current pose of the body model
dynamic_lmk_faces_idx: torch.tensor L, dtype = torch.long
The look-up table from neck rotation to faces
dynamic_lmk_b_coords: torch.tensor Lx3, dtype = torch.float32
The look-up table from neck rotation to barycentric coordinates
neck_kin_chain: list
A python list that contains the indices of the joints that form the
kinematic chain of the neck.
dtype: torch.dtype, optional
Returns
-------
dyn_lmk_faces_idx: torch.tensor, dtype = torch.long
A tensor of size BxL that contains the indices of the faces that
will be used to compute the current dynamic landmarks.
dyn_lmk_b_coords: torch.tensor, dtype = torch.float32
A tensor of size BxL that contains the indices of the faces that
will be used to compute the current dynamic landmarks.
'''
batch_size = vertices.shape[0]
aa_pose = torch.index_select(pose.view(batch_size, -1, 3), 1,
neck_kin_chain)
rot_mats = batch_rodrigues(
aa_pose.view(-1, 3), dtype=dtype).view(batch_size, -1, 3, 3)
rel_rot_mat = torch.eye(3, device=vertices.device,
dtype=dtype).unsqueeze_(dim=0)
for idx in range(len(neck_kin_chain)):
rel_rot_mat = torch.bmm(rot_mats[:, idx], rel_rot_mat)
y_rot_angle = torch.round(
torch.clamp(-rot_mat_to_euler(rel_rot_mat) * 180.0 / np.pi,
max=39)).to(dtype=torch.long)
neg_mask = y_rot_angle.lt(0).to(dtype=torch.long)
mask = y_rot_angle.lt(-39).to(dtype=torch.long)
neg_vals = mask * 78 + (1 - mask) * (39 - y_rot_angle)
y_rot_angle = (neg_mask * neg_vals +
(1 - neg_mask) * y_rot_angle)
dyn_lmk_faces_idx = torch.index_select(dynamic_lmk_faces_idx,
0, y_rot_angle)
dyn_lmk_b_coords = torch.index_select(dynamic_lmk_b_coords,
0, y_rot_angle)
return dyn_lmk_faces_idx, dyn_lmk_b_coords
def vertices2landmarks(vertices, faces, lmk_faces_idx, lmk_bary_coords):
''' Calculates landmarks by barycentric interpolation
Parameters
----------
vertices: torch.tensor BxVx3, dtype = torch.float32
The tensor of input vertices
faces: torch.tensor Fx3, dtype = torch.long
The faces of the mesh
lmk_faces_idx: torch.tensor L, dtype = torch.long
The tensor with the indices of the faces used to calculate the
landmarks.
lmk_bary_coords: torch.tensor Lx3, dtype = torch.float32
The tensor of barycentric coordinates that are used to interpolate
the landmarks
Returns
-------
landmarks: torch.tensor BxLx3, dtype = torch.float32
The coordinates of the landmarks for each mesh in the batch
'''
# Extract the indices of the vertices for each face
# BxLx3
batch_size, num_verts = vertices.shape[:2]
device = vertices.device
lmk_faces = torch.index_select(faces, 0, lmk_faces_idx.view(-1)).view(
batch_size, -1, 3)
lmk_faces += torch.arange(
batch_size, dtype=torch.long, device=device).view(-1, 1, 1) * num_verts
lmk_vertices = vertices.view(-1, 3)[lmk_faces].view(
batch_size, -1, 3, 3)
landmarks = torch.einsum('blfi,blf->bli', [lmk_vertices, lmk_bary_coords])
return landmarks
def lbs(betas, pose, v_template, shapedirs, posedirs, J_regressor, parents,
lbs_weights, pose2rot=True, dtype=torch.float32):
''' Performs Linear Blend Skinning with the given shape and pose parameters
Parameters
----------
betas : torch.tensor BxNB
The tensor of shape parameters
pose : torch.tensor Bx(J + 1) * 3
The pose parameters in axis-angle format
v_template torch.tensor BxVx3
The template mesh that will be deformed
shapedirs : torch.tensor 1xNB
The tensor of PCA shape displacements
posedirs : torch.tensor Px(V * 3)
The pose PCA coefficients
J_regressor : torch.tensor JxV
The regressor array that is used to calculate the joints from
the position of the vertices
parents: torch.tensor J
The array that describes the kinematic tree for the model
lbs_weights: torch.tensor N x V x (J + 1)
The linear blend skinning weights that represent how much the
rotation matrix of each part affects each vertex
pose2rot: bool, optional
Flag on whether to convert the input pose tensor to rotation
matrices. The default value is True. If False, then the pose tensor
should already contain rotation matrices and have a size of
Bx(J + 1)x9
dtype: torch.dtype, optional
Returns
-------
verts: torch.tensor BxVx3
The vertices of the mesh after applying the shape and pose
displacements.
joints: torch.tensor BxJx3
The joints of the model
'''
batch_size = max(betas.shape[0], pose.shape[0])
device = betas.device
# Add shape contribution
v_shaped = v_template + blend_shapes(betas, shapedirs)
# Get the joints
# NxJx3 array
J = vertices2joints(J_regressor, v_shaped)
# 3. Add pose blend shapes
# N x J x 3 x 3
ident = torch.eye(3, dtype=dtype, device=device)
if pose2rot:
rot_mats = batch_rodrigues(
pose.view(-1, 3), dtype=dtype).view([batch_size, -1, 3, 3])
pose_feature = (rot_mats[:, 1:, :, :] - ident).view([batch_size, -1])
# (N x P) x (P, V * 3) -> N x V x 3
pose_offsets = torch.matmul(pose_feature, posedirs) \
.view(batch_size, -1, 3)
else:
pose_feature = pose[:, 1:].view(batch_size, -1, 3, 3) - ident
rot_mats = pose.view(batch_size, -1, 3, 3)
pose_offsets = torch.matmul(pose_feature.view(batch_size, -1),
posedirs).view(batch_size, -1, 3)
v_posed = pose_offsets + v_shaped
# 4. Get the global joint location
J_transformed, A = batch_rigid_transform(rot_mats, J, parents, dtype=dtype)
# 5. Do skinning:
# W is N x V x (J + 1)
W = lbs_weights.unsqueeze(dim=0).expand([batch_size, -1, -1])
# (N x V x (J + 1)) x (N x (J + 1) x 16)
num_joints = J_regressor.shape[0]
T = torch.matmul(W, A.view(batch_size, num_joints, 16)) \
.view(batch_size, -1, 4, 4)
homogen_coord = torch.ones([batch_size, v_posed.shape[1], 1],
dtype=dtype, device=device)
v_posed_homo = torch.cat([v_posed, homogen_coord], dim=2)
v_homo = torch.matmul(T, torch.unsqueeze(v_posed_homo, dim=-1))
verts = v_homo[:, :, :3, 0]
return verts, J_transformed
def vertices2joints(J_regressor, vertices):
''' Calculates the 3D joint locations from the vertices
Parameters
----------
J_regressor : torch.tensor JxV
The regressor array that is used to calculate the joints from the
position of the vertices
vertices : torch.tensor BxVx3
The tensor of mesh vertices
Returns
-------
torch.tensor BxJx3
The location of the joints
'''
return torch.einsum('bik,ji->bjk', [vertices, J_regressor])
def blend_shapes(betas, shape_disps):
''' Calculates the per vertex displacement due to the blend shapes
Parameters
----------
betas : torch.tensor Bx(num_betas)
Blend shape coefficients
shape_disps: torch.tensor Vx3x(num_betas)
Blend shapes
Returns
-------
torch.tensor BxVx3
The per-vertex displacement due to shape deformation
'''
# Displacement[b, m, k] = sum_{l} betas[b, l] * shape_disps[m, k, l]
# i.e. Multiply each shape displacement by its corresponding beta and
# then sum them.
blend_shape = torch.einsum('bl,mkl->bmk', [betas, shape_disps])
return blend_shape
def batch_rodrigues(rot_vecs, epsilon=1e-8, dtype=torch.float32):
''' Calculates the rotation matrices for a batch of rotation vectors
Parameters
----------
rot_vecs: torch.tensor Nx3
array of N axis-angle vectors
Returns
-------
R: torch.tensor Nx3x3
The rotation matrices for the given axis-angle parameters
'''
batch_size = rot_vecs.shape[0]
device = rot_vecs.device
angle = torch.norm(rot_vecs + 1e-8, dim=1, keepdim=True)
rot_dir = rot_vecs / angle
cos = torch.unsqueeze(torch.cos(angle), dim=1)
sin = torch.unsqueeze(torch.sin(angle), dim=1)
# Bx1 arrays
rx, ry, rz = torch.split(rot_dir, 1, dim=1)
K = torch.zeros((batch_size, 3, 3), dtype=dtype, device=device)
zeros = torch.zeros((batch_size, 1), dtype=dtype, device=device)
K = torch.cat([zeros, -rz, ry, rz, zeros, -rx, -ry, rx, zeros], dim=1) \
.view((batch_size, 3, 3))
ident = torch.eye(3, dtype=dtype, device=device).unsqueeze(dim=0)
rot_mat = ident + sin * K + (1 - cos) * torch.bmm(K, K)
return rot_mat
def transform_mat(R, t):
''' Creates a batch of transformation matrices
Args:
- R: Bx3x3 array of a batch of rotation matrices
- t: Bx3x1 array of a batch of translation vectors
Returns:
- T: Bx4x4 Transformation matrix
'''
# No padding left or right, only add an extra row
return torch.cat([F.pad(R, [0, 0, 0, 1]),
F.pad(t, [0, 0, 0, 1], value=1)], dim=2)
def batch_rigid_transform(rot_mats, joints, parents, dtype=torch.float32):
"""
Applies a batch of rigid transformations to the joints
Parameters
----------
rot_mats : torch.tensor BxNx3x3
Tensor of rotation matrices
joints : torch.tensor BxNx3
Locations of joints
parents : torch.tensor BxN
The kinematic tree of each object
dtype : torch.dtype, optional:
The data type of the created tensors, the default is torch.float32
Returns
-------
posed_joints : torch.tensor BxNx3
The locations of the joints after applying the pose rotations
rel_transforms : torch.tensor BxNx4x4
The relative (with respect to the root joint) rigid transformations
for all the joints
"""
joints = torch.unsqueeze(joints, dim=-1)
rel_joints = joints.clone()
rel_joints[:, 1:] -= joints[:, parents[1:]]
# transforms_mat = transform_mat(
# rot_mats.view(-1, 3, 3),
# rel_joints.view(-1, 3, 1)).view(-1, joints.shape[1], 4, 4)
transforms_mat = transform_mat(
rot_mats.view(-1, 3, 3),
rel_joints.reshape(-1, 3, 1)).reshape(-1, joints.shape[1], 4, 4)
transform_chain = [transforms_mat[:, 0]]
for i in range(1, parents.shape[0]):
# Subtract the joint location at the rest pose
# No need for rotation, since it's identity when at rest
curr_res = torch.matmul(transform_chain[parents[i]],
transforms_mat[:, i])
transform_chain.append(curr_res)
transforms = torch.stack(transform_chain, dim=1)
# The last column of the transformations contains the posed joints
posed_joints = transforms[:, :, :3, 3]
# The last column of the transformations contains the posed joints
posed_joints = transforms[:, :, :3, 3]
joints_homogen = F.pad(joints, [0, 0, 0, 1])
rel_transforms = transforms - F.pad(
torch.matmul(transforms, joints_homogen), [3, 0, 0, 0, 0, 0, 0, 0])
return posed_joints, rel_transforms
@@ -0,0 +1,261 @@
"""
Author: Soubhik Sanyal
Copyright (c) 2019, Soubhik Sanyal
All rights reserved.
Loads different resnet models
"""
'''
file: Resnet.py
date: 2018_05_02
author: zhangxiong(1025679612@qq.com)
mark: copied from pytorch source code
'''
import torch.nn as nn
import torch.nn.functional as F
import torch
from torch.nn.parameter import Parameter
import torch.optim as optim
import numpy as np
import math
import torchvision
class ResNet(nn.Module):
def __init__(self, block, layers, num_classes=1000):
self.inplanes = 64
super(ResNet, self).__init__()
self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3,
bias=False)
self.bn1 = nn.BatchNorm2d(64)
self.relu = nn.ReLU(inplace=True)
self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
self.layer1 = self._make_layer(block, 64, layers[0])
self.layer2 = self._make_layer(block, 128, layers[1], stride=2)
self.layer3 = self._make_layer(block, 256, layers[2], stride=2)
self.layer4 = self._make_layer(block, 512, layers[3], stride=2)
self.avgpool = nn.AvgPool2d(7, stride=1)
# self.fc = nn.Linear(512 * block.expansion, num_classes)
for m in self.modules():
if isinstance(m, nn.Conv2d):
n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
m.weight.data.normal_(0, math.sqrt(2. / n))
elif isinstance(m, nn.BatchNorm2d):
m.weight.data.fill_(1)
m.bias.data.zero_()
def _make_layer(self, block, planes, blocks, stride=1):
downsample = None
if stride != 1 or self.inplanes != planes * block.expansion:
downsample = nn.Sequential(
nn.Conv2d(self.inplanes, planes * block.expansion,
kernel_size=1, stride=stride, bias=False),
nn.BatchNorm2d(planes * block.expansion),
)
layers = []
layers.append(block(self.inplanes, planes, stride, downsample))
self.inplanes = planes * block.expansion
for i in range(1, blocks):
layers.append(block(self.inplanes, planes))
return nn.Sequential(*layers)
def forward(self, x):
x = self.conv1(x)
x = self.bn1(x)
x = self.relu(x)
x = self.maxpool(x)
x = self.layer1(x)
x = self.layer2(x)
x = self.layer3(x)
x1 = self.layer4(x)
x2 = self.avgpool(x1)
x2 = x2.view(x2.size(0), -1)
# x = self.fc(x)
## x2: [bz, 2048] for shape
## x1: [bz, 2048, 7, 7] for texture
return x2
class Bottleneck(nn.Module):
expansion = 4
def __init__(self, inplanes, planes, stride=1, downsample=None):
super(Bottleneck, self).__init__()
self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=1, bias=False)
self.bn1 = nn.BatchNorm2d(planes)
self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=stride,
padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(planes)
self.conv3 = nn.Conv2d(planes, planes * 4, kernel_size=1, bias=False)
self.bn3 = nn.BatchNorm2d(planes * 4)
self.relu = nn.ReLU(inplace=True)
self.downsample = downsample
self.stride = stride
def forward(self, x):
residual = x
out = self.conv1(x)
out = self.bn1(out)
out = self.relu(out)
out = self.conv2(out)
out = self.bn2(out)
out = self.relu(out)
out = self.conv3(out)
out = self.bn3(out)
if self.downsample is not None:
residual = self.downsample(x)
out += residual
out = self.relu(out)
return out
def conv3x3(in_planes, out_planes, stride=1):
"""3x3 convolution with padding"""
return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride,
padding=1, bias=False)
class BasicBlock(nn.Module):
expansion = 1
def __init__(self, inplanes, planes, stride=1, downsample=None):
super(BasicBlock, self).__init__()
self.conv1 = conv3x3(inplanes, planes, stride)
self.bn1 = nn.BatchNorm2d(planes)
self.relu = nn.ReLU(inplace=True)
self.conv2 = conv3x3(planes, planes)
self.bn2 = nn.BatchNorm2d(planes)
self.downsample = downsample
self.stride = stride
def forward(self, x):
residual = x
out = self.conv1(x)
out = self.bn1(out)
out = self.relu(out)
out = self.conv2(out)
out = self.bn2(out)
if self.downsample is not None:
residual = self.downsample(x)
out += residual
out = self.relu(out)
return out
def copy_parameter_from_resnet(model, resnet_dict):
cur_state_dict = model.state_dict()
# import ipdb; ipdb.set_trace()
for name, param in list(resnet_dict.items())[0:None]:
if name not in cur_state_dict:
# print(name, ' not available in reconstructed resnet')
continue
if isinstance(param, Parameter):
param = param.data
try:
cur_state_dict[name].copy_(param)
except:
# print(name, ' is inconsistent!')
continue
# print('copy resnet state dict finished!')
# import ipdb; ipdb.set_trace()
def load_ResNet50Model():
model = ResNet(Bottleneck, [3, 4, 6, 3])
copy_parameter_from_resnet(model, torchvision.models.resnet50(pretrained = True).state_dict())
return model
def load_ResNet101Model():
model = ResNet(Bottleneck, [3, 4, 23, 3])
copy_parameter_from_resnet(model, torchvision.models.resnet101(pretrained = True).state_dict())
return model
def load_ResNet152Model():
model = ResNet(Bottleneck, [3, 8, 36, 3])
copy_parameter_from_resnet(model, torchvision.models.resnet152(pretrained = True).state_dict())
return model
# model.load_state_dict(checkpoint['model_state_dict'])
######## Unet
class DoubleConv(nn.Module):
"""(convolution => [BN] => ReLU) * 2"""
def __init__(self, in_channels, out_channels):
super().__init__()
self.double_conv = nn.Sequential(
nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),
nn.BatchNorm2d(out_channels),
nn.ReLU(inplace=True),
nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),
nn.BatchNorm2d(out_channels),
nn.ReLU(inplace=True)
)
def forward(self, x):
return self.double_conv(x)
class Down(nn.Module):
"""Downscaling with maxpool then double conv"""
def __init__(self, in_channels, out_channels):
super().__init__()
self.maxpool_conv = nn.Sequential(
nn.MaxPool2d(2),
DoubleConv(in_channels, out_channels)
)
def forward(self, x):
return self.maxpool_conv(x)
class Up(nn.Module):
"""Upscaling then double conv"""
def __init__(self, in_channels, out_channels, bilinear=True):
super().__init__()
# if bilinear, use the normal convolutions to reduce the number of channels
if bilinear:
self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
else:
self.up = nn.ConvTranspose2d(in_channels // 2, in_channels // 2, kernel_size=2, stride=2)
self.conv = DoubleConv(in_channels, out_channels)
def forward(self, x1, x2):
x1 = self.up(x1)
# input is CHW
diffY = x2.size()[2] - x1.size()[2]
diffX = x2.size()[3] - x1.size()[3]
x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,
diffY // 2, diffY - diffY // 2])
# if you have padding issues, see
# https://github.com/HaiyongJiang/U-Net-Pytorch-Unstructured-Buggy/commit/0e854509c2cea854e247a9c615f175f76fbb2e3a
# https://github.com/xiaopeng-liao/Pytorch-UNet/commit/8ebac70e633bac59fc22bb5195e513d5832fb3bd
x = torch.cat([x2, x1], dim=1)
return self.conv(x)
class OutConv(nn.Module):
def __init__(self, in_channels, out_channels):
super(OutConv, self).__init__()
self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)
def forward(self, x):
return self.conv(x)
+301
View File
@@ -0,0 +1,301 @@
# -*- coding: utf-8 -*-
#
# Max-Planck-Gesellschaft zur Förderung der Wissenschaften e.V. (MPG) is
# holder of all proprietary rights on this computer program.
# Using this computer program means that you agree to the terms
# in the LICENSE file included with this software distribution.
# Any use not explicitly granted by the LICENSE is prohibited.
#
# Copyright©2019 Max-Planck-Gesellschaft zur Förderung
# der Wissenschaften e.V. (MPG). acting on behalf of its Max Planck Institute
# for Intelligent Systems. All rights reserved.
#
# For comments or questions, please email us at deca@tue.mpg.de
# For commercial licensing contact, please contact ps-license@tuebingen.mpg.de
import os
import torch
import torch.nn as nn
import torch.nn.functional as F
from .models.encoders import PerceptualEncoder
from .utils.renderer import SRenderY, set_rasterizer
from .models.encoders import ResnetEncoder
from .models.FLAME import FLAME, FLAMETex
from .utils import util
from .utils.tensor_cropper import transform_points
from skimage.io import imread
torch.backends.cudnn.benchmark = True
import numpy as np
def recursive_to(x, target):
"""
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
class SPECTRE(nn.Module):
def __init__(self, config=None, device='cuda'):
super(SPECTRE, self).__init__()
self.cfg = config
self.device = device
self.image_size = self.cfg.dataset.image_size
self.uv_size = self.cfg.model.uv_size
self._create_model(self.cfg.model)
#self._setup_renderer(self.cfg.model)
def _setup_renderer(self, model_cfg):
set_rasterizer(self.cfg.rasterizer_type)
self.render = SRenderY(self.image_size, obj_filename=model_cfg.topology_path, uv_size=model_cfg.uv_size, rasterizer_type=self.cfg.rasterizer_type).to(self.device)
# face mask for rendering details
mask = imread(model_cfg.face_eye_mask_path).astype(np.float32)/255.; mask = torch.from_numpy(mask[:,:,0])[None,None,:,:].contiguous()
self.uv_face_eye_mask = F.interpolate(mask, [model_cfg.uv_size, model_cfg.uv_size]).to(self.device)
mask = imread(model_cfg.face_mask_path).astype(np.float32)/255.; mask = torch.from_numpy(mask[:,:,0])[None,None,:,:].contiguous()
self.uv_face_mask = F.interpolate(mask, [model_cfg.uv_size, model_cfg.uv_size]).to(self.device)
# displacement correction
fixed_dis = np.load(model_cfg.fixed_displacement_path)
self.fixed_uv_dis = torch.tensor(fixed_dis).float().to(self.device)
# mean texture
mean_texture = imread(model_cfg.mean_tex_path).astype(np.float32)/255.; mean_texture = torch.from_numpy(mean_texture.transpose(2,0,1))[None,:,:,:].contiguous()
self.mean_texture = F.interpolate(mean_texture, [model_cfg.uv_size, model_cfg.uv_size]).to(self.device)
# dense mesh template, for save detail mesh
self.dense_template = np.load(model_cfg.dense_template_path, allow_pickle=True, encoding='latin1').item()
def _create_model(self, model_cfg):
# set up parameters
self.n_param = model_cfg.n_shape + model_cfg.n_tex + model_cfg.n_exp + model_cfg.n_pose + model_cfg.n_cam + model_cfg.n_light
self.n_cond = model_cfg.n_exp + 3 # exp + jaw pose
self.num_list = [model_cfg.n_shape, model_cfg.n_tex, model_cfg.n_exp, model_cfg.n_pose, model_cfg.n_cam,
model_cfg.n_light]
self.param_dict = {i: model_cfg.get('n_' + i) for i in model_cfg.param_list}
# encoders
self.E_flame = ResnetEncoder(outsize=self.n_param).to(self.device)
self.E_expression = PerceptualEncoder(model_cfg.n_exp, model_cfg).to(self.device)
# decoders
self.flame = FLAME(model_cfg).to(self.device)
if model_cfg.use_tex:
self.flametex = FLAMETex(model_cfg).to(self.device)
# resume model
model_path = self.cfg.pretrained_modelpath
if os.path.exists(model_path):
# print(f'trained model found. load {model_path}')
checkpoint = torch.load(model_path, map_location="cpu")
checkpoint = recursive_to(checkpoint, self.device)
if 'state_dict' in checkpoint.keys():
self.checkpoint = checkpoint['state_dict']
else:
self.checkpoint = checkpoint
processed_checkpoint = {}
processed_checkpoint["E_flame"] = {}
processed_checkpoint["E_expression"] = {}
if 'deca' in list(self.checkpoint.keys())[0]:
for key in self.checkpoint.keys():
# print(key)
k = key.replace("deca.","")
if "E_flame" in key:
processed_checkpoint["E_flame"][k.replace("E_flame.","")] = self.checkpoint[key]#.replace("E_flame","")
elif "E_expression" in key:
processed_checkpoint["E_expression"][k.replace("E_expression.","")] = self.checkpoint[key]#.replace("E_flame","")
else:
pass
else:
processed_checkpoint = self.checkpoint
self.E_flame.load_state_dict(processed_checkpoint['E_flame'], strict=True)
try:
m,u = self.E_expression.load_state_dict(processed_checkpoint['E_expression'], strict=True)
# print('Missing keys', m)
# print('Unexpected keys', u)
# pass
except Exception as e:
print(f'Missing keys {e} in expression encoder weights. If starting training from scratch this is normal.')
else:
raise(f'please check model path: {model_path}')
# eval mode
self.E_flame.eval()
self.E_expression.eval()
self.E_flame.requires_grad_(False)
def decompose_code(self, code, num_dict):
''' Convert a flattened parameter vector to a dictionary of parameters
code_dict.keys() = ['shape', 'tex', 'exp', 'pose', 'cam', 'light']
'''
code_dict = {}
start = 0
for key in num_dict:
end = start + int(num_dict[key])
code_dict[key] = code[..., start:end]
start = end
if key == 'light':
dims_ = code_dict[key].ndim -1 # (to be able to handle batches of videos)
code_dict[key] = code_dict[key].reshape(*code_dict[key].shape[:dims_], 9, 3)
return code_dict
def encode(self, images):
with torch.no_grad():
parameters = self.E_flame(images)
codedict = self.decompose_code(parameters, self.param_dict)
deca_exp = codedict['exp'].clone()
deca_jaw = codedict['pose'][...,3:].clone()
codedict['images'] = images
codedict['exp'], jaw = self.E_expression(images)
codedict['pose'][..., 3:] = jaw
return codedict, deca_exp, deca_jaw
def decode(self, codedict, rendering=True, vis_lmk=True, return_vis=True,
render_orig=False, original_image=None, tform=None):
images = codedict['images']
is_video_batch = images.ndim == 5
if is_video_batch:
B, T, C, H, W = images.shape
images = images.view(B*T, C, H, W)
codedict_ = codedict
codedict = {}
for key in codedict_.keys():
# if key != 'images':
codedict[key] = codedict_[key].view(B*T, *codedict_[key].shape[2:])
batch_size = images.shape[0]
## decode
verts, landmarks2d, landmarks3d = self.flame(shape_params=codedict['shape'], expression_params=codedict['exp'],
pose_params=codedict['pose'])
if self.cfg.model.use_tex:
albedo = self.flametex(codedict['tex']).detach()
else:
albedo = torch.zeros([batch_size, 3, self.uv_size, self.uv_size], device=images.device)
landmarks3d_world = landmarks3d.clone()
## projection
landmarks2d = util.batch_orth_proj(landmarks2d, codedict['cam'])[:, :, :2];
landmarks2d[:, :, 1:] = -landmarks2d[:, :,
1:]
landmarks3d = util.batch_orth_proj(landmarks3d, codedict['cam']);
landmarks3d[:, :, 1:] = -landmarks3d[:, :,
1:]
trans_verts = util.batch_orth_proj(verts, codedict['cam']);
trans_verts[:, :, 1:] = -trans_verts[:, :, 1:]
opdict = {
'verts': verts,
'trans_verts': trans_verts,
'landmarks2d': landmarks2d,
'landmarks3d': landmarks3d,
'landmarks3d_world': landmarks3d_world,
}
if rendering and render_orig and original_image is not None and tform is not None:
points_scale = [self.image_size, self.image_size]
_, _, h, w = original_image.shape
trans_verts = transform_points(trans_verts, tform, points_scale, [h, w])
landmarks2d = transform_points(landmarks2d, tform, points_scale, [h, w])
landmarks3d = transform_points(landmarks3d, tform, points_scale, [h, w])
background = images
else:
h, w = self.image_size, self.image_size
background = None
if rendering:
if self.cfg.model.use_tex:
ops = self.render(verts, trans_verts, albedo, codedict['light'])
## output
opdict['predicted_inner_mouth'] = ops['predicted_inner_mouth']
opdict['grid'] = ops['grid']
opdict['rendered_images'] = ops['images']
opdict['alpha_images'] = ops['alpha_images']
opdict['normal_images'] = ops['normal_images']
opdict['images'] = images
else:
shape_images, _, grid, alpha_images, pos_mask = self.render.render_shape(verts, trans_verts, h=h, w=w,
images=background,
return_grid=True,
return_pos=True)
opdict['rendered_images'] = shape_images
if self.cfg.model.use_tex:
opdict['albedo'] = albedo
if vis_lmk:
landmarks3d_vis = self.visofp(ops['transformed_normals']) # /self.image_size
landmarks3d = torch.cat([landmarks3d, landmarks3d_vis], dim=2)
opdict['landmarks3d'] = landmarks3d
if is_video_batch:
for key in opdict.keys():
opdict[key] = opdict[key].view(B, T, *opdict[key].shape[1:])
if return_vis:
## render shape
shape_images, _, grid, alpha_images, pos_mask = self.render.render_shape(verts, trans_verts, h=h, w=w,
images=background, return_grid=True, return_pos=True)
# opdict['uv_texture_gt'] = uv_texture_gt
visdict = {
# 'inputs': images,
'landmarks2d': util.tensor_vis_landmarks(images, landmarks2d),
'landmarks3d': util.tensor_vis_landmarks(images, landmarks3d),
'shape_images': shape_images,
# 'rendered_images': ops['images']
}
if is_video_batch:
for key in visdict.keys():
visdict[key] = visdict[key].view(B, T, *visdict[key].shape[1:])
return opdict, visdict
else:
return opdict
def train(self):
self.E_expression.train()
self.E_flame.eval()
def eval(self):
self.E_expression.eval()
self.E_flame.eval()
def model_dict(self):
return {
'E_flame': self.E_flame.state_dict(),
'E_expression': self.E_expression.state_dict(),
}
@@ -0,0 +1,91 @@
import torch.nn as nn
import numpy as np
import torch
import torch.nn.functional as F
def l2_distance(verts1, verts2):
return torch.sqrt(((verts1 - verts2)**2).sum(2)).mean(1).mean()
### ------------------------------------- Losses/Regularizations for vertices
def batch_kp_2d_l1_loss(real_2d_kp, predicted_2d_kp, weights=None):
"""
Computes the l1 loss between the ground truth keypoints and the predicted keypoints
Inputs:
kp_gt : N x K x 3
kp_pred: N x K x 2
"""
if weights is not None:
real_2d_kp[:,:,2] = weights[None,:]*real_2d_kp[:,:,2]
kp_gt = real_2d_kp.view(-1, 3)
kp_pred = predicted_2d_kp.contiguous().view(-1, 2)
vis = kp_gt[:, 2]
k = torch.sum(vis) * 2.0 + 1e-8
dif_abs = torch.abs(kp_gt[:, :2] - kp_pred).sum(1)
return torch.matmul(dif_abs, vis) * 1.0 / k
def landmark_loss(predicted_landmarks, landmarks_gt, weight=1.):
if torch.is_tensor(landmarks_gt) is not True:
real_2d = torch.cat(landmarks_gt).cuda()
else:
real_2d = torch.cat([landmarks_gt, torch.ones((landmarks_gt.shape[0], 68, 1)).cuda()], dim=-1)
loss_lmk_2d = batch_kp_2d_l1_loss(real_2d, predicted_landmarks)
return loss_lmk_2d * weight
def weighted_landmark_loss(predicted_landmarks, landmarks_gt, weight=1.):
#smaller inner landmark weights
# (predicted_theta, predicted_verts, predicted_landmarks) = ringnet_outputs[-1]
# import ipdb; ipdb.set_trace()
real_2d = landmarks_gt
weights = torch.ones((68,)).cuda()
weights[5:7] = 2
weights[10:12] = 2
# nose points
weights[27:36] = 1.5
weights[30] = 3
weights[31] = 3
weights[35] = 3
# set mouth to zero
weights[60:68] = 0
weights[48:60] = 0
weights[48] = 0
weights[54] = 0
# weights[36:48] = 0 # these are eyes
loss_lmk_2d = batch_kp_2d_l1_loss(real_2d, predicted_landmarks, weights)
return loss_lmk_2d * weight
def rel_dis(landmarks):
lip_right = landmarks[:, [57, 51, 48, 60, 61, 62, 63], :]
lip_left = landmarks[:, [8, 33, 54, 64, 67, 66, 65], :]
# lip_right = landmarks[:, [61, 62, 63], :]
# lip_left = landmarks[:, [67, 66, 65], :]
dis = torch.sqrt(((lip_right - lip_left) ** 2).sum(2)) # [bz, 4]
return dis
def relative_landmark_loss(predicted_landmarks, landmarks_gt, weight=1.):
if torch.is_tensor(landmarks_gt) is not True:
real_2d = torch.cat(landmarks_gt)#.cuda()
else:
real_2d = torch.cat([landmarks_gt, torch.ones((landmarks_gt.shape[0], 68, 1)).to(device=predicted_landmarks.device) #.cuda()
], dim=-1)
pred_lipd = rel_dis(predicted_landmarks[:, :, :2])
gt_lipd = rel_dis(real_2d[:, :, :2])
loss = (pred_lipd - gt_lipd).abs().mean()
# loss = F.mse_loss(pred_lipd, gt_lipd)
return loss.mean()
@@ -0,0 +1,360 @@
# -*- coding: utf-8 -*-
#
# Max-Planck-Gesellschaft zur Förderung der Wissenschaften e.V. (MPG) is
# holder of all proprietary rights on this computer program.
# Using this computer program means that you agree to the terms
# in the LICENSE file included with this software distribution.
# Any use not explicitly granted by the LICENSE is prohibited.
#
# Copyright©2019 Max-Planck-Gesellschaft zur Förderung
# der Wissenschaften e.V. (MPG). acting on behalf of its Max Planck Institute
# for Intelligent Systems. All rights reserved.
#
# For comments or questions, please email us at deca@tue.mpg.de
# For commercial licensing contact, please contact ps-license@tuebingen.mpg.de
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from skimage.io import imread
import imageio
from . import util
def set_rasterizer(type = 'pytorch3d'):
if type == 'pytorch3d':
global Meshes, load_obj, rasterize_meshes
from pytorch3d.structures import Meshes
from pytorch3d.io import load_obj
from pytorch3d.renderer.mesh import rasterize_meshes
else:
NotImplementedError
class Pytorch3dRasterizer(nn.Module):
""" Borrowed from https://github.com/facebookresearch/pytorch3d
Notice:
x,y,z are in image space, normalized
can only render squared image now
"""
def __init__(self, image_size=224):
"""
use fixed raster_settings for rendering faces
"""
super().__init__()
raster_settings = {
'image_size': image_size,
'blur_radius': 0.0,
'faces_per_pixel': 1,
'bin_size': None,
'max_faces_per_bin': None,
'perspective_correct': False,
}
raster_settings = util.dict2obj(raster_settings)
self.raster_settings = raster_settings
def forward(self, vertices, faces, attributes=None, h=None, w=None):
fixed_vertices = vertices.clone()
fixed_vertices[...,:2] = -fixed_vertices[...,:2]
raster_settings = self.raster_settings
if h is None and w is None:
image_size = raster_settings.image_size
else:
image_size = [h, w]
if h>w:
fixed_vertices[..., 1] = fixed_vertices[..., 1]*h/w
else:
fixed_vertices[..., 0] = fixed_vertices[..., 0]*w/h
meshes_screen = Meshes(verts=fixed_vertices.float(), faces=faces.long())
pix_to_face, zbuf, bary_coords, dists = rasterize_meshes(
meshes_screen,
image_size=image_size,
blur_radius=raster_settings.blur_radius,
faces_per_pixel=raster_settings.faces_per_pixel,
bin_size=raster_settings.bin_size,
max_faces_per_bin=raster_settings.max_faces_per_bin,
perspective_correct=raster_settings.perspective_correct,
)
vismask = (pix_to_face > -1).float()
D = attributes.shape[-1]
attributes = attributes.clone(); attributes = attributes.view(attributes.shape[0]*attributes.shape[1], 3, attributes.shape[-1])
N, H, W, K, _ = bary_coords.shape
mask = pix_to_face == -1
pix_to_face = pix_to_face.clone()
pix_to_face[mask] = 0
idx = pix_to_face.view(N * H * W * K, 1, 1).expand(N * H * W * K, 3, D)
pixel_face_vals = attributes.gather(0, idx).view(N, H, W, K, 3, D)
pixel_vals = (bary_coords[..., None] * pixel_face_vals).sum(dim=-2)
pixel_vals[mask] = 0 # Replace masked values in output.
pixel_vals = pixel_vals[:,:,:,0].permute(0,3,1,2)
pixel_vals = torch.cat([pixel_vals, vismask[:,:,:,0][:,None,:,:]], dim=1)
# print(image_size)
# import ipdb; ipdb.set_trace()
return pixel_vals
class SRenderY(nn.Module):
def __init__(self, image_size, obj_filename, uv_size=256, rasterizer_type='pytorch3d'):
super(SRenderY, self).__init__()
self.image_size = image_size
self.uv_size = uv_size
if rasterizer_type == 'pytorch3d':
self.rasterizer = Pytorch3dRasterizer(image_size)
self.uv_rasterizer = Pytorch3dRasterizer(uv_size)
verts, faces, aux = load_obj(obj_filename)
uvcoords = aux.verts_uvs[None, ...] # (N, V, 2)
uvfaces = faces.textures_idx[None, ...] # (N, F, 3)
faces = faces.verts_idx[None,...]
else:
NotImplementedError
# faces
dense_triangles = util.generate_triangles(uv_size, uv_size)
self.register_buffer('dense_faces', torch.from_numpy(dense_triangles).long()[None,:,:])
self.register_buffer('faces', faces)
self.register_buffer('raw_uvcoords', uvcoords)
# uv coords
uvcoords = torch.cat([uvcoords, uvcoords[:,:,0:1]*0.+1.], -1) #[bz, ntv, 3]
uvcoords = uvcoords*2 - 1; uvcoords[...,1] = -uvcoords[...,1]
face_uvcoords = util.face_vertices(uvcoords, uvfaces)
self.register_buffer('uvcoords', uvcoords)
self.register_buffer('uvfaces', uvfaces)
self.register_buffer('face_uvcoords', face_uvcoords)
# shape colors, for rendering shape overlay
colors = torch.tensor([180, 180, 180])[None, None, :].repeat(1, faces.max()+1, 1).float()/255.
face_colors = util.face_vertices(colors, faces)
self.register_buffer('face_colors', face_colors)
## SH factors for lighting
pi = np.pi
constant_factor = torch.tensor([1/np.sqrt(4*pi), ((2*pi)/3)*(np.sqrt(3/(4*pi))), ((2*pi)/3)*(np.sqrt(3/(4*pi))),\
((2*pi)/3)*(np.sqrt(3/(4*pi))), (pi/4)*(3)*(np.sqrt(5/(12*pi))), (pi/4)*(3)*(np.sqrt(5/(12*pi))),\
(pi/4)*(3)*(np.sqrt(5/(12*pi))), (pi/4)*(3/2)*(np.sqrt(5/(12*pi))), (pi/4)*(1/2)*(np.sqrt(5/(4*pi)))]).float()
self.register_buffer('constant_factor', constant_factor)
def forward(self, vertices, transformed_vertices, albedos, lights=None, light_type='point'):
'''
-- Texture Rendering
vertices: [batch_size, V, 3], vertices in world space, for calculating normals, then shading
transformed_vertices: [batch_size, V, 3], range:normalized to [-1,1], projected vertices in image space (that is aligned to the iamge pixel), for rasterization
albedos: [batch_size, 3, h, w], uv map
lights:
spherical homarnic: [N, 9(shcoeff), 3(rgb)]
points/directional lighting: [N, n_lights, 6(xyzrgb)]
light_type:
point or directional
'''
batch_size = vertices.shape[0]
## rasterizer near 0 far 100. move mesh so minz larger than 0
transformed_vertices[:,:,2] = transformed_vertices[:,:,2] + 10
# attributes
face_vertices = util.face_vertices(vertices, self.faces.expand(batch_size, -1, -1))
normals = util.vertex_normals(vertices, self.faces.expand(batch_size, -1, -1)); face_normals = util.face_vertices(normals, self.faces.expand(batch_size, -1, -1))
transformed_normals = util.vertex_normals(transformed_vertices, self.faces.expand(batch_size, -1, -1)); transformed_face_normals = util.face_vertices(transformed_normals, self.faces.expand(batch_size, -1, -1))
attributes = torch.cat([self.face_uvcoords.expand(batch_size, -1, -1, -1),
transformed_face_normals.detach(),
face_vertices.detach(),
face_normals],
-1)
# rasterize
rendering = self.rasterizer(transformed_vertices, self.faces.expand(batch_size, -1, -1), attributes)
####
# vis mask
alpha_images = rendering[:, -1, :, :][:, None, :, :].detach()
# albedo
uvcoords_images = rendering[:, :3, :, :]; grid = (uvcoords_images).permute(0, 2, 3, 1)[:, :, :, :2]
albedo_images = F.grid_sample(albedos, grid, align_corners=False)
# visible mask for pixels with positive normal direction
transformed_normal_map = rendering[:, 3:6, :, :].detach()
pos_mask = (transformed_normal_map[:, 2:, :, :] < -0.05).float()
# rasterize with colors, and we get the pixels rgb
mouth_mask = rendering[:, 5:6, :, :]
# print(mouth_mask.min(), mouth_mask.max())
mouth_mask = 1-torch.where(mouth_mask < -0.05, torch.Tensor([0]).float().cuda(), mouth_mask)
# print(mouth_mask.max(), mouth_mask.min())
# mouth_mask = (transformed_normal_map[:, 2:, :, :] < 0.15).float()
# shading
normal_images = rendering[:, 9:12, :, :]
if lights is not None:
if lights.shape[1] == 9:
shading_images = self.add_SHlight(normal_images, lights)
else:
if light_type=='point':
vertice_images = rendering[:, 6:9, :, :].detach()
shading = self.add_pointlight(vertice_images.permute(0,2,3,1).reshape([batch_size, -1, 3]), normal_images.permute(0,2,3,1).reshape([batch_size, -1, 3]), lights)
shading_images = shading.reshape([batch_size, albedo_images.shape[2], albedo_images.shape[3], 3]).permute(0,3,1,2)
else:
shading = self.add_directionlight(normal_images.permute(0,2,3,1).reshape([batch_size, -1, 3]), lights)
shading_images = shading.reshape([batch_size, albedo_images.shape[2], albedo_images.shape[3], 3]).permute(0,3,1,2)
images = albedo_images*shading_images
else:
images = albedo_images
shading_images = images.detach()*0.
outputs = {
'images': images*alpha_images,
'albedo_images': albedo_images*alpha_images,
'alpha_images': alpha_images,
'pos_mask': pos_mask,
'shading_images': shading_images,
'grid': grid,
'normals': normals,
'normal_images': normal_images*alpha_images,
'transformed_normals': transformed_normals,
'predicted_inner_mouth': mouth_mask
}
return outputs
def add_SHlight(self, normal_images, sh_coeff):
'''
sh_coeff: [bz, 9, 3]
'''
N = normal_images
sh = torch.stack([
N[:,0]*0.+1., N[:,0], N[:,1], \
N[:,2], N[:,0]*N[:,1], N[:,0]*N[:,2],
N[:,1]*N[:,2], N[:,0]**2 - N[:,1]**2, 3*(N[:,2]**2) - 1
],
1) # [bz, 9, h, w]
sh = sh*self.constant_factor[None,:,None,None]
shading = torch.sum(sh_coeff[:,:,:,None,None]*sh[:,:,None,:,:], 1) # [bz, 9, 3, h, w]
return shading
def add_pointlight(self, vertices, normals, lights):
'''
vertices: [bz, nv, 3]
lights: [bz, nlight, 6]
returns:
shading: [bz, nv, 3]
'''
light_positions = lights[:,:,:3]; light_intensities = lights[:,:,3:]
directions_to_lights = F.normalize(light_positions[:,:,None,:] - vertices[:,None,:,:], dim=3)
# normals_dot_lights = torch.clamp((normals[:,None,:,:]*directions_to_lights).sum(dim=3), 0., 1.)
normals_dot_lights = (normals[:,None,:,:]*directions_to_lights).sum(dim=3)
shading = normals_dot_lights[:,:,:,None]*light_intensities[:,:,None,:]
return shading.mean(1)
def add_directionlight(self, normals, lights):
'''
normals: [bz, nv, 3]
lights: [bz, nlight, 6]
returns:
shading: [bz, nv, 3]
'''
light_direction = lights[:,:,:3]; light_intensities = lights[:,:,3:]
directions_to_lights = F.normalize(light_direction[:,:,None,:].expand(-1,-1,normals.shape[1],-1), dim=3)
# normals_dot_lights = torch.clamp((normals[:,None,:,:]*directions_to_lights).sum(dim=3), 0., 1.)
# normals_dot_lights = (normals[:,None,:,:]*directions_to_lights).sum(dim=3)
normals_dot_lights = torch.clamp((normals[:,None,:,:]*directions_to_lights).sum(dim=3), 0., 1.)
shading = normals_dot_lights[:,:,:,None]*light_intensities[:,:,None,:]
return shading.mean(1)
def render_shape(self, vertices, transformed_vertices, colors = None, images=None, detail_normal_images=None,
lights=None, return_grid=False, uv_detail_normals=None, h=None, w=None, return_pos=False):
'''
-- rendering shape with detail normal map
'''
batch_size = vertices.shape[0]
# use these lights if rendering DAD model https://github.com/PinataFarms/DAD-3DHeads
# [
# [-1,-1,-1],
# [1,-1,-1],
# [-1,+1,-1],
# [1,+1,-1],
# [0,0,-1]
# ]
if lights is None:
light_positions = torch.tensor(
[
[-1,1,1],
[1,1,1],
[-1,-1,1],
[1,-1,1],
[0,0,1]
]
)[None,:,:].expand(batch_size, -1, -1).float()
light_intensities = torch.ones_like(light_positions).float()*1.7
lights = torch.cat((light_positions, light_intensities), 2).to(vertices.device)
transformed_vertices[:,:,2] = transformed_vertices[:,:,2] + 10
# Attributes
face_vertices = util.face_vertices(vertices, self.faces.expand(batch_size, -1, -1))
normals = util.vertex_normals(vertices, self.faces.expand(batch_size, -1, -1)); face_normals = util.face_vertices(normals, self.faces.expand(batch_size, -1, -1))
transformed_normals = util.vertex_normals(transformed_vertices, self.faces.expand(batch_size, -1, -1)); transformed_face_normals = util.face_vertices(transformed_normals, self.faces.expand(batch_size, -1, -1))
if colors is None:
colors = self.face_colors.expand(batch_size, -1, -1, -1)
attributes = torch.cat([colors,
transformed_face_normals.detach(),
face_vertices.detach(),
face_normals,
self.face_uvcoords.expand(batch_size, -1, -1, -1)],
-1)
# rasterize
# import ipdb; ipdb.set_trace()
rendering = self.rasterizer(transformed_vertices, self.faces.expand(batch_size, -1, -1), attributes, h, w)
####
alpha_images = rendering[:, -1, :, :][:, None, :, :].detach()
# albedo
albedo_images = rendering[:, :3, :, :]
# mask
transformed_normal_map = rendering[:, 3:6, :, :].detach()
pos_mask = (transformed_normal_map[:, 2:, :, :] < 0.15).float()
mouth_mask = rendering[:, 5:6, :, :]
# shading
normal_images = rendering[:, 9:12, :, :].detach()
vertice_images = rendering[:, 6:9, :, :].detach()
if detail_normal_images is not None:
normal_images = detail_normal_images
shading = self.add_directionlight(normal_images.permute(0,2,3,1).reshape([batch_size, -1, 3]), lights)
shading_images = shading.reshape([batch_size, albedo_images.shape[2], albedo_images.shape[3], 3]).permute(0,3,1,2).contiguous()
shaded_images = albedo_images*shading_images
alpha_images = alpha_images*pos_mask
if images is None:
shape_images = shaded_images*alpha_images + torch.zeros_like(shaded_images).to(vertices.device)*(1-alpha_images)
else:
shape_images = shaded_images*alpha_images + images*(1-alpha_images)
if return_grid:
uvcoords_images = rendering[:, 12:15, :, :];
grid = (uvcoords_images).permute(0, 2, 3, 1)[:, :, :, :2]
if return_pos:
return shape_images, normal_images, grid, alpha_images, mouth_mask
else:
return shape_images, normal_images, grid, alpha_images
else:
return shape_images
def world2uv(self, vertices):
'''
warp vertices from world space to uv space
vertices: [bz, V, 3]
uv_vertices: [bz, 3, h, w]
'''
batch_size = vertices.shape[0]
face_vertices = util.face_vertices(vertices, self.faces.expand(batch_size, -1, -1))
uv_vertices = self.uv_rasterizer(self.uvcoords.expand(batch_size, -1, -1), self.uvfaces.expand(batch_size, -1, -1), face_vertices)[:, :3]
return uv_vertices
@@ -0,0 +1,374 @@
import torch
''' Rotation Converter
Repre: euler angle(3), angle axis(3), rotation matrix(3x3), quaternion(4)
ref: https://kornia.readthedocs.io/en/v0.1.2/_modules/torchgeometry/core/conversions.html#
"pi",
"rad2deg",
"deg2rad",
# "angle_axis_to_rotation_matrix", batch_rodrigues
"rotation_matrix_to_angle_axis",
"rotation_matrix_to_quaternion",
"quaternion_to_angle_axis",
# "angle_axis_to_quaternion",
euler2quat_conversion_sanity_batch
ref: smplx/lbs
batch_rodrigues: axis angle -> matrix
#
'''
pi = torch.Tensor([3.14159265358979323846])
def rad2deg(tensor):
"""Function that converts angles from radians to degrees.
See :class:`~torchgeometry.RadToDeg` for details.
Args:
tensor (Tensor): Tensor of arbitrary shape.
Returns:
Tensor: Tensor with same shape as input.
Example:
>>> input = tgm.pi * torch.rand(1, 3, 3)
>>> output = tgm.rad2deg(input)
"""
if not torch.is_tensor(tensor):
raise TypeError("Input type is not a torch.Tensor. Got {}"
.format(type(tensor)))
return 180. * tensor / pi.to(tensor.device).type(tensor.dtype)
def deg2rad(tensor):
"""Function that converts angles from degrees to radians.
See :class:`~torchgeometry.DegToRad` for details.
Args:
tensor (Tensor): Tensor of arbitrary shape.
Returns:
Tensor: Tensor with same shape as input.
Examples::
>>> input = 360. * torch.rand(1, 3, 3)
>>> output = tgm.deg2rad(input)
"""
if not torch.is_tensor(tensor):
raise TypeError("Input type is not a torch.Tensor. Got {}"
.format(type(tensor)))
return tensor * pi.to(tensor.device).type(tensor.dtype) / 180.
######### to quaternion
def euler_to_quaternion(r):
x = r[..., 0]
y = r[..., 1]
z = r[..., 2]
z = z/2.0
y = y/2.0
x = x/2.0
cz = torch.cos(z)
sz = torch.sin(z)
cy = torch.cos(y)
sy = torch.sin(y)
cx = torch.cos(x)
sx = torch.sin(x)
quaternion = torch.zeros_like(r.repeat(1,2))[..., :4].to(r.device)
quaternion[..., 0] += cx*cy*cz - sx*sy*sz
quaternion[..., 1] += cx*sy*sz + cy*cz*sx
quaternion[..., 2] += cx*cz*sy - sx*cy*sz
quaternion[..., 3] += cx*cy*sz + sx*cz*sy
return quaternion
def rotation_matrix_to_quaternion(rotation_matrix, eps=1e-6):
"""Convert 3x4 rotation matrix to 4d quaternion vector
This algorithm is based on algorithm described in
https://github.com/KieranWynn/pyquaternion/blob/master/pyquaternion/quaternion.py#L201
Args:
rotation_matrix (Tensor): the rotation matrix to convert.
Return:
Tensor: the rotation in quaternion
Shape:
- Input: :math:`(N, 3, 4)`
- Output: :math:`(N, 4)`
Example:
>>> input = torch.rand(4, 3, 4) # Nx3x4
>>> output = tgm.rotation_matrix_to_quaternion(input) # Nx4
"""
if not torch.is_tensor(rotation_matrix):
raise TypeError("Input type is not a torch.Tensor. Got {}".format(
type(rotation_matrix)))
if len(rotation_matrix.shape) > 3:
raise ValueError(
"Input size must be a three dimensional tensor. Got {}".format(
rotation_matrix.shape))
# if not rotation_matrix.shape[-2:] == (3, 4):
# raise ValueError(
# "Input size must be a N x 3 x 4 tensor. Got {}".format(
# rotation_matrix.shape))
rmat_t = torch.transpose(rotation_matrix, 1, 2)
mask_d2 = rmat_t[:, 2, 2] < eps
mask_d0_d1 = rmat_t[:, 0, 0] > rmat_t[:, 1, 1]
mask_d0_nd1 = rmat_t[:, 0, 0] < -rmat_t[:, 1, 1]
t0 = 1 + rmat_t[:, 0, 0] - rmat_t[:, 1, 1] - rmat_t[:, 2, 2]
q0 = torch.stack([rmat_t[:, 1, 2] - rmat_t[:, 2, 1],
t0, rmat_t[:, 0, 1] + rmat_t[:, 1, 0],
rmat_t[:, 2, 0] + rmat_t[:, 0, 2]], -1)
t0_rep = t0.repeat(4, 1).t()
t1 = 1 - rmat_t[:, 0, 0] + rmat_t[:, 1, 1] - rmat_t[:, 2, 2]
q1 = torch.stack([rmat_t[:, 2, 0] - rmat_t[:, 0, 2],
rmat_t[:, 0, 1] + rmat_t[:, 1, 0],
t1, rmat_t[:, 1, 2] + rmat_t[:, 2, 1]], -1)
t1_rep = t1.repeat(4, 1).t()
t2 = 1 - rmat_t[:, 0, 0] - rmat_t[:, 1, 1] + rmat_t[:, 2, 2]
q2 = torch.stack([rmat_t[:, 0, 1] - rmat_t[:, 1, 0],
rmat_t[:, 2, 0] + rmat_t[:, 0, 2],
rmat_t[:, 1, 2] + rmat_t[:, 2, 1], t2], -1)
t2_rep = t2.repeat(4, 1).t()
t3 = 1 + rmat_t[:, 0, 0] + rmat_t[:, 1, 1] + rmat_t[:, 2, 2]
q3 = torch.stack([t3, rmat_t[:, 1, 2] - rmat_t[:, 2, 1],
rmat_t[:, 2, 0] - rmat_t[:, 0, 2],
rmat_t[:, 0, 1] - rmat_t[:, 1, 0]], -1)
t3_rep = t3.repeat(4, 1).t()
mask_c0 = mask_d2 * mask_d0_d1.float()
mask_c1 = mask_d2 * (1 - mask_d0_d1.float())
mask_c2 = (1 - mask_d2.float()) * mask_d0_nd1
mask_c3 = (1 - mask_d2.float()) * (1 - mask_d0_nd1.float())
mask_c0 = mask_c0.view(-1, 1).type_as(q0)
mask_c1 = mask_c1.view(-1, 1).type_as(q1)
mask_c2 = mask_c2.view(-1, 1).type_as(q2)
mask_c3 = mask_c3.view(-1, 1).type_as(q3)
q = q0 * mask_c0 + q1 * mask_c1 + q2 * mask_c2 + q3 * mask_c3
q /= torch.sqrt(t0_rep * mask_c0 + t1_rep * mask_c1 + # noqa
t2_rep * mask_c2 + t3_rep * mask_c3) # noqa
q *= 0.5
return q
# def angle_axis_to_quaternion(theta):
# batch_size = theta.shape[0]
# l1norm = torch.norm(theta + 1e-8, p=2, dim=1)
# angle = torch.unsqueeze(l1norm, -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
def angle_axis_to_quaternion(angle_axis: torch.Tensor) -> torch.Tensor:
"""Convert an angle axis to a quaternion.
Adapted from ceres C++ library: ceres-solver/include/ceres/rotation.h
Args:
angle_axis (torch.Tensor): tensor with angle axis.
Return:
torch.Tensor: tensor with quaternion.
Shape:
- Input: :math:`(*, 3)` where `*` means, any number of dimensions
- Output: :math:`(*, 4)`
Example:
>>> angle_axis = torch.rand(2, 4) # Nx4
>>> quaternion = tgm.angle_axis_to_quaternion(angle_axis) # Nx3
"""
if not torch.is_tensor(angle_axis):
raise TypeError("Input type is not a torch.Tensor. Got {}".format(
type(angle_axis)))
if not angle_axis.shape[-1] == 3:
raise ValueError("Input must be a tensor of shape Nx3 or 3. Got {}"
.format(angle_axis.shape))
# unpack input and compute conversion
a0: torch.Tensor = angle_axis[..., 0:1]
a1: torch.Tensor = angle_axis[..., 1:2]
a2: torch.Tensor = angle_axis[..., 2:3]
theta_squared: torch.Tensor = a0 * a0 + a1 * a1 + a2 * a2
theta: torch.Tensor = torch.sqrt(theta_squared)
half_theta: torch.Tensor = theta * 0.5
mask: torch.Tensor = theta_squared > 0.0
ones: torch.Tensor = torch.ones_like(half_theta)
k_neg: torch.Tensor = 0.5 * ones
k_pos: torch.Tensor = torch.sin(half_theta) / theta
k: torch.Tensor = torch.where(mask, k_pos, k_neg)
w: torch.Tensor = torch.where(mask, torch.cos(half_theta), ones)
quaternion: torch.Tensor = torch.zeros_like(angle_axis)
quaternion[..., 0:1] += a0 * k
quaternion[..., 1:2] += a1 * k
quaternion[..., 2:3] += a2 * k
return torch.cat([w, quaternion], dim=-1)
#### quaternion to
def quaternion_to_rotation_matrix(quat):
"""Convert quaternion coefficients to rotation matrix.
Args:
quat: size = [B, 4] 4 <===>(w, x, y, z)
Returns:
Rotation matrix corresponding to the quaternion -- size = [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 quaternion_to_angle_axis(quaternion: torch.Tensor):
"""Convert quaternion vector to angle axis of rotation. TODO: CORRECT
Adapted from ceres C++ library: ceres-solver/include/ceres/rotation.h
Args:
quaternion (torch.Tensor): tensor with quaternions.
Return:
torch.Tensor: tensor with angle axis of rotation.
Shape:
- Input: :math:`(*, 4)` where `*` means, any number of dimensions
- Output: :math:`(*, 3)`
Example:
>>> quaternion = torch.rand(2, 4) # Nx4
>>> angle_axis = tgm.quaternion_to_angle_axis(quaternion) # Nx3
"""
if not torch.is_tensor(quaternion):
raise TypeError("Input type is not a torch.Tensor. Got {}".format(
type(quaternion)))
if not quaternion.shape[-1] == 4:
raise ValueError("Input must be a tensor of shape Nx4 or 4. Got {}"
.format(quaternion.shape))
# unpack input and compute conversion
q1: torch.Tensor = quaternion[..., 1]
q2: torch.Tensor = quaternion[..., 2]
q3: torch.Tensor = quaternion[..., 3]
sin_squared_theta: torch.Tensor = q1 * q1 + q2 * q2 + q3 * q3
sin_theta: torch.Tensor = torch.sqrt(sin_squared_theta)
cos_theta: torch.Tensor = quaternion[..., 0]
two_theta: torch.Tensor = 2.0 * torch.where(
cos_theta < 0.0,
torch.atan2(-sin_theta, -cos_theta),
torch.atan2(sin_theta, cos_theta))
k_pos: torch.Tensor = two_theta / sin_theta
k_neg: torch.Tensor = 2.0 * torch.ones_like(sin_theta).to(quaternion.device)
k: torch.Tensor = torch.where(sin_squared_theta > 0.0, k_pos, k_neg)
angle_axis: torch.Tensor = torch.zeros_like(quaternion).to(quaternion.device)[..., :3]
angle_axis[..., 0] += q1 * k
angle_axis[..., 1] += q2 * k
angle_axis[..., 2] += q3 * k
return angle_axis
#### batch converter
def batch_euler2axis(r):
return quaternion_to_angle_axis(euler_to_quaternion(r))
def batch_euler2matrix(r):
return quaternion_to_rotation_matrix(euler_to_quaternion(r))
def batch_matrix2euler(rot_mats):
# Calculates rotation matrix to euler angles
# Careful for extreme cases of eular angles like [0.0, pi, 0.0]
### only y?
# TODO:
sy = torch.sqrt(rot_mats[:, 0, 0] * rot_mats[:, 0, 0] +
rot_mats[:, 1, 0] * rot_mats[:, 1, 0])
return torch.atan2(-rot_mats[:, 2, 0], sy)
def batch_matrix2axis(rot_mats):
return quaternion_to_angle_axis(rotation_matrix_to_quaternion(rot_mats))
def batch_axis2matrix(theta):
# angle axis to rotation matrix
# theta N x 3
# return quat2mat(quat)
# batch_rodrigues
return quaternion_to_rotation_matrix(angle_axis_to_quaternion(theta))
def batch_axis2euler(theta):
return batch_matrix2euler(batch_axis2matrix(theta))
def batch_axis2euler(r):
return rot_mat_to_euler(batch_rodrigues(r))
def batch_orth_proj(X, camera):
'''
X is N x num_pquaternion_to_angle_axisoints x 3
'''
camera = camera.clone().view(-1, 1, 3)
X_trans = X[:, :, :2] + camera[:, :, 1:]
X_trans = torch.cat([X_trans, X[:,:,2:]], 2)
Xn = (camera[:, :, 0:1] * X_trans)
return Xn
def batch_rodrigues(rot_vecs, epsilon=1e-8, dtype=torch.float32):
''' same as batch_matrix2axis
Calculates the rotation matrices for a batch of rotation vectors
Parameters
----------
rot_vecs: torch.tensor Nx3
array of N axis-angle vectors
Returns
-------
R: torch.tensor Nx3x3
The rotation matrices for the given axis-angle parameters
'''
batch_size = rot_vecs.shape[0]
device = rot_vecs.device
angle = torch.norm(rot_vecs + 1e-8, dim=1, keepdim=True)
rot_dir = rot_vecs / angle
cos = torch.unsqueeze(torch.cos(angle), dim=1)
sin = torch.unsqueeze(torch.sin(angle), dim=1)
# Bx1 arrays
rx, ry, rz = torch.split(rot_dir, 1, dim=1)
K = torch.zeros((batch_size, 3, 3), dtype=dtype, device=device)
zeros = torch.zeros((batch_size, 1), dtype=dtype, device=device)
K = torch.cat([zeros, -rz, ry, rz, zeros, -rx, -ry, rx, zeros], dim=1) \
.view((batch_size, 3, 3))
ident = torch.eye(3, dtype=dtype, device=device).unsqueeze(dim=0)
rot_mat = ident + sin * K + (1 - cos) * torch.bmm(K, K)
return rot_mat
@@ -0,0 +1,136 @@
'''
crop
for torch tensor
Given image, bbox(center, bboxsize)
return: cropped image, tform(used for transform the keypoint accordingly)
only support crop to squared images
'''
import torch
from kornia.geometry.transform.imgwarp import (
warp_perspective, get_perspective_transform, warp_affine
)
def points2bbox(points, points_scale=None):
if points_scale:
assert points_scale[0]==points_scale[1]
points = points.clone()
points[:,:,:2] = (points[:,:,:2]*0.5 + 0.5)*points_scale[0]
min_coords, _ = torch.min(points, dim=1)
xmin, ymin = min_coords[:, 0], min_coords[:, 1]
max_coords, _ = torch.max(points, dim=1)
xmax, ymax = max_coords[:, 0], max_coords[:, 1]
center = torch.stack([xmax + xmin, ymax + ymin], dim=-1) * 0.5
width = (xmax - xmin)
height = (ymax - ymin)
# Convert the bounding box to a square box
size = torch.max(width, height).unsqueeze(-1)
return center, size
def augment_bbox(center, bbox_size, scale=[1.0, 1.0], trans_scale=0.):
batch_size = center.shape[0]
trans_scale = (torch.rand([batch_size, 2], device=center.device)*2. -1.) * trans_scale
center = center + trans_scale*bbox_size # 0.5
scale = torch.rand([batch_size,1], device=center.device) * (scale[1] - scale[0]) + scale[0]
size = bbox_size*scale
return center, size
def crop_tensor(image, center, bbox_size, crop_size, interpolation = 'bilinear', align_corners=False):
''' for batch image
Args:
image (torch.Tensor): the reference tensor of shape BXHxWXC.
center: [bz, 2]
bboxsize: [bz, 1]
crop_size;
interpolation (str): Interpolation flag. Default: 'bilinear'.
align_corners (bool): mode for grid_generation. Default: False. See
https://pytorch.org/docs/stable/nn.functional.html#torch.nn.functional.interpolate for details
Returns:
cropped_image
tform
'''
dtype = image.dtype
device = image.device
batch_size = image.shape[0]
# points: top-left, top-right, bottom-right, bottom-left
src_pts = torch.zeros([4,2], dtype=dtype, device=device).unsqueeze(0).expand(batch_size, -1, -1).contiguous()
src_pts[:, 0, :] = center - bbox_size*0.5 # / (self.crop_size - 1)
src_pts[:, 1, 0] = center[:, 0] + bbox_size[:, 0] * 0.5
src_pts[:, 1, 1] = center[:, 1] - bbox_size[:, 0] * 0.5
src_pts[:, 2, :] = center + bbox_size * 0.5
src_pts[:, 3, 0] = center[:, 0] - bbox_size[:, 0] * 0.5
src_pts[:, 3, 1] = center[:, 1] + bbox_size[:, 0] * 0.5
DST_PTS = torch.tensor([[
[0, 0],
[crop_size - 1, 0],
[crop_size - 1, crop_size - 1],
[0, crop_size - 1],
]], dtype=dtype, device=device).expand(batch_size, -1, -1)
# estimate transformation between points
dst_trans_src = get_perspective_transform(src_pts, DST_PTS)
# simulate broadcasting
# dst_trans_src = dst_trans_src.expand(batch_size, -1, -1)
# warp images
cropped_image = warp_affine(
image, dst_trans_src[:, :2, :], (crop_size, crop_size),
flags=interpolation, align_corners=align_corners)
tform = torch.transpose(dst_trans_src, 2, 1)
# tform = torch.inverse(dst_trans_src)
return cropped_image, tform
class Cropper(object):
def __init__(self, crop_size, scale=[1,1], trans_scale = 0.):
self.crop_size = crop_size
self.scale = scale
self.trans_scale = trans_scale
def crop(self, image, points, points_scale=None):
# points to bbox
center, bbox_size = points2bbox(points.clone(), points_scale)
# argument bbox. TODO: add rotation?
center, bbox_size = augment_bbox(center, bbox_size, scale=self.scale, trans_scale=self.trans_scale)
# crop
cropped_image, tform = crop_tensor(image, center, bbox_size, self.crop_size)
return cropped_image, tform
def transform_points(self, points, tform, points_scale=None, normalize = True):
points_2d = points[:,:,:2]
#'input points must use original range'
if points_scale:
assert points_scale[0]==points_scale[1]
points_2d = (points_2d*0.5 + 0.5)*points_scale[0]
batch_size, n_points, _ = points.shape
trans_points_2d = torch.bmm(
torch.cat([points_2d, torch.ones([batch_size, n_points, 1], device=points.device, dtype=points.dtype)], dim=-1),
tform
)
trans_points = torch.cat([trans_points_2d[:,:,:2], points[:,:,2:]], dim=-1)
if normalize:
trans_points[:,:,:2] = trans_points[:,:,:2]/self.crop_size*2 - 1
return trans_points
def transform_points(points, tform, points_scale=None, out_scale=None):
points_2d = points[:,:,:2]
#'input points must use original range'
if points_scale:
assert points_scale[0]==points_scale[1]
points_2d = (points_2d*0.5 + 0.5)*points_scale[0]
# import ipdb; ipdb.set_trace()
batch_size, n_points, _ = points.shape
trans_points_2d = torch.bmm(
torch.cat([points_2d, torch.ones([batch_size, n_points, 1], device=points.device, dtype=points.dtype)], dim=-1),
tform
)
if out_scale: # h,w of output image size
trans_points_2d[:,:,0] = trans_points_2d[:,:,0]/out_scale[1]*2 - 1
trans_points_2d[:,:,1] = trans_points_2d[:,:,1]/out_scale[0]*2 - 1
trans_points = torch.cat([trans_points_2d[:,:,:2], points[:,:,2:]], dim=-1)
return trans_points
@@ -0,0 +1,310 @@
# -*- coding: utf-8 -*-
#
# Max-Planck-Gesellschaft zur Förderung der Wissenschaften e.V. (MPG) is
# holder of all proprietary rights on this computer program.
# Using this computer program means that you agree to the terms
# in the LICENSE file included with this software distribution.
# Any use not explicitly granted by the LICENSE is prohibited.
#
# Copyright©2019 Max-Planck-Gesellschaft zur Förderung
# der Wissenschaften e.V. (MPG). acting on behalf of its Max Planck Institute
# for Intelligent Systems. All rights reserved.
#
# For comments or questions, please email us at deca@tue.mpg.de
# For commercial licensing contact, please contact ps-license@tuebingen.mpg.de
import os, sys
import torch
import torchvision
import torch.nn.functional as F
import torch.nn as nn
from torch.utils.data import DataLoader
import numpy as np
from time import time
from skimage.io import imread
import cv2
import pickle
from loguru import logger
from datetime import datetime
from tqdm import tqdm
from .utils.renderer import SRenderY
from .models.encoders import ResnetEncoder
from .models.FLAME import FLAME, FLAMETex
from .models.decoders import Generator
from .utils import util
from .utils.rotation_converter import batch_euler2axis
from .datasets import datasets
from .utils.config import cfg
torch.backends.cudnn.benchmark = True
from .utils import lossfunc
from .datasets import build_datasets
class Trainer(object):
def __init__(self, model, config=None, device='cuda:0'):
if config is None:
self.cfg = cfg
else:
self.cfg = config
self.device = device
self.batch_size = self.cfg.dataset.batch_size
self.image_size = self.cfg.dataset.image_size
self.uv_size = self.cfg.model.uv_size
# deca model
self.deca = model
self.E_flame = self.deca.E_flame
self.flametex = self.deca.flametex
self.configure_optimizers()
self.load_checkpoint()
# initialize loss
self.mrf_loss = lossfunc.IDMRFLoss()
self.id_loss = lossfunc.VGGFace2Loss()
self.face_attr_mask = util.load_local_mask(image_size=self.cfg.model.uv_size, mode='bbx')
# intizalize loggers
logger.add(os.path.join(self.cfg.output_dir, self.cfg.train.log_dir, 'train.log'))
if self.cfg.train.write_summary:
from torch.utils.tensorboard import SummaryWriter
self.writer = SummaryWriter(log_dir=os.path.join(self.cfg.output_dir, self.cfg.train.log_dir))
def configure_optimizers(self):
self.opt = torch.optim.Adam(
self.E_flame.parameters(),
lr=self.cfg.train.lr,
amsgrad=False)
def load_checkpoint(self):
model_dict = self.deca.model_dict()
# resume training, including model weight, opt, steps
if self.cfg.train.resume and os.path.exists(os.path.join(self.cfg.output_dir, 'model.tar')):
checkpoint = torch.load(os.path.join(self.cfg.output_dir, 'model.tar'))
for key in model_dict.keys():
if key in checkpoint.keys():
util.copy_state_dict(model_dict[key], checkpoint[key])
util.copy_state_dict(self.opt.state_dict(), checkpoint['opt'])
self.global_step = checkpoint['global_step']
logger.info(f"resume training from {os.path.join(self.cfg.output_dir, 'model.tar')}")
logger.info(f"training start from step {self.global_step}")
# load model weights only
elif os.path.exists(self.cfg.pretrained_modelpath):
checkpoint = torch.load(self.cfg.pretrained_modelpath)
for key in model_dict.keys():
if key in checkpoint.keys():
util.copy_state_dict(model_dict[key], checkpoint[key])
self.global_step = 0
else:
logger.info('model path not found, start training from scratch')
self.global_step = 0
def training_step(self, batch, batch_nb):
self.deca.train()
# [B, K, 3, size, size] ==> [BxK, 3, size, size]
images = batch['image'].cuda(); images = images.view(-1, images.shape[-3], images.shape[-2], images.shape[-1])
lmk = batch['landmark'].cuda(); lmk = lmk.view(-1, lmk.shape[-2], lmk.shape[-1])
masks = batch['mask'].cuda(); masks = masks.view(-1, images.shape[-2], images.shape[-1])
#-- encoder
codedict = self.deca.encode(images)
### shape constraints
if self.cfg.loss.shape_consistency == 'exchange':
'''
make sure s0, s1 is something to make shape close
the difference from ||so - s1|| is
the later encourage s0, s1 is cloase in l2 space, but not really ensure shape will be close
'''
new_order = np.array([np.random.permutation(self.K) + i*self.K for i in range(self.batch_size)])
new_order = new_order.flatten()
shapecode = codedict['shape']
shapecode_new = shapecode[new_order]
codedict[key] = torch.cat([shapecode, shapecode_new], dim=0)
for key in ['tex', 'exp', 'pose', 'cam', 'light', 'images']:
code = codedict[key]
codedict[key] = torch.cat([code, code], dim=0)
## append gt
images = torch.cat([images, images], dim=0)# images = images.view(-1, images.shape[-3], images.shape[-2], images.shape[-1])
lmk = torch.cat([lmk, lmk], dim=0) #lmk = lmk.view(-1, lmk.shape[-2], lmk.shape[-1])
masks = torch.cat([masks, masks], dim=0)
# import ipdb; ipdb.set_trace()
batch_size = images.shape[0]
#-- decoder
opdict = self.deca.decode(codedict, vis_lmk=False, return_vis=False, use_detail=False)
#------ rendering
# mask
mask_face_eye = F.grid_sample(self.deca.uv_face_eye_mask.expand(batch_size,-1,-1,-1), opdict['grid'].detach(), align_corners=False)
# images
predicted_images = opdict['rendered_images']*mask_face_eye*opdict['alpha_images']
opdict['predicted_images'] = predicted_images
opdict['images'] = images
opdict['lmk'] = lmk
#### ----------------------- Losses
losses = {}
############################# base shape
predicted_landmarks = opdict['landmarks2d']
if self.cfg.loss.useWlmk:
losses['landmark'] = lossfunc.weighted_landmark_loss(predicted_landmarks, lmk)*self.cfg.loss.lmk
else:
losses['landmark'] = lossfunc.landmark_loss(predicted_landmarks, lmk)*self.cfg.loss.lmk
if self.cfg.loss.eyed > 0.:
losses['eye_distance'] = lossfunc.eyed_loss(predicted_landmarks, lmk)*self.cfg.loss.eyed
if self.cfg.loss.lipd > 0.:
losses['lip_distance'] = lossfunc.lipd_loss(predicted_landmarks, lmk)*self.cfg.loss.lipd
if self.cfg.loss.useSeg:
masks = masks[:,None,:,:]
else:
masks = mask_face_eye*opdict['alpha_images']
if self.cfg.loss.photo > 0.:
losses['photometric_texture'] = (masks*(predicted_images - images).abs()).mean()*self.cfg.loss.photo
if self.cfg.loss.id > 0.:
shading_images = self.deca.render.add_SHlight(opdict['normal_images'], codedict['light'].detach())
albedo_images = F.grid_sample(opdict['albedo'].detach(), opdict['grid'], align_corners=False)
overlay = albedo_images*shading_images*mask_face_eye + images*(1-mask_face_eye)
losses['identity'] = self.id_loss(overlay, images) * self.cfg.loss.id
losses['shape_reg'] = (torch.sum(codedict['shape']**2)/2)*self.cfg.loss.reg_shape
losses['expression_reg'] = (torch.sum(codedict['exp']**2)/2)*self.cfg.loss.reg_exp
losses['tex_reg'] = (torch.sum(codedict['tex']**2)/2)*self.cfg.loss.reg_tex
losses['light_reg'] = ((torch.mean(codedict['light'], dim=2)[:,:,None] - codedict['light'])**2).mean()*self.cfg.loss.reg_light
if self.cfg.model.jaw_type == 'euler':
# reg on jaw pose
losses['reg_jawpose_roll'] = (torch.sum(codedict['euler_jaw_pose'][:,-1]**2)/2)*10.
losses['reg_jawpose_close'] = (torch.sum(F.relu(-codedict['euler_jaw_pose'][:,0])**2)/2)*10.
#########################################################
all_loss = 0.
losses_key = losses.keys()
# losses_key = ['landmark', 'shape_reg', 'expression_reg']
for key in losses_key:
all_loss = all_loss + losses[key]
losses['all_loss'] = all_loss
return losses, opdict
def validation_step(self):
self.deca.eval()
try:
batch = next(self.val_iter)
except:
self.val_iter = iter(self.val_dataloader)
batch = next(self.val_iter)
images = batch['image'].cuda(); images = images.view(-1, images.shape[-3], images.shape[-2], images.shape[-1])
with torch.no_grad():
codedict = self.deca.encode(images)
opdict, visdict = self.deca.decode(codedict)
savepath = os.path.join(self.cfg.output_dir, self.cfg.train.val_vis_dir, f'{self.global_step:08}.jpg')
util.visualize_grid(visdict, savepath)
def evaluate(self):
''' NOW validation
'''
os.makedirs(os.path.join(self.cfg.output_dir, 'NOW_validation'), exist_ok=True)
savefolder = os.path.join(self.cfg.output_dir, 'NOW_validation', f'step_{self.global_step:08}')
os.makedirs(savefolder, exist_ok=True)
self.deca.eval()
# run now validation images
from .datasets.now import NoWDataset #, NoWVal_old
dataset = NoWDataset(scale=(self.cfg.dataset.scale_min + self.cfg.dataset.scale_max)/2)
dataloader = DataLoader(dataset, batch_size=8, shuffle=False,
num_workers=8,
pin_memory=True,
drop_last=False)
faces = self.deca.flame.faces_tensor.cpu().numpy()
for i, batch in enumerate(tqdm(dataloader, desc='now evaluation ')):
images = batch['image'].cuda()
imagename = batch['imagename']
with torch.no_grad():
codedict = self.deca.encode(images)
codedict['exp'][:] = 0.
codedict['pose'][:] = 0.
opdict, visdict = self.deca.decode(codedict)
#-- save results for evaluation
verts = opdict['verts'].cpu().numpy()
landmark_51 = opdict['landmarks3d_world'][:, 17:]
landmark_7 = landmark_51[:,[19, 22, 25, 28, 16, 31, 37]]
landmark_7 = landmark_7.cpu().numpy()
for k in range(images.shape[0]):
os.makedirs(os.path.join(savefolder, imagename[k]), exist_ok=True)
# save mesh
util.write_obj(os.path.join(savefolder, f'{imagename[k]}.obj'), vertices=verts[k], faces=faces)
# save 7 landmarks for alignment
np.save(os.path.join(savefolder, f'{imagename[k]}.npy'), landmark_7[k])
# visualize results to check
util.visualize_grid(visdict, os.path.join(savefolder, f'{i}.jpg'))
# exit()
## then please run main.py in https://github.com/soubhiksanyal/now_evaluation, it will take around 0.5h to get the metric results
def prepare_data(self):
self.train_dataset = build_datasets.build_train(self.cfg.dataset)
self.val_dataset = build_datasets.build_val(self.cfg.dataset)
logger.info('---- training data numbers: ', len(self.train_dataset))
self.train_dataloader = DataLoader(self.train_dataset, batch_size=self.batch_size, shuffle=True,
num_workers=self.cfg.dataset.num_workers,
pin_memory=True,
drop_last=True)
self.val_dataloader = DataLoader(self.val_dataset, batch_size=8, shuffle=True,
num_workers=8,
pin_memory=True,
drop_last=False)
self.val_iter = iter(self.val_dataloader)
def fit(self):
self.prepare_data()
iters_every_epoch = int(len(self.train_dataset)/self.batch_size)
start_epoch = self.global_step//iters_every_epoch
for epoch in tqdm(range(start_epoch, self.cfg.train.max_epochs)):
for step, batch in enumerate(tqdm(self.train_dataloader)):
losses, opdict = self.training_step(batch, step)
import ipdb; ipdb.set_trace()
if self.global_step % self.cfg.train.log_steps == 0:
loss_info = f"ExpName: {self.cfg.exp_name} \nEpoch: {epoch}, Iter: {step}/{iters_every_epoch}, Time: {datetime.now().strftime('%Y-%m-%d-%H:%M:%S')} \n"
for k, v in losses.items():
loss_info = loss_info + f'{k}: {v:.4f}, '
if self.cfg.train.write_summary:
self.writer.add_scalar('train_loss/'+k, v, global_step=self.global_step)
logger.info(loss_info)
if self.global_step % self.cfg.train.vis_steps == 0:
visind = list(range(8))
shape_images = self.deca.render.render_shape(opdict['verts'], opdict['trans_verts'])
# import ipdb; ipdb.set_trace()
visdict = {
'inputs': opdict['images'][visind],
'landmarks2d_gt': util.tensor_vis_landmarks(opdict['images'][visind], opdict['lmk'][visind], isScale=True),
'landmarks2d': util.tensor_vis_landmarks(opdict['images'][visind], opdict['landmarks2d'][visind], isScale=True),
'shape_images': shape_images[visind],
'predicted_images': opdict['predicted_images'][visind]
}
savepath = os.path.join(self.cfg.output_dir, self.cfg.train.vis_dir, f'{self.global_step:06}.jpg')
util.visualize_grid(visdict, savepath)
if self.global_step % self.cfg.train.checkpoint_steps == 0:
model_dict = self.deca.model_dict()
model_dict['opt'] = self.opt.state_dict()
model_dict['global_step'] = self.global_step
model_dict['batch_size'] = self.batch_size
torch.save(model_dict, os.path.join(self.cfg.output_dir, 'model' + '.tar'))
if self.global_step % self.cfg.train.val_steps == 0:
self.validation_step()
if self.global_step % self.cfg.train.eval_steps == 0:
self.evaluate()
all_loss = losses['all_loss']
self.opt.zero_grad(); all_loss.backward(); self.opt.step()
self.global_step += 1
if self.global_step > self.cfg.train.max_steps:
break
@@ -0,0 +1,715 @@
# -*- coding: utf-8 -*-
#
# Max-Planck-Gesellschaft zur Förderung der Wissenschaften e.V. (MPG) is
# holder of all proprietary rights on this computer program.
# Using this computer program means that you agree to the terms
# in the LICENSE file included with this software distribution.
# Any use not explicitly granted by the LICENSE is prohibited.
#
# Copyright©2019 Max-Planck-Gesellschaft zur Förderung
# der Wissenschaften e.V. (MPG). acting on behalf of its Max Planck Institute
# for Intelligent Systems. All rights reserved.
#
# For comments or questions, please email us at deca@tue.mpg.de
# For commercial licensing contact, please contact ps-license@tuebingen.mpg.de
import numpy as np
import torch
import torch.nn.functional as F
import math
from collections import OrderedDict
import os
from scipy.ndimage import morphology
from skimage.io import imsave
import cv2
import torchvision
def upsample_mesh(vertices, normals, faces, displacement_map, texture_map, dense_template):
''' Credit to Timo
upsampling coarse mesh (with displacment map)
vertices: vertices of coarse mesh, [nv, 3]
normals: vertex normals, [nv, 3]
faces: faces of coarse mesh, [nf, 3]
texture_map: texture map, [256, 256, 3]
displacement_map: displacment map, [256, 256]
dense_template:
Returns:
dense_vertices: upsampled vertices with details, [number of dense vertices, 3]
dense_colors: vertex color, [number of dense vertices, 3]
dense_faces: [number of dense faces, 3]
'''
img_size = dense_template['img_size']
dense_faces = dense_template['f']
x_coords = dense_template['x_coords']
y_coords = dense_template['y_coords']
valid_pixel_ids = dense_template['valid_pixel_ids']
valid_pixel_3d_faces = dense_template['valid_pixel_3d_faces']
valid_pixel_b_coords = dense_template['valid_pixel_b_coords']
pixel_3d_points = vertices[valid_pixel_3d_faces[:, 0], :] * valid_pixel_b_coords[:, 0][:, np.newaxis] + \
vertices[valid_pixel_3d_faces[:, 1], :] * valid_pixel_b_coords[:, 1][:, np.newaxis] + \
vertices[valid_pixel_3d_faces[:, 2], :] * valid_pixel_b_coords[:, 2][:, np.newaxis]
vertex_normals = normals
pixel_3d_normals = vertex_normals[valid_pixel_3d_faces[:, 0], :] * valid_pixel_b_coords[:, 0][:, np.newaxis] + \
vertex_normals[valid_pixel_3d_faces[:, 1], :] * valid_pixel_b_coords[:, 1][:, np.newaxis] + \
vertex_normals[valid_pixel_3d_faces[:, 2], :] * valid_pixel_b_coords[:, 2][:, np.newaxis]
pixel_3d_normals = pixel_3d_normals / np.linalg.norm(pixel_3d_normals, axis=-1)[:, np.newaxis]
displacements = displacement_map[y_coords[valid_pixel_ids].astype(int), x_coords[valid_pixel_ids].astype(int)]
dense_colors = texture_map[y_coords[valid_pixel_ids].astype(int), x_coords[valid_pixel_ids].astype(int)]
offsets = np.einsum('i,ij->ij', displacements, pixel_3d_normals)
dense_vertices = pixel_3d_points + offsets
return dense_vertices, dense_colors, dense_faces
# borrowed from https://github.com/YadiraF/PRNet/blob/master/utils/write.py
def write_obj(obj_name,
vertices,
faces,
colors=None,
texture=None,
uvcoords=None,
uvfaces=None,
inverse_face_order=False,
normal_map=None,
):
''' Save 3D face model with texture.
Ref: https://github.com/patrikhuber/eos/blob/bd00155ebae4b1a13b08bf5a991694d682abbada/include/eos/core/Mesh.hpp
Args:
obj_name: str
vertices: shape = (nver, 3)
colors: shape = (nver, 3)
faces: shape = (ntri, 3)
texture: shape = (uv_size, uv_size, 3)
uvcoords: shape = (nver, 2) max value<=1
'''
if os.path.splitext(obj_name)[-1] != '.obj':
obj_name = obj_name + '.obj'
mtl_name = obj_name.replace('.obj', '.mtl')
texture_name = obj_name.replace('.obj', '.png')
material_name = 'FaceTexture'
faces = faces.copy()
# mesh lab start with 1, python/c++ start from 0
faces += 1
if inverse_face_order:
faces = faces[:, [2, 1, 0]]
if uvfaces is not None:
uvfaces = uvfaces[:, [2, 1, 0]]
# write obj
with open(obj_name, 'w') as f:
# first line: write mtlib(material library)
# f.write('# %s\n' % os.path.basename(obj_name))
# f.write('#\n')
# f.write('\n')
if texture is not None:
f.write('mtllib %s\n\n' % os.path.basename(mtl_name))
# write vertices
if colors is None:
for i in range(vertices.shape[0]):
f.write('v {} {} {}\n'.format(vertices[i, 0], vertices[i, 1], vertices[i, 2]))
else:
for i in range(vertices.shape[0]):
f.write('v {} {} {} {} {} {}\n'.format(vertices[i, 0], vertices[i, 1], vertices[i, 2], colors[i, 0], colors[i, 1], colors[i, 2]))
# write uv coords
if texture is None:
for i in range(faces.shape[0]):
f.write('f {} {} {}\n'.format(faces[i, 2], faces[i, 1], faces[i, 0]))
else:
for i in range(uvcoords.shape[0]):
f.write('vt {} {}\n'.format(uvcoords[i,0], uvcoords[i,1]))
f.write('usemtl %s\n' % material_name)
# write f: ver ind/ uv ind
uvfaces = uvfaces + 1
for i in range(faces.shape[0]):
f.write('f {}/{} {}/{} {}/{}\n'.format(
# faces[i, 2], uvfaces[i, 2],
# faces[i, 1], uvfaces[i, 1],
# faces[i, 0], uvfaces[i, 0]
faces[i, 0], uvfaces[i, 0],
faces[i, 1], uvfaces[i, 1],
faces[i, 2], uvfaces[i, 2]
)
)
# write mtl
with open(mtl_name, 'w') as f:
f.write('newmtl %s\n' % material_name)
s = 'map_Kd {}\n'.format(os.path.basename(texture_name)) # map to image
f.write(s)
if normal_map is not None:
name, _ = os.path.splitext(obj_name)
normal_name = f'{name}_normals.png'
f.write(f'disp {normal_name}')
# out_normal_map = normal_map / (np.linalg.norm(
# normal_map, axis=-1, keepdims=True) + 1e-9)
# out_normal_map = (out_normal_map + 1) * 0.5
cv2.imwrite(
normal_name,
# (out_normal_map * 255).astype(np.uint8)[:, :, ::-1]
normal_map
)
cv2.imwrite(texture_name, texture)
## load obj, similar to load_obj from pytorch3d
def load_obj(obj_filename):
""" Ref: https://github.com/facebookresearch/pytorch3d/blob/25c065e9dafa90163e7cec873dbb324a637c68b7/pytorch3d/io/obj_io.py
Load a mesh from a file-like object.
"""
with open(obj_filename, 'r') as f:
lines = [line.strip() for line in f]
verts, uvcoords = [], []
faces, uv_faces = [], []
# startswith expects each line to be a string. If the file is read in as
# bytes then first decode to strings.
if lines and isinstance(lines[0], bytes):
lines = [el.decode("utf-8") for el in lines]
for line in lines:
tokens = line.strip().split()
if line.startswith("v "): # Line is a vertex.
vert = [float(x) for x in tokens[1:4]]
if len(vert) != 3:
msg = "Vertex %s does not have 3 values. Line: %s"
raise ValueError(msg % (str(vert), str(line)))
verts.append(vert)
elif line.startswith("vt "): # Line is a texture.
tx = [float(x) for x in tokens[1:3]]
if len(tx) != 2:
raise ValueError(
"Texture %s does not have 2 values. Line: %s" % (str(tx), str(line))
)
uvcoords.append(tx)
elif line.startswith("f "): # Line is a face.
# Update face properties info.
face = tokens[1:]
face_list = [f.split("/") for f in face]
for vert_props in face_list:
# Vertex index.
faces.append(int(vert_props[0]))
if len(vert_props) > 1:
if vert_props[1] != "":
# Texture index is present e.g. f 4/1/1.
uv_faces.append(int(vert_props[1]))
verts = torch.tensor(verts, dtype=torch.float32)
uvcoords = torch.tensor(uvcoords, dtype=torch.float32)
faces = torch.tensor(faces, dtype=torch.long); faces = faces.reshape(-1, 3) - 1
uv_faces = torch.tensor(uv_faces, dtype=torch.long); uv_faces = uv_faces.reshape(-1, 3) - 1
return (
verts,
uvcoords,
faces,
uv_faces
)
# ---------------------------- process/generate vertices, normals, faces
def generate_triangles(h, w, margin_x=2, margin_y=5, mask = None):
# quad layout:
# 0 1 ... w-1
# w w+1
#.
# w*h
triangles = []
for x in range(margin_x, w-1-margin_x):
for y in range(margin_y, h-1-margin_y):
triangle0 = [y*w + x, y*w + x + 1, (y+1)*w + x]
triangle1 = [y*w + x + 1, (y+1)*w + x + 1, (y+1)*w + x]
triangles.append(triangle0)
triangles.append(triangle1)
triangles = np.array(triangles)
triangles = triangles[:,[0,2,1]]
return triangles
# borrowed from https://github.com/daniilidis-group/neural_renderer/blob/master/neural_renderer/vertices_to_faces.py
def face_vertices(vertices, faces):
"""
:param vertices: [batch size, number of vertices, 3]
:param faces: [batch size, number of faces, 3]
:return: [batch size, number of faces, 3, 3]
"""
assert (vertices.ndimension() == 3)
assert (faces.ndimension() == 3)
assert (vertices.shape[0] == faces.shape[0])
assert (vertices.shape[2] == 3)
assert (faces.shape[2] == 3)
bs, nv = vertices.shape[:2]
bs, nf = faces.shape[:2]
device = vertices.device
faces = faces + (torch.arange(bs, dtype=torch.int32).to(device) * nv)[:, None, None]
vertices = vertices.reshape((bs * nv, 3))
# pytorch only supports long and byte tensors for indexing
return vertices[faces.long()]
def vertex_normals(vertices, faces):
"""
:param vertices: [batch size, number of vertices, 3]
:param faces: [batch size, number of faces, 3]
:return: [batch size, number of vertices, 3]
"""
assert (vertices.ndimension() == 3)
assert (faces.ndimension() == 3)
assert (vertices.shape[0] == faces.shape[0])
assert (vertices.shape[2] == 3)
assert (faces.shape[2] == 3)
bs, nv = vertices.shape[:2]
bs, nf = faces.shape[:2]
device = vertices.device
normals = torch.zeros(bs * nv, 3).to(device)
faces = faces + (torch.arange(bs, dtype=torch.int32).to(device) * nv)[:, None, None] # expanded faces
vertices_faces = vertices.reshape((bs * nv, 3))[faces.long()]
faces = faces.reshape(-1, 3)
vertices_faces = vertices_faces.reshape(-1, 3, 3)
normals.index_add_(0, faces[:, 1].long(),
torch.cross(vertices_faces[:, 2] - vertices_faces[:, 1], vertices_faces[:, 0] - vertices_faces[:, 1]))
normals.index_add_(0, faces[:, 2].long(),
torch.cross(vertices_faces[:, 0] - vertices_faces[:, 2], vertices_faces[:, 1] - vertices_faces[:, 2]))
normals.index_add_(0, faces[:, 0].long(),
torch.cross(vertices_faces[:, 1] - vertices_faces[:, 0], vertices_faces[:, 2] - vertices_faces[:, 0]))
normals = F.normalize(normals, eps=1e-6, dim=1)
normals = normals.reshape((bs, nv, 3))
# pytorch only supports long and byte tensors for indexing
return normals
def batch_orth_proj(X, camera):
''' orthgraphic projection
X: 3d vertices, [bz, n_point, 3]
camera: scale and translation, [bz, 3], [scale, tx, ty]
'''
camera = camera.clone().view(-1, 1, 3)
X_trans = X[:, :, :2] + camera[:, :, 1:]
X_trans = torch.cat([X_trans, X[:,:,2:]], 2)
shape = X_trans.shape
Xn = (camera[:, :, 0:1] * X_trans)
return Xn
# -------------------------------------- image processing
# borrowed from: https://torchgeometry.readthedocs.io/en/latest/_modules/kornia/filters
def gaussian(window_size, sigma):
def gauss_fcn(x):
return -(x - window_size // 2)**2 / float(2 * sigma**2)
gauss = torch.stack(
[torch.exp(torch.tensor(gauss_fcn(x))) for x in range(window_size)])
return gauss / gauss.sum()
def get_gaussian_kernel(kernel_size: int, sigma: float):
r"""Function that returns Gaussian filter coefficients.
Args:
kernel_size (int): filter size. It should be odd and positive.
sigma (float): gaussian standard deviation.
Returns:
Tensor: 1D tensor with gaussian filter coefficients.
Shape:
- Output: :math:`(\text{kernel_size})`
Examples::
>>> kornia.image.get_gaussian_kernel(3, 2.5)
tensor([0.3243, 0.3513, 0.3243])
>>> kornia.image.get_gaussian_kernel(5, 1.5)
tensor([0.1201, 0.2339, 0.2921, 0.2339, 0.1201])
"""
if not isinstance(kernel_size, int) or kernel_size % 2 == 0 or \
kernel_size <= 0:
raise TypeError("kernel_size must be an odd positive integer. "
"Got {}".format(kernel_size))
window_1d = gaussian(kernel_size, sigma)
return window_1d
def get_gaussian_kernel2d(kernel_size, sigma):
r"""Function that returns Gaussian filter matrix coefficients.
Args:
kernel_size (Tuple[int, int]): filter sizes in the x and y direction.
Sizes should be odd and positive.
sigma (Tuple[int, int]): gaussian standard deviation in the x and y
direction.
Returns:
Tensor: 2D tensor with gaussian filter matrix coefficients.
Shape:
- Output: :math:`(\text{kernel_size}_x, \text{kernel_size}_y)`
Examples::
>>> kornia.image.get_gaussian_kernel2d((3, 3), (1.5, 1.5))
tensor([[0.0947, 0.1183, 0.0947],
[0.1183, 0.1478, 0.1183],
[0.0947, 0.1183, 0.0947]])
>>> kornia.image.get_gaussian_kernel2d((3, 5), (1.5, 1.5))
tensor([[0.0370, 0.0720, 0.0899, 0.0720, 0.0370],
[0.0462, 0.0899, 0.1123, 0.0899, 0.0462],
[0.0370, 0.0720, 0.0899, 0.0720, 0.0370]])
"""
if not isinstance(kernel_size, tuple) or len(kernel_size) != 2:
raise TypeError("kernel_size must be a tuple of length two. Got {}"
.format(kernel_size))
if not isinstance(sigma, tuple) or len(sigma) != 2:
raise TypeError("sigma must be a tuple of length two. Got {}"
.format(sigma))
ksize_x, ksize_y = kernel_size
sigma_x, sigma_y = sigma
kernel_x = get_gaussian_kernel(ksize_x, sigma_x)
kernel_y = get_gaussian_kernel(ksize_y, sigma_y)
kernel_2d = torch.matmul(
kernel_x.unsqueeze(-1), kernel_y.unsqueeze(-1).t())
return kernel_2d
def gaussian_blur(x, kernel_size=(3,3), sigma=(0.8,0.8)):
b, c, h, w = x.shape
kernel = get_gaussian_kernel2d(kernel_size, sigma).to(x.device).to(x.dtype)
kernel = kernel.repeat(c, 1, 1, 1)
padding = [(k - 1) // 2 for k in kernel_size]
return F.conv2d(x, kernel, padding=padding, stride=1, groups=c)
def _compute_binary_kernel(window_size):
r"""Creates a binary kernel to extract the patches. If the window size
is HxW will create a (H*W)xHxW kernel.
"""
window_range = window_size[0] * window_size[1]
kernel: torch.Tensor = torch.zeros(window_range, window_range)
for i in range(window_range):
kernel[i, i] += 1.0
return kernel.view(window_range, 1, window_size[0], window_size[1])
def median_blur(x, kernel_size=(3,3)):
b, c, h, w = x.shape
kernel = _compute_binary_kernel(kernel_size).to(x.device).to(x.dtype)
kernel = kernel.repeat(c, 1, 1, 1)
padding = [(k - 1) // 2 for k in kernel_size]
features = F.conv2d(x, kernel, padding=padding, stride=1, groups=c)
features = features.view(b,c,-1,h,w)
median = torch.median(features, dim=2)[0]
return median
def get_laplacian_kernel2d(kernel_size: int):
r"""Function that returns Gaussian filter matrix coefficients.
Args:
kernel_size (int): filter size should be odd.
Returns:
Tensor: 2D tensor with laplacian filter matrix coefficients.
Shape:
- Output: :math:`(\text{kernel_size}_x, \text{kernel_size}_y)`
Examples::
>>> kornia.image.get_laplacian_kernel2d(3)
tensor([[ 1., 1., 1.],
[ 1., -8., 1.],
[ 1., 1., 1.]])
>>> kornia.image.get_laplacian_kernel2d(5)
tensor([[ 1., 1., 1., 1., 1.],
[ 1., 1., 1., 1., 1.],
[ 1., 1., -24., 1., 1.],
[ 1., 1., 1., 1., 1.],
[ 1., 1., 1., 1., 1.]])
"""
if not isinstance(kernel_size, int) or kernel_size % 2 == 0 or \
kernel_size <= 0:
raise TypeError("ksize must be an odd positive integer. Got {}"
.format(kernel_size))
kernel = torch.ones((kernel_size, kernel_size))
mid = kernel_size // 2
kernel[mid, mid] = 1 - kernel_size ** 2
kernel_2d: torch.Tensor = kernel
return kernel_2d
def laplacian(x):
# https://torchgeometry.readthedocs.io/en/latest/_modules/kornia/filters/laplacian.html
b, c, h, w = x.shape
kernel_size = 3
kernel = get_laplacian_kernel2d(kernel_size).to(x.device).to(x.dtype)
kernel = kernel.repeat(c, 1, 1, 1)
padding = (kernel_size - 1) // 2
return F.conv2d(x, kernel, padding=padding, stride=1, groups=c)
def angle2matrix(angles):
''' get rotation matrix from three rotation angles(degree). right-handed.
Args:
angles: [batch_size, 3] tensor containing X, Y, and Z angles.
x: pitch. positive for looking down.
y: yaw. positive for looking left.
z: roll. positive for tilting head right.
Returns:
R: [batch_size, 3, 3]. rotation matrices.
'''
angles = angles*(np.pi)/180.
s = torch.sin(angles)
c = torch.cos(angles)
cx, cy, cz = (c[:, 0], c[:, 1], c[:, 2])
sx, sy, sz = (s[:, 0], s[:, 1], s[:, 2])
zeros = torch.zeros_like(s[:, 0]).to(angles.device)
ones = torch.ones_like(s[:, 0]).to(angles.device)
# Rz.dot(Ry.dot(Rx))
R_flattened = torch.stack(
[
cz * cy, cz * sy * sx - sz * cx, cz * sy * cx + sz * sx,
sz * cy, sz * sy * sx + cz * cx, sz * sy * cx - cz * sx,
-sy, cy * sx, cy * cx,
],
dim=0) #[batch_size, 9]
R = torch.reshape(R_flattened, (-1, 3, 3)) #[batch_size, 3, 3]
return R
def binary_erosion(tensor, kernel_size=5):
# tensor: [bz, 1, h, w].
device = tensor.device
mask = tensor.cpu().numpy()
structure=np.ones((kernel_size,kernel_size))
new_mask = mask.copy()
for i in range(mask.shape[0]):
new_mask[i,0] = morphology.binary_erosion(mask[i,0], structure)
return torch.from_numpy(new_mask.astype(np.float32)).to(device)
def flip_image(src_image, kps):
'''
purpose:
flip a image given by src_image and the 2d keypoints
flip_mode:
0: horizontal flip
>0: vertical flip
<0: horizontal & vertical flip
'''
h, w = src_image.shape[0], src_image.shape[1]
src_image = cv2.flip(src_image, 1)
if kps is not None:
kps[:, 0] = w - 1 - kps[:, 0]
kp_map = [5, 4, 3, 2, 1, 0, 11, 10, 9, 8, 7, 6, 12, 13]
kps[:, :] = kps[kp_map]
return src_image, kps
# -------------------------------------- io
def copy_state_dict(cur_state_dict, pre_state_dict, prefix='', load_name=None):
def _get_params(key):
key = prefix + key
if key in pre_state_dict:
return pre_state_dict[key]
return None
for k in cur_state_dict.keys():
if load_name is not None:
if load_name not in k:
continue
v = _get_params(k)
try:
if v is None:
# print('parameter {} not found'.format(k))
continue
cur_state_dict[k].copy_(v)
except:
# print('copy param {} failed'.format(k))
continue
def check_mkdir(path):
if not os.path.exists(path):
print('creating %s' % path)
os.makedirs(path)
def check_mkdirlist(pathlist):
for path in pathlist:
if not os.path.exists(path):
print('creating %s' % path)
os.makedirs(path)
def tensor2video(tensor,gray=False):
video = tensor.detach().cpu().numpy()
video = video*255.
video = np.maximum(np.minimum(video, 255), 0)
if not gray:
video = video.transpose(0,2,3,1) #[:,:,[2,1,0]]
return video.astype(np.uint8).copy()
def tensor2image(tensor):
image = tensor.detach().cpu().numpy()
image = image*255.
image = np.maximum(np.minimum(image, 255), 0)
image = image.transpose(1,2,0)[:,:,[2,1,0]]
return image.astype(np.uint8).copy()
def dict2obj(d):
# if isinstance(d, list):
# d = [dict2obj(x) for x in d]
if not isinstance(d, dict):
return d
class C(object):
pass
o = C()
for k in d:
o.__dict__[k] = dict2obj(d[k])
return o
class Struct(object):
def __init__(self, **kwargs):
for key, val in kwargs.items():
setattr(self, key, val)
# original saved file with DataParallel
def remove_module(state_dict):
# create new OrderedDict that does not contain `module.`
new_state_dict = OrderedDict()
for k, v in state_dict.items():
name = k[7:] # remove `module.`
new_state_dict[name] = v
return new_state_dict
def dict_tensor2npy(tensor_dict):
npy_dict = {}
for key in tensor_dict:
npy_dict[key] = tensor_dict[key][0].cpu().numpy()
return npy_dict
# ---------------------------------- visualization
end_list = np.array([17, 22, 27, 42, 48, 31, 36, 68], dtype = np.int32) - 1
def plot_kpts(image, kpts, color = 'r'):
''' Draw 68 key points
Args:
image: the input image
kpt: (68, 3).
'''
if color == 'r':
c = (255, 0, 0)
elif color == 'g':
c = (0, 255, 0)
elif color == 'b':
c = (255, 0, 0)
image = image.copy()
kpts = kpts.copy()
radius = max(int(min(image.shape[0], image.shape[1])/200), 1)
radius = 1
for i in range(kpts.shape[0]):
st = kpts[i, :2]
if kpts.shape[1]==4:
if kpts[i, 3] > 0.5:
c = (0, 255, 0)
else:
c = (0, 0, 255)
if i in end_list:
continue
ed = kpts[i + 1, :2]
image = cv2.line(image, (int(st[0]), int(st[1])), (int(ed[0]), int(ed[1])), (255, 255, 255), radius)
image = cv2.circle(image,(int(st[0]), int(st[1])), radius, c, -1)
return image
def plot_verts(image, kpts, color = 'r'):
''' Draw 68 key points
Args:
image: the input image
kpt: (68, 3).
'''
if color == 'r':
c = (255, 0, 0)
elif color == 'g':
c = (0, 255, 0)
elif color == 'b':
c = (0, 0, 255)
elif color == 'y':
c = (0, 255, 255)
image = image.copy()
for i in range(kpts.shape[0]):
st = kpts[i, :2]
image = cv2.circle(image,(int(st[0]), int(st[1])), 1, c, 2)
return image
def tensor_vis_landmarks(images, landmarks, gt_landmarks=None, color = 'g', isScale=True):
# visualize landmarks
vis_landmarks = []
images = images.cpu().numpy()
predicted_landmarks = landmarks.detach().cpu().numpy()
if gt_landmarks is not None:
gt_landmarks_np = gt_landmarks.detach().cpu().numpy()
for i in range(images.shape[0]):
image = images[i]
image = image.transpose(1,2,0)[:,:,[2,1,0]].copy(); image = (image*255)
if isScale:
predicted_landmark = predicted_landmarks[i]
predicted_landmark[...,0] = predicted_landmark[...,0]*image.shape[1]/2 + image.shape[1]/2
predicted_landmark[...,1] = predicted_landmark[...,1]*image.shape[0]/2 + image.shape[0]/2
else:
predicted_landmark = predicted_landmarks[i]
if predicted_landmark.shape[0] == 68:
image_landmarks = plot_kpts(image, predicted_landmark, color)
if gt_landmarks is not None:
image_landmarks = plot_verts(image_landmarks, gt_landmarks_np[i]*image.shape[0]/2 + image.shape[0]/2, 'r')
else:
image_landmarks = plot_verts(image, predicted_landmark, color)
if gt_landmarks is not None:
image_landmarks = plot_verts(image_landmarks, gt_landmarks_np[i]*image.shape[0]/2 + image.shape[0]/2, 'r')
vis_landmarks.append(image_landmarks)
vis_landmarks = np.stack(vis_landmarks)
vis_landmarks = torch.from_numpy(vis_landmarks[:,:,:,[2,1,0]].transpose(0,3,1,2))/255.#, dtype=torch.float32)
return vis_landmarks
############### for training
def load_local_mask(image_size=256, mode='bbx'):
if mode == 'bbx':
# UV space face attributes bbx in size 2048 (l r t b)
# face = np.array([512, 1536, 512, 1536]) #
face = np.array([400, 1648, 400, 1648])
# if image_size == 512:
# face = np.array([400, 400+512*2, 400, 400+512*2])
# face = np.array([512, 512+512*2, 512, 512+512*2])
forehead = np.array([550, 1498, 430, 700+50])
eye_nose = np.array([490, 1558, 700, 1050+50])
mouth = np.array([574, 1474, 1050, 1550])
ratio = image_size / 2048.
face = (face * ratio).astype(np.int)
forehead = (forehead * ratio).astype(np.int)
eye_nose = (eye_nose * ratio).astype(np.int)
mouth = (mouth * ratio).astype(np.int)
regional_mask = np.array([face, forehead, eye_nose, mouth])
return regional_mask
def visualize_grid(visdict, savepath=None, size=224, dim=1, return_gird=True):
'''
image range should be [0,1]
dim: 2 for horizontal. 1 for vertical
'''
assert dim == 1 or dim==2
grids = {}
for key in visdict:
_,_,h,w = visdict[key].shape
if dim == 2:
new_h = size; new_w = int(w*size/h)
elif dim == 1:
new_h = int(h*size/w); new_w = size
grids[key] = torchvision.utils.make_grid(F.interpolate(visdict[key], [new_h, new_w]).detach().cpu(), nrow=12)
grid = torch.cat(list(grids.values()), dim)
grid_image = (grid.numpy().transpose(1,2,0).copy()*255)[:,:,[2,1,0]]
grid_image = np.minimum(np.maximum(grid_image, 0), 255).astype(np.uint8)
if savepath:
cv2.imwrite(savepath, grid_image)
if return_gird:
return grid_image
@@ -0,0 +1,54 @@
#! /usr/bin/env python
# -*- coding: utf-8 -*-
# Copyright 2021 Imperial College London (Pingchuan Ma)
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
import collections
from motiondiff_modules.spectre.external.ibug.face_detection import RetinaFacePredictor
from motiondiff_modules.spectre.external.ibug.face_alignment import FANPredictor
from .utils import get_landmarks
from .utils import extract_opencv_generator
class FaceTracker(object):
"""FaceTracker."""
def __init__(self, device="cuda:0"):
"""__init__.
:param device: str, contain the device on which a torch.Tensor is or will be allocated.
"""
# Create a RetinaFace detector using Resnet50 backbone.
self.face_detector = RetinaFacePredictor(
device=device,
threshold=0.8,
model=RetinaFacePredictor.get_model('resnet50')
)
# Create FAN for alignmentm, default model is '2dfan2'
alignment_weights = None
self.landmark_detector = FANPredictor(device=device, model=alignment_weights)
def tracker(self, filename):
"""tracker.
:param filename: str, the filename for the video
"""
face_info = collections.defaultdict(list)
frame_gen = extract_opencv_generator(filename)
while True:
try:
frame = frame_gen.__next__()
except StopIteration:
break
# -- face detection
detected_faces = self.face_detector(frame, rgb=False)
# -- face alignment
landmarks, scores = self.landmark_detector(frame, detected_faces, rgb=False)
face_info['bbox'].append(detected_faces)
face_info['landmarks'].append(landmarks)
face_info['landmarks_scores'].append(scores)
return get_landmarks(face_info)
@@ -0,0 +1,54 @@
#! /usr/bin/env python
# -*- coding: utf-8 -*-
# Copyright 2021 Imperial College London (Pingchuan Ma)
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
import cv2
def extract_opencv_generator(filename):
"""extract_opencv_generator.
:param filename: str, the filename for video.
"""
cap = cv2.VideoCapture(filename)
while(cap.isOpened()):
ret, frame = cap.read() # BGR
if ret:
yield frame
else:
break
cap.release()
def get_landmarks(multi_sub_landmarks):
"""get_landmarks.
:param multi_sub_landmarks: dict, a dictionary contains landmarks, bbox, and landmarks_scores.
"""
landmarks = [None] * len( multi_sub_landmarks["landmarks"])
for frame_idx in range(len(landmarks)):
if len(multi_sub_landmarks["landmarks"][frame_idx]) == 0:
continue
else:
# -- decide person id using maximal bounding box 0: Left, 1: top, 2: right, 3: bottom, probability
max_bbox_person_id = 0
max_bbox_len = multi_sub_landmarks["bbox"][frame_idx][max_bbox_person_id][2] + \
multi_sub_landmarks["bbox"][frame_idx][max_bbox_person_id][3] - \
multi_sub_landmarks["bbox"][frame_idx][max_bbox_person_id][0] - \
multi_sub_landmarks["bbox"][frame_idx][max_bbox_person_id][1]
landmark_scores = multi_sub_landmarks["landmarks_scores"][frame_idx][max_bbox_person_id]
for temp_person_id in range(1, len(multi_sub_landmarks["bbox"][frame_idx])):
temp_bbox_len = multi_sub_landmarks["bbox"][frame_idx][temp_person_id][2] + \
multi_sub_landmarks["bbox"][frame_idx][temp_person_id][3] - \
multi_sub_landmarks["bbox"][frame_idx][temp_person_id][0] - \
multi_sub_landmarks["bbox"][frame_idx][temp_person_id][1]
if temp_bbox_len > max_bbox_len:
max_bbox_person_id = temp_person_id
max_bbox_len = temp_bbox_len
landmark_scores = multi_sub_landmarks['landmarks_scores'][frame_idx][temp_person_id]
if landmark_scores[17:].min() >= 0.2:
landmarks[frame_idx] = multi_sub_landmarks["landmarks"][frame_idx][max_bbox_person_id]
return landmarks
@@ -0,0 +1,124 @@
import os
import cv2
import time
import numpy as np
import torch
from argparse import ArgumentParser
import sys
def extract(video, tmpl='%06d.jpg'):
os.makedirs(video.replace(".mp4", ""),exist_ok=True)
cmd = 'ffmpeg -i \"{}\" -threads 1 -q:v 0 \"{}/%06d.jpg\"'.format(video,
video.replace(".mp4", ""))
os.system(cmd)
# os.system("ffmpeg -i {} {} -y".format(videopath, videopath.replace(".mp4",".wav")))
# -*- coding: utf-8 -*-
import os, sys
import cv2
import numpy as np
from time import time
from scipy.io import savemat
import argparse
from tqdm import tqdm
import torch
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
from decalib.deca import DECA
from decalib.datasets import datasets
from decalib.utils import util
from decalib.utils.config import cfg as deca_cfg
import pickle
def video2sequence(video_path, videofolder):
os.makedirs(videofolder, exist_ok=True)
video_name = os.path.splitext(os.path.split(video_path)[-1])[0]
vidcap = cv2.VideoCapture(video_path)
success,image = vidcap.read()
count = 0
imagepath_list = []
while success:
imagepath = os.path.join(videofolder, f'{video_name}_frame{count:05d}.jpg')
cv2.imwrite(imagepath, image) # save frame as JPEG file
success,image = vidcap.read()
count += 1
imagepath_list.append(imagepath)
print('video frames are stored in {}'.format(videofolder))
return imagepath_list
from multiprocessing import Pool
from tqdm import tqdm
def main():
# Parse command-line arguments
parser = ArgumentParser()
root = "/gpu-data3/filby/LRS3/pretrain"
l = list(os.listdir("/gpu-data3/filby/LRS3/pretrain"))
test_list = []
for folder in l:
for file in os.listdir(os.path.join("/gpu-data3/filby/LRS3/pretrain",folder)):
if file.endswith(".txt"):
test_list.append([os.path.join("/gpu-data3/filby/LRS3/pretrain",folder,file.replace(".txt",".mp4")),os.path.join("/gpu-data3/filby/LRS3/pretrain",folder,file.replace(".txt",".mp4"))])
# print(test_list[0])
extract(test_list[0])
raise
p = Pool(12)
for _ in tqdm(p.imap_unordered(video2sequence, test_list), total=len(test_list)):
pass
main()
# import os
# import cv2
# import time
# import numpy as np
# import torch
# from argparse import ArgumentParser
#
# import sys
# sys.path.append("face_parsing")
#
#
# def extract_wav(videopath):
# # print(videopath)
#
# os.system("ffmpeg -i {} {} -y".format(videopath, videopath.replace("/videos/","/wavs/").replace(".mp4",".wav")))
#
# from multiprocessing import Pool
# from tqdm import tqdm
#
# def main():
# # Parse command-line arguments
# parser = ArgumentParser()
#
# root = "/gpu-data3/filby/MEAD/rendered/train/MEAD/videos"
#
# p = Pool(20)
#
# test_list = []
# for file in os.listdir(root):
# test_list.append(os.path.join(root,file))
#
# # print(test_list)
# # extract_wav(test_list[0])
# for _ in tqdm(p.imap_unordered(extract_wav, test_list), total=len(test_list)):
# pass
#
#
# main()
@@ -0,0 +1,52 @@
# -*- coding: utf-8 -*-
import os, sys
import cv2
import argparse
from tqdm import tqdm
from multiprocessing import Pool
def video2sequence(video_path):
videofolder = os.path.splitext(video_path)[0]
os.makedirs(videofolder, exist_ok=True)
vidcap = cv2.VideoCapture(video_path)
success,image = vidcap.read()
count = 0
imagepath_list = []
while success:
imagepath = os.path.join(videofolder, f'%06d.jpg'%count)
cv2.imwrite(imagepath, image) # save frame as JPEG file
success,image = vidcap.read()
count += 1
imagepath_list.append(imagepath)
print('video frames are stored in {}'.format(videofolder))
return videofolder
def extract_audio(video_path):
os.system("ffmpeg -i {} {} -y".format(video_path, video_path.replace(".mp4",".wav")))
def main(args):
video_list = []
for mode in ["trainval","test"]:
for folder in os.listdir(os.path.join(args.dataset_path,mode)):
for file in os.listdir(os.path.join(args.dataset_path,mode,folder)):
if file.endswith(".mp4"):
video_list.append(os.path.join(args.dataset_path,mode,folder,file))
p = Pool(12)
for _ in tqdm(p.imap_unordered(video2sequence, video_list), total=len(video_list)):
pass
for _ in tqdm(p.imap_unordered(extract_audio, video_list), total=len(video_list)):
pass
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('--dataset_path', default='./data/LRS3', type=str, help='path to dataset')
main(parser.parse_args())
@@ -0,0 +1,81 @@
import os
import cv2
import time
import numpy as np
import torch
from argparse import ArgumentParser
import sys
sys.path.append("face_parsing")
def extract_wav(videopath):
print(videopath)
os.system("ffmpeg -i {} {} -y".format(videopath, videopath.replace(".mp4",".wav")))
from multiprocessing import Pool
from tqdm import tqdm
def main():
# Parse command-line arguments
parser = ArgumentParser()
root = "/raid/gretsinas/LRS3/test"
p = Pool(12)
l = list(os.listdir("/raid/gretsinas/LRS3/test"))
test_list = []
for folder in l:
for file in os.listdir(os.path.join("/raid/gretsinas/LRS3/test",folder)):
if file.endswith(".txt"):
test_list.append(os.path.join("/raid/gretsinas/LRS3/test",folder,file.replace(".txt",".mp4")))
# print(test_list)
# extract_wav(test_list[0])
for _ in tqdm(p.imap_unordered(extract_wav, test_list), total=len(test_list)):
pass
main()
# import os
# import cv2
# import time
# import numpy as np
# import torch
# from argparse import ArgumentParser
#
# import sys
# sys.path.append("face_parsing")
#
#
# def extract_wav(videopath):
# # print(videopath)
#
# os.system("ffmpeg -i {} {} -y".format(videopath, videopath.replace("/videos/","/wavs/").replace(".mp4",".wav")))
#
# from multiprocessing import Pool
# from tqdm import tqdm
#
# def main():
# # Parse command-line arguments
# parser = ArgumentParser()
#
# root = "/gpu-data3/filby/MEAD/rendered/train/MEAD/videos"
#
# p = Pool(20)
#
# test_list = []
# for file in os.listdir(root):
# test_list.append(os.path.join(root,file))
#
# # print(test_list)
# # extract_wav(test_list[0])
# for _ in tqdm(p.imap_unordered(extract_wav, test_list), total=len(test_list)):
# pass
#
#
# main()
@@ -0,0 +1,114 @@
#! /usr/bin/env python
# -*- coding: utf-8 -*-
import os
import torch
from phonemizer.backend import EspeakBackend
from phonemizer.separator import Separator
separator = Separator(phone='-', word=' ')
backend = EspeakBackend('en-us', words_mismatch='ignore', with_stress=False)
import cv2
# phonemes to visemes map. this was created using Amazon Polly
# https://docs.aws.amazon.com/polly/latest/dg/polly-dg.pdf
def get_phoneme_to_viseme_map():
pho2vi = {}
# pho2vi_counts = {}
all_vis = []
p2v = "data/phonemes2visemes.csv"
with open(p2v) as file:
lines = file.readlines()
# for line in lines[2:29]+lines[30:50]:
for line in lines:
if line.split(",")[0] in pho2vi:
if line.split(",")[4].strip() != pho2vi[line.split(",")[0]]:
print('error')
pho2vi[line.split(",")[0]] = line.split(",")[4].strip()
all_vis.append(line.split(",")[4].strip())
# pho2vi_counts[line.split(",")[0]] = 0
return pho2vi, all_vis
pho2vi, all_vis = get_phoneme_to_viseme_map()
def convert_text_to_visemes(text):
phonemized = backend.phonemize([text], separator=separator)[0]
text = ""
for word in phonemized.split(" "):
visemized = []
for phoneme in word.split("-"):
if phoneme == "":
continue
try:
visemized.append(pho2vi[phoneme.strip()])
if pho2vi[phoneme.strip()] not in all_vis:
all_vis.append(pho2vi[phoneme.strip()])
# pho2vi_counts[phoneme.strip()] += 1
except:
print('Count not find', phoneme)
continue
text += " " + "".join(visemized)
return text
def save2avi(filename, data=None, fps=25):
"""save2avi. - function taken from Visual Speech Recognition repository
:param filename: str, the filename to save the video (.avi).
:param data: numpy.ndarray, the data to be saved.
:param fps: the chosen frames per second.
"""
assert data is not None, "data is {}".format(data)
os.makedirs(os.path.dirname(filename), exist_ok=True)
fourcc = cv2.VideoWriter_fourcc("F", "F", "V", "1")
writer = cv2.VideoWriter(filename, fourcc, fps, (data[0].shape[1], data[0].shape[0]), 0)
for frame in data:
writer.write(frame)
writer.release()
def predict_text(lipreader, mouth_sequence):
from external.Visual_Speech_Recognition_for_Multiple_Languages.espnet.asr.asr_utils import add_results_to_json
lipreader.model.eval()
with torch.no_grad():
enc_feats, _ = lipreader.model.encoder(mouth_sequence, None)
enc_feats = enc_feats.squeeze(0)
nbest_hyps = lipreader.beam_search(
x=enc_feats,
maxlenratio=lipreader.maxlenratio,
minlenratio=lipreader.minlenratio
)
nbest_hyps = [
h.asdict() for h in nbest_hyps[: min(len(nbest_hyps), lipreader.nbest)]
]
transcription = add_results_to_json(nbest_hyps, lipreader.char_list)
return transcription.replace("<eos>", "")
def predict_text_deca(lipreader, mouth_sequence):
from external.Visual_Speech_Recognition_for_Multiple_Languages.espnet.asr.asr_utils import add_results_to_json
lipreader.model.eval()
with torch.no_grad():
enc_feats, _ = lipreader.model.encoder(mouth_sequence, None)
enc_feats = enc_feats.squeeze(0)
ys_hat = lipreader.model.ctc.ctc_lo(enc_feats)
# print(ys_hat)
ys_hat = ys_hat.argmax(1)
ys_hat = torch.unique_consecutive(ys_hat, dim=-1)
ys = [lipreader.model.args.char_list[x] for x in ys_hat if x != 0]
ys = "".join(ys)
ys = ys.replace("<space>", " ")
return ys.replace("<eos>", "")
@@ -0,0 +1,172 @@
from argparse import Namespace
from fairseq import checkpoint_utils, tasks, utils
import os
from fairseq.dataclass.configs import GenerationConfig
from utils.lipread_utils import convert_text_to_visemes
from jiwer import wer, cer
# WARNING
# Run this file with additional command line arguments e.g. python apply_lip_read.py test test due to something stupid by fairseq
class AverageMeter(object):
"""Computes and stores the average and current value"""
def __init__(self):
self.reset()
def reset(self):
self.val = 0
self.avg = 0
self.sum = 0
self.count = 0
def update(self, val, n=1):
self.val = val
self.sum += val * n
self.count += n
self.avg = self.sum / self.count
def run_lipreading(videos, transcriptions):
"""
:param videos: list of videos
:param transcriptions: list of transcriptions
:return:
"""
ckpt_path = "../av_hubert/data/self_large_vox_433h.pt" # download this from https://facebookresearch.github.io/av_hubert/
utils.import_user_module(Namespace(user_dir='external/av_hubert/avhubert'))
modalities = ["video"]
gen_subset = "test"
gen_cfg = GenerationConfig(beam=1)
models, saved_cfg, task = checkpoint_utils.load_model_ensemble_and_task([ckpt_path])
models = [model.eval().cuda() for model in models]
saved_cfg.task.modalities = modalities
import cv2,tempfile
total_wer = AverageMeter()
total_cer = AverageMeter()
total_werv = AverageMeter()
total_cerv = AverageMeter()
for idx,video_path in enumerate(videos):
num_frames = int(cv2.VideoCapture(video_path).get(cv2.CAP_PROP_FRAME_COUNT))
data_dir = tempfile.mkdtemp()
tsv_cont = ["/\n", f"test-0\t{video_path}\t{None}\t{num_frames}\t{int(16_000*num_frames/25)}\n"]
label_cont = ["DUMMY\n"]
with open(f"{data_dir}/test.tsv", "w") as fo:
fo.write("".join(tsv_cont))
with open(f"{data_dir}/test.wrd", "w") as fo:
fo.write("".join(label_cont))
saved_cfg.task.data = data_dir
saved_cfg.task.label_dir = data_dir
task = tasks.setup_task(saved_cfg.task)
task.load_dataset(gen_subset, task_cfg=saved_cfg.task)
generator = task.build_generator(models, gen_cfg)
def decode_fn(x):
dictionary = task.target_dictionary
symbols_ignore = generator.symbols_to_strip_from_output
symbols_ignore.add(dictionary.pad())
return task.datasets[gen_subset].label_processors[0].decode(x, symbols_ignore)
itr = task.get_batch_iterator(dataset=task.dataset(gen_subset)).next_epoch_itr(shuffle=False)
sample = next(itr)
sample = utils.move_to_cuda(sample)
hypos = task.inference_step(generator, models, sample)
hypo = hypos[0][0]['tokens'].int().cpu()
hypo = decode_fn(hypo).upper()
groundtruth = transcriptions[idx].upper()
w = wer(groundtruth, hypo)
c = cer(groundtruth, hypo)
# ---------- convert to visemes -------- #
vg = convert_text_to_visemes(groundtruth)
v = convert_text_to_visemes(hypo)
print(hypo)
print(groundtruth)
print(v)
print(vg)
# -------------------------------------- #
wv = wer(vg, v)
cv = cer(vg, v)
total_wer.update(w)
total_cer.update(c)
total_werv.update(wv)
total_cerv.update(cv)
print(
f"progress: {idx + 1}/{len(videos)}\tcur WER: {total_wer.val * 100:.1f}\t"
f"cur CER: {total_cer.val * 100:.1f}\t"
f"count: {total_cer.count}\t"
f"avg WER: {total_wer.avg * 100:.1f}\tavg CER: {total_cer.avg * 100:.1f}\t"
f"avg WERV: {total_werv.avg * 100:.1f}\tavg CERV: {total_cerv.avg * 100:.1f}"
)
import glob
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--videos", type=str, required=True, help="path to videos (regex style)")
parser.add_argument("--LRS3_path", type=str, default="/gpu-data3/filby/LRS3", help="path to LRS3")
args = parser.parse_args()
video_list = glob.glob(args.videos)
assert len(video_list) > 0, "No videos found"
transcriptions = []
print('Found {} videos'.format(len(video_list)))
# LRS3
for video in video_list:
video_name = os.path.basename(video).replace("_mouth","")
# print(video_name)
subject = video_name.split("_")[1]
clip = video_name.split("_")[2].replace(".avi", ".txt")
text = open(os.path.join(f"/{args.LRS3_path}/test/", subject,clip)).readlines()[0].replace("Text:","").strip()
transcriptions.append(text)
# if running on TCDTIMIT uncomment the following:
# for video in video_list:
# video_name = os.path.basename(video).replace("_mouth", "")
# # print(video_name)
# subject = video_name.split("_")[0]
# clip = video_name.split(".")[0].split("_")[1].upper() + ".txt"
#
# text = open(os.path.join(f"/gpu-data3/filby/EAVTTS/TCDTIMITprocessing/downloadTCDTIMIT/volunteers", subject, 'Clips', 'straightcam', clip)).readlines()
# text = " ".join([x.split()[2].strip() for x in text])
#
# transcriptions.append(text)
# if running on MEAD uncomment the following:
# gt = open("data/list_full_mead_annotated.txt").readlines()
# gt_dic = {}
# for line in gt:
# gt_dic[line.split()[0]] = " ".join(line.split()[1:])
# for video in video_list:
# video_name = os.path.basename(video).replace("_mouth", "")
# # print(video_name)
# subject = video_name.split("_")[0]
# clip = video_name.split(".")[0].split("_")[1].upper() + ".txt"
#
# text = gt_dic[video_name.split(".")[0].replace("_mouth","")]
#
# transcriptions.append(text)
run_lipreading(video_list, transcriptions)