durgappc/infinitetalk
0
1"""2Model Manager for InfiniteTalk3Handles lazy loading and caching of models from HuggingFace Hub4"""5 6import os7import torch8from huggingface_hub import snapshot_download9from pathlib import Path10import logging11 12logging.basicConfig(level=logging.INFO)13logger = logging.getLogger(__name__)14 15 16class ModelManager:17 """Manages model loading and caching"""18 19 def __init__(self, cache_dir=None):20 """21 Initialize Model Manager22 23 Args:24 cache_dir: Directory for caching models. Defaults to HF_HOME or /data/.huggingface25 """26 if cache_dir is None:27 cache_dir = os.environ.get("HF_HOME", "/data/.huggingface")28 29 self.cache_dir = Path(cache_dir)30 self.cache_dir.mkdir(parents=True, exist_ok=True)31 32 self.models = {}33 self.model_paths = {34 "wan": None,35 "infinitetalk": None,36 "wav2vec": None37 }38 39 def download_model(self, repo_id, subfolder=None, filename=None):40 """41 Download model from HuggingFace Hub with caching42 43 Args:44 repo_id: HuggingFace repository ID (e.g., "Kijai/WanVideo_comfy")45 subfolder: Optional subfolder within the repository46 filename: Optional specific file to download47 48 Returns:49 Path to downloaded model directory50 """51 try:52 logger.info(f"Downloading {repo_id} from HuggingFace Hub...")53 54 download_kwargs = {55 "repo_id": repo_id,56 "cache_dir": str(self.cache_dir),57 "resume_download": True,58 }59 60 if subfolder:61 download_kwargs["allow_patterns"] = f"{subfolder}/*"62 if filename:63 download_kwargs["allow_patterns"] = filename64 65 model_path = snapshot_download(**download_kwargs)66 67 if subfolder:68 model_path = os.path.join(model_path, subfolder)69 70 logger.info(f"Model downloaded successfully to {model_path}")71 return model_path72 73 except Exception as e:74 logger.error(f"Error downloading model {repo_id}: {e}")75 raise76 77 def get_wan_model_path(self):78 """Get or download Wan2.1 I2V model"""79 if self.model_paths["wan"] is None:80 logger.info("Downloading Wan2.1-I2V-14B-480P model...")81 # This will download the full model - adjust repo_id based on actual HF location82 self.model_paths["wan"] = self.download_model(83 repo_id="Kijai/WanVideo_comfy",84 subfolder="wan2_1_i2v_14B_480P"85 )86 return self.model_paths["wan"]87 88 def get_infinitetalk_weights_path(self):89 """Get or download InfiniteTalk weights"""90 if self.model_paths["infinitetalk"] is None:91 logger.info("Downloading InfiniteTalk weights...")92 self.model_paths["infinitetalk"] = self.download_model(93 repo_id="MeiGen-AI/InfiniteTalk",94 subfolder="single"95 )96 return self.model_paths["infinitetalk"]97 98 def get_wav2vec_model_path(self):99 """Get or download Wav2Vec2 audio encoder"""100 if self.model_paths["wav2vec"] is None:101 logger.info("Downloading Wav2Vec2 audio encoder...")102 self.model_paths["wav2vec"] = self.download_model(103 repo_id="TencentGameMate/chinese-wav2vec2-base"104 )105 return self.model_paths["wav2vec"]106 107 def load_wan_model(self, size="infinitetalk-480", device="cuda", offload_model=True):108 """109 Load Wan InfiniteTalk pipeline for inference110 111 Args:112 size: Model size configuration (infinitetalk-480 or infinitetalk-720)113 device: Device to load model on114 offload_model: Whether to offload model to CPU between forwards115 116 Returns:117 Loaded InfiniteTalkPipeline118 """119 if "wan_pipeline" not in self.models:120 import wan121 from wan.configs import WAN_CONFIGS122 123 model_path = self.get_wan_model_path()124 infinitetalk_path = self.get_infinitetalk_weights_path()125 infinitetalk_weights = os.path.join(infinitetalk_path, "infinitetalk.safetensors")126 127 logger.info(f"Loading InfiniteTalk pipeline from {model_path}...")128 129 # Get configuration for infinitetalk-14B130 task = "infinitetalk-14B"131 cfg = WAN_CONFIGS[task]132 133 # Create InfiniteTalk pipeline134 # This matches the initialization in generate_infinitetalk.py135 pipeline = wan.InfiniteTalkPipeline(136 config=cfg,137 checkpoint_dir=model_path,138 quant_dir=None, # No quantization for now139 device_id=device if isinstance(device, int) else 0,140 rank=0, # Single GPU141 t5_fsdp=False,142 dit_fsdp=False,143 use_usp=False,144 t5_cpu=False,145 lora_dir=None,146 lora_scales=None,147 quant=None,148 dit_path=None,149 infinitetalk_dir=infinitetalk_weights150 )151 152 # Enable memory management for low VRAM if needed153 # pipeline.enable_vram_management(num_persistent_param_in_dit=0)154 155 self.models["wan_pipeline"] = pipeline156 logger.info("InfiniteTalk pipeline loaded successfully")157 158 return self.models["wan_pipeline"]159 160 def load_audio_encoder(self, device="cuda"):161 """162 Load Wav2Vec2 audio encoder163 164 Args:165 device: Device to load model on166 167 Returns:168 Audio encoder model and feature extractor169 """170 if "audio_encoder" not in self.models:171 from transformers import Wav2Vec2FeatureExtractor172 from src.audio_analysis.wav2vec2 import Wav2Vec2Model173 174 wav2vec_path = self.get_wav2vec_model_path()175 176 logger.info(f"Loading audio encoder from {wav2vec_path}...")177 178 feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(wav2vec_path)179 audio_encoder = Wav2Vec2Model.from_pretrained(wav2vec_path)180 audio_encoder.to(device)181 audio_encoder.eval()182 183 self.models["audio_encoder"] = (audio_encoder, feature_extractor)184 logger.info("Audio encoder loaded successfully")185 186 return self.models["audio_encoder"]187 188 def unload_model(self, model_name):189 """Unload a specific model to free memory"""190 if model_name in self.models:191 del self.models[model_name]192 torch.cuda.empty_cache()193 logger.info(f"Unloaded {model_name}")194 195 def clear_all(self):196 """Unload all models"""197 self.models.clear()198 torch.cuda.empty_cache()199 logger.info("All models unloaded")200 