handle xformers

This commit is contained in:
Vladimir Mandic
2023-10-08 08:10:34 -04:00
parent 564d04d9f4
commit dca4efb3ad
5 changed files with 26 additions and 13 deletions
+1 -1
View File
@@ -643,7 +643,7 @@ def process_images(p: StableDiffusionProcessing) -> Processed:
if not shared.opts.cuda_compile:
modules.sd_models.apply_token_merging(p.sd_model, p.get_token_merging_ratio())
modules.sd_hijack_freeu.apply_freeu(p.sd_model, shared.backend == shared.Backend.ORIGINAL)
modules.sd_hijack_freeu.apply_freeu(p, shared.backend == shared.Backend.ORIGINAL)
if shared.cmd_opts.profile:
"""
+6 -4
View File
@@ -126,11 +126,12 @@ def ratio_to_region(width: float, offset: float, n: int):
return round(start), round(end), inverted
def apply_freeu(model, backend_original):
def apply_freeu(p, backend_original):
global state_enabled # pylint: disable=global-statement
global cat_original # pylint: disable=global-statement
if backend_original:
if opts.freeu_enabled:
p.extra_generation_params['FreeU'] = f'b1={opts.freeu_b1} b2={opts.freeu_b2} s1={opts.freeu_s1} s2={opts.freeu_s2}'
if not state_enabled: # otherwise already patched
cat_original = th.cat
th.cat = functools.partial(free_u_cat_hijack, original_function=th.cat)
@@ -139,12 +140,13 @@ def apply_freeu(model, backend_original):
if cat_original is not None:
th.cat = cat_original
state_enabled = False
elif hasattr(model, 'enable_freeu'):
elif hasattr(p.sd_model, 'enable_freeu'):
if opts.freeu_enabled:
model.enable_freeu(s1=opts.freeu_s1, s2=opts.freeu_s2, b1=opts.freeu_b1, b2=opts.freeu_b2)
p.extra_generation_params['FreeU'] = f'b1={opts.freeu_b1} b2={opts.freeu_b2} s1={opts.freeu_s1} s2={opts.freeu_s2}'
p.sd_model.enable_freeu(s1=opts.freeu_s1, s2=opts.freeu_s2, b1=opts.freeu_b1, b2=opts.freeu_b2)
state_enabled = True
elif state_enabled:
model.disable_freeu()
p.sd_model.disable_freeu()
state_enabled = False
if opts.freeu_enabled:
log.info(f'Applying free-u: b1={opts.freeu_b1} b2={opts.freeu_b2} s1={opts.freeu_s1} s2={opts.freeu_s2}')
+2
View File
@@ -142,6 +142,7 @@ def context_hypertile_vae(p):
return nullcontext()
else:
shared.log.info(f'Applying hypertile: vae={shared.opts.hypertile_vae_tile}')
p.extra_generation_params['Hypertile VAE'] = shared.opts.hypertile_vae_tile
return split_attention(vae, tile_size=shared.opts.hypertile_vae_tile, min_tile_size=128, swap_size=1)
@@ -161,6 +162,7 @@ def context_hypertile_unet(p):
return nullcontext()
else:
shared.log.info(f'Applying hypertile: unet={shared.opts.hypertile_unet_tile}')
p.extra_generation_params['Hypertile UNet'] = shared.opts.hypertile_unet_tile
return split_attention(unet, tile_size=shared.opts.hypertile_unet_tile, min_tile_size=128, swap_size=1)