From 6f55f6a3f2c1c9d4674f6d46d9aae23959e2e8b7 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 26 Jul 2026 14:40:27 +0300 Subject: [PATCH] Speedup model loading --- comfy/ldm/sam3d_body/model/model.py | 9 ++++----- comfy/ldm/sam3d_body/model/prompt.py | 9 +++++---- comfy_extras/sam3d_body/utils.py | 24 ++++++++++-------------- 3 files changed, 19 insertions(+), 23 deletions(-) diff --git a/comfy/ldm/sam3d_body/model/model.py b/comfy/ldm/sam3d_body/model/model.py index db5193366..f2a7b410d 100644 --- a/comfy/ldm/sam3d_body/model/model.py +++ b/comfy/ldm/sam3d_body/model/model.py @@ -5,6 +5,7 @@ import torch.nn as nn import torch.nn.functional as F import comfy.model_management +from comfy.ops import cast_to_input from comfy.ldm.sam3.sam import PositionEmbeddingRandom from comfy.image_encoders.dino3 import DINOV3_VITH_CONFIG, DINOv3ViTModel @@ -330,7 +331,7 @@ class SAM3DBody(nn.Module): """Append a token block from `embedding_weight` (+ zero-block in token_augment). Returns (token_embeddings, token_augment, start_idx).""" start_idx = token_embeddings.shape[1] - block = embedding_weight.to(token_embeddings)[None, :, :].repeat(batch_size, 1, 1) + block = cast_to_input(embedding_weight, token_embeddings)[None, :, :].repeat(batch_size, 1, 1) token_embeddings = torch.cat([token_embeddings, block], dim=1) token_augment = torch.cat([token_augment, torch.zeros_like(block)], dim=1) return token_embeddings, token_augment, start_idx @@ -388,10 +389,8 @@ class SAM3DBody(nn.Module): batch_size = image_embeddings.shape[0] - # .to(image_embeddings) moves weights CPU→GPU under dynamic loading - # (they stay on CPU until first use). - init_pose = init_pose_emb.weight.to(image_embeddings).expand(batch_size, -1).unsqueeze(1) - init_camera = init_camera_emb.weight.to(image_embeddings).expand(batch_size, -1).unsqueeze(1) + init_pose = cast_to_input(init_pose_emb.weight, image_embeddings).expand(batch_size, -1).unsqueeze(1) + init_camera = cast_to_input(init_camera_emb.weight, image_embeddings).expand(batch_size, -1).unsqueeze(1) init_estimate = torch.cat([init_pose, init_camera], dim=-1) # B x 1 x (404 + 3) init_input = torch.cat([condition_info.view(batch_size, 1, -1), init_estimate], dim=-1) diff --git a/comfy/ldm/sam3d_body/model/prompt.py b/comfy/ldm/sam3d_body/model/prompt.py index 5b2ec232d..2e7fafad6 100644 --- a/comfy/ldm/sam3d_body/model/prompt.py +++ b/comfy/ldm/sam3d_body/model/prompt.py @@ -11,6 +11,7 @@ import torch import torch.nn as nn from comfy.ldm.cascade.common import LayerNorm2d_op +from comfy.ops import cast_to_input from comfy.ldm.sam3.sam import PositionEmbeddingRandom from .transformer import TransformerDecoderLayer @@ -77,11 +78,11 @@ class PromptEncoder(nn.Module): weight_dtype = self.invalid_point_embed.weight.dtype point_embedding = self.pe_layer._encode(points.to(torch.float)).to(weight_dtype) point_embedding[labels == -2] = 0.0 # invalid points - point_embedding[labels == -2] += self.invalid_point_embed.weight.to(point_embedding) + point_embedding[labels == -2] += cast_to_input(self.invalid_point_embed.weight, point_embedding) point_embedding[labels == -1] = 0.0 - point_embedding[labels == -1] += self.not_a_point_embed.weight.to(point_embedding) + point_embedding[labels == -1] += cast_to_input(self.not_a_point_embed.weight, point_embedding) for i in range(self.num_body_joints): - point_embedding[labels == i] += self.point_embeddings[i].weight.to(point_embedding) + point_embedding[labels == i] += cast_to_input(self.point_embeddings[i].weight, point_embedding) point_mask = labels > -2 return point_embedding, point_mask @@ -136,7 +137,7 @@ class PromptEncoder(nn.Module): def get_mask_embeddings(self, masks: torch.Tensor, bs: int = 1, size: Tuple[int, int] = (16, 16)) -> torch.Tensor: """Embeds mask inputs. Caller casts both outputs to its working dtype.""" - no_mask_embeddings = self.no_mask_embed.weight.reshape(1, -1, 1, 1).expand(bs, -1, size[0], size[1]) + no_mask_embeddings = cast_to_input(self.no_mask_embed.weight, masks).reshape(1, -1, 1, 1).expand(bs, -1, size[0], size[1]) mask_embeddings = self.mask_downscaling(masks) return mask_embeddings, no_mask_embeddings diff --git a/comfy_extras/sam3d_body/utils.py b/comfy_extras/sam3d_body/utils.py index d8d0e9eb4..bcab02a14 100644 --- a/comfy_extras/sam3d_body/utils.py +++ b/comfy_extras/sam3d_body/utils.py @@ -471,27 +471,23 @@ def _compute_face_mask(model, disp_threshold_m: float = 1e-4) -> np.ndarray: device = head.scale_mean.device num_face = head.num_face_comps - zeros = lambda *s: torch.zeros(1, *s, device=device) - neutral_kw = dict( + # One batched forward over [neutral, axis_0 .. axis_71] + batch = num_face + 1 + zeros = lambda *s: torch.zeros(batch, *s, device=device) + expr = torch.zeros(batch, num_face, device=device) + expr[1:] = torch.eye(num_face, device=device) + verts = head.mhr_forward( global_trans=zeros(3), global_rot=zeros(3), body_pose_params=zeros(130), hand_pose_params=zeros(head.num_hand_comps * 2), scale_params=zeros(head.num_scale_comps), shape_params=zeros(head.num_shape_comps), - expr_params=zeros(num_face), - ) - v0 = head.mhr_forward(**neutral_kw).cpu().numpy()[0] # (Nv, 3) + expr_params=expr, + ).cpu().numpy() # (batch, Nv, 3) - face_mask = np.zeros(v0.shape[0], dtype=bool) - for axis in range(num_face): - expr = zeros(num_face) - expr[0, axis] = 1.0 - kw = dict(neutral_kw) - kw["expr_params"] = expr - v = head.mhr_forward(**kw).cpu().numpy()[0] - face_mask |= (np.linalg.norm(v - v0, axis=1) > disp_threshold_m) - return face_mask + disp = np.linalg.norm(verts[1:] - verts[0][None], axis=2) # (num_face, Nv) + return (disp > disp_threshold_m).any(axis=0) def jet_colormap(s: np.ndarray) -> np.ndarray: