Update model management and add DWPose offloading
This commit is contained in:
+25
-17
@@ -12,6 +12,8 @@ from fashn_human_parser import CATEGORY_TO_BODY_COVERAGE, FashnHumanParser
|
||||
from PIL import Image
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
import comfy.model_management
|
||||
|
||||
from .dwpose import DWposeDetector, draw_pose
|
||||
from .preprocessing import (
|
||||
BODY_COVERAGE_TO_FASHN_LABELS,
|
||||
@@ -59,20 +61,22 @@ class TryOnPipeline:
|
||||
def __init__(
|
||||
self,
|
||||
weights_dir: str,
|
||||
device: Optional[str] = None,
|
||||
logger: Optional[logging.Logger] = None,
|
||||
):
|
||||
self.weights_dir = os.path.abspath(weights_dir)
|
||||
self.logger = logger or setup_logger("TryOnPipeline", level=logging.INFO)
|
||||
|
||||
# Setup device
|
||||
self.device = torch.device(device if device else ("cuda" if torch.cuda.is_available() else "cpu"))
|
||||
self.logger.info(f"Using device: {self.device}")
|
||||
self.offload_device = torch.device("cpu")
|
||||
self.device = comfy.model_management.get_torch_device()
|
||||
|
||||
|
||||
# Setup inference dtype
|
||||
self.inference_dtype = torch.float32
|
||||
if self.device.type == "cuda" and torch.cuda.is_bf16_supported():
|
||||
if comfy.model_management.should_use_bf16():
|
||||
self.inference_dtype = torch.bfloat16
|
||||
elif comfy.model_management.should_use_fp16():
|
||||
self.inference_dtype = torch.float16
|
||||
self.logger.info(f"Using dtype: {self.inference_dtype}")
|
||||
|
||||
# Validate weights exist
|
||||
@@ -80,7 +84,6 @@ class TryOnPipeline:
|
||||
|
||||
# Load models
|
||||
self._setup_tryon_model()
|
||||
self._setup_pose_model()
|
||||
self._setup_hp_model()
|
||||
|
||||
# Setup transforms (derived from model input shape)
|
||||
@@ -117,30 +120,32 @@ class TryOnPipeline:
|
||||
self.logger.info(f"Loading TryOnModel from {model_path}")
|
||||
|
||||
self.tryon_model = TryOnModel()
|
||||
state_dict = load_checkpoint(model_path, device=str(self.device))
|
||||
state_dict = load_checkpoint(model_path, device="cpu")
|
||||
self.tryon_model.load_state_dict(state_dict)
|
||||
self.tryon_model.to(self.device, dtype=self.inference_dtype).eval()
|
||||
self.tryon_model.to(self.offload_device, dtype=self.inference_dtype).eval()
|
||||
|
||||
self.logger.info("TryOnModel loaded")
|
||||
|
||||
def _setup_pose_model(self):
|
||||
"""Load DWPose model."""
|
||||
def _get_dwpose_detector(self):
|
||||
"""
|
||||
Helper to init ONNX model on the fly.
|
||||
ONNX sessions are hard to move, so we create it when needed.
|
||||
"""
|
||||
dwpose_dir = os.path.join(self.weights_dir, "dwpose")
|
||||
self.logger.info(f"Loading DWPose from {dwpose_dir}")
|
||||
|
||||
dwpose_device = f"cuda:{self.device.index or 0}" if self.device.type == "cuda" else "cpu"
|
||||
self.pose_model = DWposeDetector(checkpoints_dir=dwpose_dir, device=dwpose_device)
|
||||
|
||||
self.logger.info("DWPose loaded")
|
||||
return DWposeDetector(checkpoints_dir=dwpose_dir, device=dwpose_device)
|
||||
|
||||
def _setup_hp_model(self):
|
||||
"""Load human parsing model."""
|
||||
self.logger.info("Loading FashnHumanParser")
|
||||
|
||||
hp_device = "cuda" if self.device.type == "cuda" else "cpu"
|
||||
self.hp_model = FashnHumanParser(device=hp_device)
|
||||
self.hp_model = FashnHumanParser(device="cpu")
|
||||
|
||||
self.logger.info("FashnHumanParser loaded")
|
||||
if hasattr(self.hp_model, "model"):
|
||||
self.hp_model.model.to(self.offload_device)
|
||||
|
||||
self.hp_model.device = self.offload_device
|
||||
self.logger.info(f"FashnHumanParser loaded on {self.offload_device}")
|
||||
|
||||
@torch.inference_mode()
|
||||
def _sample(
|
||||
@@ -243,6 +248,9 @@ class TryOnPipeline:
|
||||
Returns:
|
||||
PipelineOutput with `images` list containing generated PIL Images.
|
||||
"""
|
||||
|
||||
self.pose_model = self._get_dwpose_detector()
|
||||
|
||||
# Set seed
|
||||
torch.manual_seed(seed)
|
||||
if self.device.type == "cuda":
|
||||
|
||||
@@ -95,8 +95,9 @@ class FashnVtonInference:
|
||||
if hasattr(pipeline, "hp_model"):
|
||||
if hasattr(pipeline.hp_model, "model"):
|
||||
pipeline.hp_model.model.to(device)
|
||||
pipeline.hp_model.device = device
|
||||
|
||||
pbar = comfy.utils.ProgressBar(num_timesteps)
|
||||
pbar = comfy.utils.ProgressBar(steps)
|
||||
|
||||
def progress_callback(step, total_steps):
|
||||
pbar.update_absolute(step + 1, total_steps)
|
||||
@@ -135,6 +136,9 @@ class FashnVtonInference:
|
||||
if hasattr(pipeline, "hp_model"):
|
||||
if hasattr(pipeline.hp_model, "model"):
|
||||
pipeline.hp_model.model.to("cpu")
|
||||
pipeline.hp_model.device = device
|
||||
if hasattr(pipeline, "pose_model"):
|
||||
del pipeline.pose_model
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
Reference in New Issue
Block a user