Merge pull request #15 from kijai/main

Add normals output option
This commit is contained in:
Fannovel16
2024-03-27 07:02:40 +07:00
committed by GitHub
4 changed files with 57 additions and 8 deletions
+9 -3
View File
@@ -11,12 +11,15 @@ if os.name == 'posix' and "DISPLAY" not in os.environ:
import torch
from mogen.smpl.simplify_loc2rot import joints2smpl
import pyrender
from pyrender.shader_program import ShaderProgramCache
from shapely import geometry
import trimesh
from pyrender.constants import RenderFlags
from comfy.model_management import get_torch_device
from tqdm import tqdm
shader_dir = os.path.join(os.path.dirname(__file__), 'shaders')
class WeakPerspectiveCamera(pyrender.Camera):
def __init__(self,
scale,
@@ -156,7 +159,7 @@ def render(motions):
out = np.stack(vid, axis=0)
return out
def render_from_smpl(thetas, yfov, move_x, move_y, move_z, draw_platform=True, depth_only=False, smpl_model_path=None):
def render_from_smpl(thetas, yfov, move_x, move_y, move_z, draw_platform=True, depth_only=False, normals=False, smpl_model_path=None):
rot2xyz = Rotation2xyz(device=get_torch_device(), smpl_model_path=smpl_model_path)
faces = rot2xyz.smpl_model.faces
@@ -262,12 +265,15 @@ def render_from_smpl(thetas, yfov, move_x, move_y, move_z, draw_platform=True, d
depth = r.render(scene, flags=RenderFlags.DEPTH_ONLY)
color = np.zeros([960, 960, 3])
else:
color, depth = r.render(scene, flags=RenderFlags.RGBA)
if normals:
r._renderer._program_cache = ShaderProgramCache(shader_dir=shader_dir)
color, depth = r.render(scene, flags=RenderFlags.RGBA)
# Image.fromarray(color).save(outdir+name+'_'+str(i)+'.png')
vid.append(color)
vid_depth.append(depth)
r.delete()
r = None
return np.stack(vid, axis=0), np.stack(vid_depth, axis=0)
+13
View File
@@ -0,0 +1,13 @@
#version 330 core
in vec3 frag_position;
in vec3 frag_normal;
out vec4 frag_color;
void main()
{
vec3 normal = normalize(frag_normal);
frag_color = vec4(normal * 0.5 + 0.5, 1.0);
}
+24
View File
@@ -0,0 +1,24 @@
#version 330 core
// Vertex Attributes
layout(location = 0) in vec3 position;
layout(location = NORMAL_LOC) in vec3 normal;
layout(location = INST_M_LOC) in mat4 inst_m;
// Uniforms
uniform mat4 M;
uniform mat4 V;
uniform mat4 P;
// Outputs
out vec3 frag_position;
out vec3 frag_normal;
void main()
{
gl_Position = P * V * M * inst_m * vec4(position, 1);
frag_position = vec3(M * inst_m * vec4(position, 1.0));
mat4 N = transpose(inverse(M * inst_m));
frag_normal = normalize(vec3(N * vec4(normal, 0.0)));
}
+11 -5
View File
@@ -90,18 +90,21 @@ class RenderSMPLMesh:
"move_y": ("FLOAT", {"default": 0,"min": -500, "max": 500, "step": 0.1}),
"move_z": ("FLOAT", {"default": 0,"min": -500, "max": 500, "step": 0.1}),
"background_hex_color": ("STRING", {"default": "#FFFFFF", "mutiline": False})
},
"optional": {
"normals": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("IMAGE", "IMAGE")
RETURN_TYPES = ("IMAGE", "IMAGE", "MASK")
RETURN_NAMES = ("IMAGE", "DEPTH_MAP")
CATEGORY = "MotionDiff/smpl"
FUNCTION = "render"
def render(self, smpl, yfov, move_x, move_y, move_z, draw_platform, depth_only, background_hex_color):
def render(self, smpl, yfov, move_x, move_y, move_z, draw_platform, depth_only, background_hex_color, normals=False):
smpl_model_path, thetas, _ = smpl
color_frames, depth_frames = render_from_smpl(
thetas.to(get_torch_device()),
yfov, move_x, move_y, move_z, draw_platform,depth_only,
yfov, move_x, move_y, move_z, draw_platform,depth_only, normals,
smpl_model_path=smpl_model_path
)
bg_color = ImageColor.getcolor(background_hex_color, "RGB")
@@ -112,7 +115,10 @@ class RenderSMPLMesh:
(color_frames[..., 2] == 1.)
]
color_frames[..., :3][white_mask] = torch.Tensor(bg_color)
white_mask_tensor = torch.stack(white_mask, dim=0)
white_mask_tensor = white_mask_tensor.float() / white_mask_tensor.max()
white_mask_tensor = 1.0 - white_mask_tensor.permute(1, 2, 3, 0).squeeze(dim=-1)
print(white_mask_tensor.shape)
#Normalize to [0, 1]
normalized_depth = (depth_frames - depth_frames.min()) / (depth_frames.max() - depth_frames.min())
#Pyrender's depths are the distance in meters to the camera, which is the inverse of depths in normal context
@@ -121,7 +127,7 @@ class RenderSMPLMesh:
#https://github.com/Fannovel16/comfyui_controlnet_aux/blob/main/src/controlnet_aux/util.py#L24
depth_frames = [torch.from_numpy(np.concatenate([x, x, x], axis=2)) for x in normalized_depth[..., None]]
depth_frames = torch.stack(depth_frames, dim=0)
return (color_frames, depth_frames,)
return (color_frames, depth_frames, white_mask_tensor,)
class SMPLLoader:
@classmethod