diff --git a/modules/hidiffusion/__init__.py b/modules/hidiffusion/__init__.py index 649af8ba8..795ecb75f 100644 --- a/modules/hidiffusion/__init__.py +++ b/modules/hidiffusion/__init__.py @@ -30,7 +30,7 @@ def apply_hidiffusion(p, model_type): hidiffusion.apply_hidiffusion(shared.sd_model, apply_raunet=shared.opts.hidiffusion_raunet, apply_window_attn=shared.opts.hidiffusion_attn, model_type=model_type) p.extra_generation_params['HiDiffusion'] = f'{shared.opts.hidiffusion_raunet}/{shared.opts.hidiffusion_attn}/{shared.opts.hidiffusion_steps > 0}:{shared.opts.hidiffusion_steps}' t1 = time.time() - shared.log.debug(f'HiDiffusion apply: raunet={shared.opts.hidiffusion_raunet} attn={shared.opts.hidiffusion_attn} aggressive={shared.opts.hidiffusion_steps > 0}:{shared.opts.hidiffusion_steps} t1={shared.opts.hidiffusion_t1} t2={shared.opts.hidiffusion_t2} time={t1-t0:.2f}') + shared.log.debug(f'HiDiffusion apply: raunet={shared.opts.hidiffusion_raunet} attn={shared.opts.hidiffusion_attn} aggressive={shared.opts.hidiffusion_steps > 0}:{shared.opts.hidiffusion_steps} t1={shared.opts.hidiffusion_t1} t2={shared.opts.hidiffusion_t2} time={t1-t0:.2f} type={shared.sd_model_type} width={p.width} height={p.height}') def remove_hidiffusion(p): diff --git a/modules/hidiffusion/hidiffusion.py b/modules/hidiffusion/hidiffusion.py index fbd394235..e24bce020 100644 --- a/modules/hidiffusion/hidiffusion.py +++ b/modules/hidiffusion/hidiffusion.py @@ -400,11 +400,12 @@ def make_diffusers_cross_attn_up_block(block_class: Type[torch.nn.Module]) -> Ty encoder_attention_mask: Optional[torch.FloatTensor] = None, ) -> torch.FloatTensor: - # TODO hidiffusion breaking hidden_shapes on 3-rd generate - if self.timestep == 0 and (hidden_states.shape[-1] != res_hidden_states_tuple[0].shape[-1] or hidden_states.shape[-2] != res_hidden_states_tuple[0].shape[-2]): - rescale = min(res_hidden_states_tuple[0].shape[-2] / hidden_states.shape[-2], res_hidden_states_tuple[0].shape[-1] / hidden_states.shape[-1]) - log.debug(f"HiDiffusion rescale: {hidden_states.shape} => {res_hidden_states_tuple[0].shape} scale={rescale}") - hidden_states = F.interpolate(hidden_states, scale_factor=rescale, mode='bicubic') + def fix_scale(first, second): # TODO hidiffusion breaks hidden_scale.shape on 3rd generate with sdxl + if (first.shape[-1] != second.shape[-1] or first.shape[-2] != second.shape[-2]): + rescale = min(second.shape[-2] / first.shape[-2], second.shape[-1] / first.shape[-1]) + # log.debug(f"HiDiffusion rescale: {hidden_states.shape} => {res_hidden_states_tuple[0].shape} scale={rescale}") + return F.interpolate(first, scale_factor=rescale, mode='bicubic') + return first self.max_timestep = self.info['pipeline']._num_timesteps ori_H, ori_W = self.info['size'] @@ -445,6 +446,7 @@ def make_diffusers_cross_attn_up_block(block_class: Type[torch.nn.Module]) -> Ty # pop res hidden states res_hidden_states = res_hidden_states_tuple[-1] res_hidden_states_tuple = res_hidden_states_tuple[:-1] + hidden_states = fix_scale(hidden_states, res_hidden_states) hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) hidden_states = resnet(hidden_states, temb) hidden_states = attn(