CoolFace
Apppublic

durgappc/infinitetalk

sourceHugging Faceapache-2.0updated 8mo agoView on Hugging Face
0likes
model_loader.py200 linesDownload Raw Back to utils
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