diff --git a/acestep/pipeline_ace_step.py b/acestep/pipeline_ace_step.py index 6f2cf24..7afd14f 100644 --- a/acestep/pipeline_ace_step.py +++ b/acestep/pipeline_ace_step.py @@ -823,6 +823,7 @@ class ACEStepPipeline: encoder_text_hidden_states_null=None, use_erg_lyric=False, use_erg_diffusion=False, + shift=3.0, retake_random_generators=None, retake_variance=0.5, add_retake_noise=False, @@ -831,7 +832,6 @@ class ACEStepPipeline: repaint_start=0, repaint_end=0, src_latents=None, - shift=3.0, ): logger.info( @@ -1308,7 +1308,7 @@ class ACEStepPipeline: sample=target_latents, return_dict=False, omega=omega_scale, - generator=random_generators, + generator=random_generators[0], )[0] if is_extend: diff --git a/acestep/schedulers/scheduling_flow_match_res_multistep.py b/acestep/schedulers/scheduling_flow_match_res_multistep.py index 6130c04..df28714 100644 --- a/acestep/schedulers/scheduling_flow_match_res_multistep.py +++ b/acestep/schedulers/scheduling_flow_match_res_multistep.py @@ -27,7 +27,7 @@ from diffusers.schedulers.scheduling_utils import SchedulerMixin logger = logging.get_logger(__name__) # pylint: disable=invalid-name -def get_ancestral_step(sigma_from, sigma_to, eta=1.): +def get_ancestral_step(sigma_from, sigma_to, eta=0.0): """Calculates the noise level (sigma_down) to step down to and the amount of noise to add (sigma_up) when doing an ancestral sampling step.""" if not eta: @@ -353,7 +353,7 @@ class FlowMatchResMultiStepScheduler(SchedulerMixin, ConfigMixin): x = sigma_fn(h) * x + h * (b1 * denoised + b2 * self.old_denoised) if self.sigmas[self.step_index + 1] > 0: - init_noise = torch.randn_like(x, device=x.device, generator=generator, dtype=x.dtype) + init_noise = torch.randn(x.size(), dtype=x.dtype, layout=x.layout, device=x.device, generator=generator) x = x + init_noise * s_noise * sigma_up self.old_denoised = denoised diff --git a/acestep/ui/components.py b/acestep/ui/components.py index 72bee1e..3de21a1 100644 --- a/acestep/ui/components.py +++ b/acestep/ui/components.py @@ -822,6 +822,7 @@ def create_text2music_ui( use_erg_tag, use_erg_lyric, use_erg_diffusion, + shift, oss_steps, guidance_scale_text, guidance_scale_lyric,