mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
add checks
This commit is contained in:
@@ -1,4 +1,5 @@
|
||||
import os
|
||||
import sys
|
||||
import json
|
||||
import shutil
|
||||
import tempfile
|
||||
@@ -45,11 +46,13 @@ class PipelineBase(TorchCompatibleModule, diffusers.DiffusionPipeline, metaclass
|
||||
|
||||
module = getattr(self, name)
|
||||
|
||||
if isinstance(module, optimum.onnxruntime.modeling_diffusion._ORTDiffusionModelPart): # pylint: disable=protected-access
|
||||
device = extract_device(args, kwargs)
|
||||
if device is None:
|
||||
return self
|
||||
module.session = move_inference_session(module.session, device)
|
||||
if "optimum.onnxruntime" in sys.modules:
|
||||
import optimum.onnxruntime
|
||||
if isinstance(module, optimum.onnxruntime.modeling_diffusion._ORTDiffusionModelPart): # pylint: disable=protected-access
|
||||
device = extract_device(args, kwargs)
|
||||
if device is None:
|
||||
return self
|
||||
module.session = move_inference_session(module.session, device)
|
||||
|
||||
if not isinstance(module, diffusers.OnnxRuntimeModel):
|
||||
continue
|
||||
|
||||
@@ -74,12 +74,16 @@ def get_pipelines():
|
||||
'InstaFlow': getattr(diffusers, 'StableDiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser
|
||||
'SegMoE': getattr(diffusers, 'StableDiffusionPipeline', None), # dynamically redefined and loaded in sd_models.load_diffuser
|
||||
}
|
||||
if hasattr(diffusers, 'OnnxStableDiffusionXLPipeline'):
|
||||
if hasattr(diffusers, 'OnnxStableDiffusionPipeline'):
|
||||
onnx_pipelines = {
|
||||
'ONNX Stable Diffusion': getattr(diffusers, 'OnnxStableDiffusionPipeline', None),
|
||||
'ONNX Stable Diffusion Img2Img': getattr(diffusers, 'OnnxStableDiffusionImg2ImgPipeline', None),
|
||||
'ONNX Stable Diffusion Inpaint': getattr(diffusers, 'OnnxStableDiffusionInpaintPipeline', None),
|
||||
'ONNX Stable Diffusion Upscale': getattr(diffusers, 'OnnxStableDiffusionUpscalePipeline', None),
|
||||
}
|
||||
pipelines.update(onnx_pipelines)
|
||||
if hasattr(diffusers, 'OnnxStableDiffusionXLPipeline'):
|
||||
onnx_pipelines = {
|
||||
'ONNX Stable Diffusion XL': getattr(diffusers, 'OnnxStableDiffusionXLPipeline', None),
|
||||
'ONNX Stable Diffusion XL Img2Img': getattr(diffusers, 'OnnxStableDiffusionXLImg2ImgPipeline', None),
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user