Add SpectreFaceRecon and move vert normalization of Human4D to its internal
This commit is contained in:
@@ -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,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:
|
||||
"""
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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.
|
||||
@@ -0,0 +1,144 @@
|
||||
<div align="center">
|
||||
|
||||
# SPECTRE: Visual Speech-Aware Perceptual 3D Facial Expression Reconstruction from Videos
|
||||
|
||||
[](https://arxiv.org/abs/2207.11094)
|
||||
[](https://filby89.github.io/spectre/)
|
||||
<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},
|
||||
}
|
||||
```
|
||||
|
||||
|
||||
@@ -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.
Binary file not shown.
Binary file not shown.
File diff suppressed because it is too large
Load Diff
Binary file not shown.
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
|
||||
|
Binary file not shown.
Binary file not shown.
|
After Width: | Height: | Size: 11 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 12 KiB |
@@ -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)
|
||||
+168
@@ -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'
|
||||
+1
@@ -0,0 +1 @@
|
||||
from .retina_face_predictor import RetinaFacePredictor
|
||||
+332
@@ -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
|
||||
|
||||
|
||||
+41
@@ -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
|
||||
}
|
||||
+33
@@ -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
|
||||
+39
@@ -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
|
||||
+117
@@ -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
|
||||
+137
@@ -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
|
||||
Vendored
+110
@@ -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
|
||||
+70
@@ -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
|
||||
BIN
Binary file not shown.
+78
@@ -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
|
||||
+90
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user