fix hidiffusion again

This commit is contained in:
Vladimir Mandic
2024-05-23 12:58:24 -04:00
parent 84137b1e17
commit 7d7eba74b7
2 changed files with 8 additions and 6 deletions
+1 -1
View File
@@ -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):
+7 -5
View File
@@ -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(