diff --git a/modules/sd_hijack.py b/modules/sd_hijack.py index 4af457dfd..f9652d215 100644 --- a/modules/sd_hijack.py +++ b/modules/sd_hijack.py @@ -104,20 +104,6 @@ class StableDiffusionModelHijack: m.cond_stage_model.model.token_embedding = EmbeddingsWithFixes(m.cond_stage_model.model.token_embedding, self) m.cond_stage_model = sd_hijack_open_clip.FrozenOpenCLIPEmbedderWithCustomWords(m.cond_stage_model, self) - """ - try: - import time - t0 = time.time() - torch._dynamo.config.verbose = True - torch.backends.cudnn.benchmark = True - m.model = torch.compile(m.model, mode="max-autotune", fullgraph=False) - m = torch.compile(m, mode="max-autotune", fullgraph=False) - t1 = time.time() - print(f"Model compiled in {round(t1 - t0, 2)} sec") - except Exception as err: - print(f"Model compile not supported: {err}") - """ - self.optimization_method = apply_optimizations() self.clip = m.cond_stage_model diff --git a/modules/sd_models.py b/modules/sd_models.py index 3ea1dea7f..099de7017 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -400,17 +400,23 @@ def load_model(checkpoint_info=None): sd_hijack.model_hijack.hijack(sd_model) - sd_model.eval() + if shared.cmd_opts.compile is not None: + try: + import time + import torch._dynamo as dynamo # must be imported explicitly or namespace is not found + torch._dynamo.config.verbose=True + torch.backends.cudnn.benchmark = True + t0 = time.time() + # script = sd_model.model.to_torchscript(method="trace") + # script = torch.jit.script(sd_model.model.eval()) + # sd_model.model = torch.compile(script, backend=shared.cmd_opts.compile) + sd_model.model = torch.compile(sd_model, backend=shared.cmd_opts.compile) + t1 = time.time() + print(f"Model compiled using backend {shared.cmd_opts.compile} in {round(t1 - t0, 2)} sec") + except Exception as err: + print(f"Model compile not supported: {err}") - """ - try: - t0 = time.time() - sd_model = torch.compile(sd_model, mode="max-autotune", fullgraph=True) - t1 = time.time() - print(f"Model compiled in {round(t1 - t0, 2)} sec") - except Exception as err: - print(f"Model compile not supported: {err}") - """ + sd_model.eval() shared.sd_model = sd_model diff --git a/modules/shared.py b/modules/shared.py index b32540b13..5f6bd7021 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -101,6 +101,7 @@ parser.add_argument("--cors-allow-origins-regex", type=str, help="Allowed CORS o parser.add_argument("--tls-keyfile", type=str, help="Partially enables TLS, requires --tls-certfile to fully function", default=None) parser.add_argument("--tls-certfile", type=str, help="Partially enables TLS, requires --tls-keyfile to fully function", default=None) parser.add_argument("--server-name", type=str, help="Sets hostname of server", default=None) +parser.add_argument("--compile", type=str, help="Use Torch Dynamo compile with specified backend", default=None) script_loading.preload_extensions(extensions.extensions_dir, parser) script_loading.preload_extensions(extensions.extensions_builtin_dir, parser)