fix model auto download
This commit is contained in:
+71
-15
@@ -9,6 +9,7 @@ from loguru import logger
|
||||
from tqdm import tqdm
|
||||
import json
|
||||
import math
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
# from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
@@ -17,8 +18,6 @@ from diffusers.pipelines.stable_diffusion_3.pipeline_stable_diffusion_3 import r
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from transformers import UMT5EncoderModel, AutoTokenizer
|
||||
|
||||
from hf_download import download_repo
|
||||
|
||||
from language_segmentation import LangSegment
|
||||
from music_dcae.music_dcae_pipeline import MusicDCAE
|
||||
from models.ace_step_transformer import ACEStepTransformer2DModel
|
||||
@@ -56,16 +55,13 @@ REPO_ID = "ACE-Step/ACE-Step-v1-3.5B"
|
||||
# class ACEStepPipeline(DiffusionPipeline):
|
||||
class ACEStepPipeline:
|
||||
|
||||
def __init__(self, checkpoint_dir=None, device_id=0, dtype="bfloat16", text_encoder_checkpoint_path=None, torch_compile=False, **kwargs):
|
||||
# check checkpoint dir exist
|
||||
def __init__(self, checkpoint_dir=None, device_id=0, dtype="bfloat16", text_encoder_checkpoint_path=None, persistent_storage_path=None, torch_compile=False, **kwargs):
|
||||
if not checkpoint_dir:
|
||||
checkpoint_dir = os.path.join(os.path.dirname(__file__), "checkpoints")
|
||||
if not os.path.exists(checkpoint_dir):
|
||||
# huggingface download
|
||||
download_repo(
|
||||
repo_id=REPO_ID,
|
||||
save_path=checkpoint_dir
|
||||
)
|
||||
if persistent_storage_path is None:
|
||||
checkpoint_dir = os.path.join(os.path.dirname(__file__), "checkpoints")
|
||||
else:
|
||||
checkpoint_dir = os.path.join(persistent_storage_path, "checkpoints")
|
||||
ensure_directory_exists(checkpoint_dir)
|
||||
|
||||
self.checkpoint_dir = checkpoint_dir
|
||||
device = torch.device(f"cuda:{device_id}") if torch.cuda.is_available() else torch.device("cpu")
|
||||
@@ -76,12 +72,73 @@ class ACEStepPipeline:
|
||||
|
||||
def load_checkpoint(self, checkpoint_dir=None):
|
||||
device = self.device
|
||||
dcae_checkpoint_path = os.path.join(checkpoint_dir, "music_dcae_f8c8")
|
||||
vocoder_checkpoint_path = os.path.join(checkpoint_dir, "music_vocoder")
|
||||
|
||||
dcae_model_path = os.path.join(checkpoint_dir, "music_dcae_f8c8")
|
||||
vocoder_model_path = os.path.join(checkpoint_dir, "music_vocoder")
|
||||
ace_step_model_path = os.path.join(checkpoint_dir, "ace_step_transformer")
|
||||
text_encoder_model_path = os.path.join(checkpoint_dir, "umt5-base")
|
||||
|
||||
files_exist = (
|
||||
os.path.exists(os.path.join(dcae_model_path, "config.json")) and
|
||||
os.path.exists(os.path.join(dcae_model_path, "diffusion_pytorch_model.safetensors")) and
|
||||
os.path.exists(os.path.join(vocoder_model_path, "config.json")) and
|
||||
os.path.exists(os.path.join(vocoder_model_path, "diffusion_pytorch_model.safetensors")) and
|
||||
os.path.exists(os.path.join(ace_step_model_path, "config.json")) and
|
||||
os.path.exists(os.path.join(ace_step_model_path, "diffusion_pytorch_model.safetensors")) and
|
||||
os.path.exists(os.path.join(text_encoder_model_path, "config.json")) and
|
||||
os.path.exists(os.path.join(text_encoder_model_path, "model.safetensors")) and
|
||||
os.path.exists(os.path.join(text_encoder_model_path, "special_tokens_map.json")) and
|
||||
os.path.exists(os.path.join(text_encoder_model_path, "tokenizer_config.json")) and
|
||||
os.path.exists(os.path.join(text_encoder_model_path, "tokenizer.json"))
|
||||
)
|
||||
|
||||
if not files_exist:
|
||||
logger.info(f"Checkpoint directory {checkpoint_dir} is not complete, downloading from Hugging Face Hub")
|
||||
|
||||
# download music dcae model
|
||||
os.makedirs(dcae_model_path, exist_ok=True)
|
||||
hf_hub_download(repo_id=REPO_ID, subfolder="music_dcae_f8c8",
|
||||
filename="config.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
|
||||
hf_hub_download(repo_id=REPO_ID, subfolder="music_dcae_f8c8",
|
||||
filename="diffusion_pytorch_model.safetensors", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
|
||||
|
||||
# download vocoder model
|
||||
os.makedirs(vocoder_model_path, exist_ok=True)
|
||||
hf_hub_download(repo_id=REPO_ID, subfolder="music_vocoder",
|
||||
filename="config.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
|
||||
hf_hub_download(repo_id=REPO_ID, subfolder="music_vocoder",
|
||||
filename="diffusion_pytorch_model.safetensors", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
|
||||
|
||||
# download ace_step transformer model
|
||||
os.makedirs(ace_step_model_path, exist_ok=True)
|
||||
hf_hub_download(repo_id=REPO_ID, subfolder="ace_step_transformer",
|
||||
filename="config.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
|
||||
hf_hub_download(repo_id=REPO_ID, subfolder="ace_step_transformer",
|
||||
filename="diffusion_pytorch_model.safetensors", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
|
||||
|
||||
# download text encoder model
|
||||
os.makedirs(text_encoder_model_path, exist_ok=True)
|
||||
hf_hub_download(repo_id=REPO_ID, subfolder="umt5-base",
|
||||
filename="config.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
|
||||
hf_hub_download(repo_id=REPO_ID, subfolder="umt5-base",
|
||||
filename="model.safetensors", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
|
||||
hf_hub_download(repo_id=REPO_ID, subfolder="umt5-base",
|
||||
filename="special_tokens_map.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
|
||||
hf_hub_download(repo_id=REPO_ID, subfolder="umt5-base",
|
||||
filename="tokenizer_config.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
|
||||
hf_hub_download(repo_id=REPO_ID, subfolder="umt5-base",
|
||||
filename="tokenizer.json", local_dir=checkpoint_dir, local_dir_use_symlinks=False)
|
||||
|
||||
logger.info("Models downloaded")
|
||||
|
||||
dcae_checkpoint_path = dcae_model_path
|
||||
vocoder_checkpoint_path = vocoder_model_path
|
||||
ace_step_checkpoint_path = ace_step_model_path
|
||||
text_encoder_checkpoint_path = text_encoder_model_path
|
||||
|
||||
self.music_dcae = MusicDCAE(dcae_checkpoint_path=dcae_checkpoint_path, vocoder_checkpoint_path=vocoder_checkpoint_path)
|
||||
self.music_dcae.to(device).eval().to(self.dtype)
|
||||
|
||||
ace_step_checkpoint_path = os.path.join(checkpoint_dir, "ace_step_transformer")
|
||||
self.ace_step_transformer = ACEStepTransformer2DModel.from_pretrained(ace_step_checkpoint_path)
|
||||
self.ace_step_transformer.to(device).eval().to(self.dtype)
|
||||
|
||||
@@ -97,7 +154,6 @@ class ACEStepPipeline:
|
||||
])
|
||||
self.lang_segment = lang_segment
|
||||
self.lyric_tokenizer = VoiceBpeTokenizer()
|
||||
text_encoder_checkpoint_path = os.path.join(checkpoint_dir, "umt5-base")
|
||||
text_encoder_model = UMT5EncoderModel.from_pretrained(text_encoder_checkpoint_path).eval()
|
||||
text_encoder_model = text_encoder_model.to(device).to(self.dtype)
|
||||
text_encoder_model.requires_grad_(False)
|
||||
|
||||
Reference in New Issue
Block a user