From ba49c5e283d6ed63b72837aee409bcb5831faf41 Mon Sep 17 00:00:00 2001 From: smthemex <138738845+smthemex@users.noreply.github.com> Date: Fri, 29 May 2026 19:23:09 +0800 Subject: [PATCH] Update __init__.py --- .../longcat_video/audio_process/__init__.py | 18 +++++++++++++----- 1 file changed, 13 insertions(+), 5 deletions(-) diff --git a/LongCat_Video/longcat_video/audio_process/__init__.py b/LongCat_Video/longcat_video/audio_process/__init__.py index f98cf6c..cb5ed45 100644 --- a/LongCat_Video/longcat_video/audio_process/__init__.py +++ b/LongCat_Video/longcat_video/audio_process/__init__.py @@ -1,17 +1,25 @@ from .wav2vec2 import Wav2Vec2ModelWrapper -from transformers import Wav2Vec2FeatureExtractor, WhisperModel, AutoFeatureExtractor - -def get_audio_encoder(checkpoint_path, model_type="avatar-v1.0"): +from transformers import Wav2Vec2FeatureExtractor, WhisperModel, AutoFeatureExtractor,WhisperConfig +import os +import torch +from safetensors.torch import load_file as safe_load_file +def get_audio_encoder(checkpoint_path, model_type="avatar-v1.0",config_dir=""): if model_type == "avatar-v1.0": model = Wav2Vec2ModelWrapper(checkpoint_path) model.feature_extractor._freeze_parameters() return model if model_type == "avatar-v1.5": - model = WhisperModel.from_pretrained(checkpoint_path).eval() + #model = WhisperModel.from_pretrained(checkpoint_path).eval() + config_=WhisperConfig.from_pretrained(config_dir) + model=WhisperModel(config_) + sd=torch.load(checkpoint_path) if checkpoint_path.endswith(".pt") else safe_load_file(checkpoint_path) + model.load_state_dict(sd, strict=False) + del sd + model.eval() model.requires_grad_(False) return model -def get_audio_feature_extractor(checkpoint_path, model_type="avatar-v1.0"): +def get_audio_feature_extractor(checkpoint_path, model_type="avatar-v1.0",): if model_type == "avatar-v1.0": return Wav2Vec2FeatureExtractor(checkpoint_path, local_files_only=True) if model_type == "avatar-v1.5":