From fa633f1fc8b83747bdf944f2bbcdb3ec46db16b1 Mon Sep 17 00:00:00 2001 From: billwuhao Date: Fri, 9 May 2025 19:03:24 +0800 Subject: [PATCH] Fix invalid seed --- pipeline_ace_step.py | 39 +++++++++++++++++++++++++-------------- pyproject.toml | 2 +- 2 files changed, 26 insertions(+), 15 deletions(-) diff --git a/pipeline_ace_step.py b/pipeline_ace_step.py index d067c6a..39ba397 100644 --- a/pipeline_ace_step.py +++ b/pipeline_ace_step.py @@ -223,26 +223,37 @@ class ACEStepPipeline: return last_hidden_states def set_seeds(self, batch_size, manual_seeds=None): - seeds = None + processed_input_seeds = None if manual_seeds is not None: if isinstance(manual_seeds, str): if "," in manual_seeds: - seeds = list(map(int, manual_seeds.split(","))) + processed_input_seeds = list(map(int, manual_seeds.split(","))) elif manual_seeds.isdigit(): - seeds = int(manual_seeds) - - random_generators = [torch.Generator(device=self.device) for _ in range(batch_size)] + processed_input_seeds = int(manual_seeds) + elif isinstance(manual_seeds, list) and all(isinstance(s, int) for s in manual_seeds): + if len(manual_seeds) > 0: + processed_input_seeds = list(manual_seeds) + elif isinstance(manual_seeds, int): + processed_input_seeds = manual_seeds + random_generators = [ + torch.Generator(device=self.device) for _ in range(batch_size) + ] actual_seeds = [] for i in range(batch_size): - seed = None - if seeds is None: - seed = torch.randint(0, 2**32, (1,)).item() - if isinstance(seeds, int): - seed = seeds - if isinstance(seeds, list): - seed = seeds[i] - random_generators[i].manual_seed(seed) - actual_seeds.append(seed) + current_seed_for_generator = None + if processed_input_seeds is None: + current_seed_for_generator = torch.randint(0, 2**32, (1,)).item() + elif isinstance(processed_input_seeds, int): + current_seed_for_generator = processed_input_seeds + elif isinstance(processed_input_seeds, list): + if i < len(processed_input_seeds): + current_seed_for_generator = processed_input_seeds[i] + else: + current_seed_for_generator = processed_input_seeds[-1] + if current_seed_for_generator is None: + current_seed_for_generator = torch.randint(0, 2**32, (1,)).item() + random_generators[i].manual_seed(current_seed_for_generator) + actual_seeds.append(current_seed_for_generator) return random_generators, actual_seeds def get_lang(self, text): diff --git a/pyproject.toml b/pyproject.toml index 02eefe0..48fcd84 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "ace-step" description = "ACE-Step: A Step Towards Music Generation Foundation Model" -version = "1.0.3" +version = "1.0.4" license = {file = "LICENSE"} dependencies = ["diffusers", "librosa", "pypinyin", "pytorch_lightning", "soundfile", "hangul-romanize", "num2words", "spacy"]