MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / SDImagePipeline

Class SDImagePipeline

diffsynth/pipelines/sd_image.py:14–191  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

12
13
14class SDImagePipeline(BasePipeline):
15
16 def __init__(self, device="cuda", torch_dtype=torch.float16):
17 super().__init__(device=device, torch_dtype=torch_dtype)
18 self.scheduler = EnhancedDDIMScheduler()
19 self.prompter = SDPrompter()
20 # models
21 self.text_encoder: SDTextEncoder = None
22 self.unet: SDUNet = None
23 self.vae_decoder: SDVAEDecoder = None
24 self.vae_encoder: SDVAEEncoder = None
25 self.controlnet: MultiControlNetManager = None
26 self.ipadapter_image_encoder: IpAdapterCLIPImageEmbedder = None
27 self.ipadapter: SDIpAdapter = None
28 self.model_names = ['text_encoder', 'unet', 'vae_decoder', 'vae_encoder', 'controlnet', 'ipadapter_image_encoder', 'ipadapter']
29
30
31 def denoising_model(self):
32 return self.unet
33
34
35 def fetch_models(self, model_manager: ModelManager, controlnet_config_units: List[ControlNetConfigUnit]=[], prompt_refiner_classes=[]):
36 # Main models
37 self.text_encoder = model_manager.fetch_model("sd_text_encoder")
38 self.unet = model_manager.fetch_model("sd_unet")
39 self.vae_decoder = model_manager.fetch_model("sd_vae_decoder")
40 self.vae_encoder = model_manager.fetch_model("sd_vae_encoder")
41 self.prompter.fetch_models(self.text_encoder)
42 self.prompter.load_prompt_refiners(model_manager, prompt_refiner_classes)
43
44 # ControlNets
45 controlnet_units = []
46 for config in controlnet_config_units:
47 controlnet_unit = ControlNetUnit(
48 Annotator(config.processor_id, device=self.device),
49 model_manager.fetch_model("sd_controlnet", config.model_path),
50 config.scale
51 )
52 controlnet_units.append(controlnet_unit)
53 self.controlnet = MultiControlNetManager(controlnet_units)
54
55 # IP-Adapters
56 self.ipadapter = model_manager.fetch_model("sd_ipadapter")
57 self.ipadapter_image_encoder = model_manager.fetch_model("sd_ipadapter_clip_image_encoder")
58
59
60 @staticmethod
61 def from_model_manager(model_manager: ModelManager, controlnet_config_units: List[ControlNetConfigUnit]=[], prompt_refiner_classes=[], device=None):
62 pipe = SDImagePipeline(
63 device=model_manager.device if device is None else device,
64 torch_dtype=model_manager.torch_dtype,
65 )
66 pipe.fetch_models(model_manager, controlnet_config_units, prompt_refiner_classes=[])
67 return pipe
68
69
70 def encode_image(self, image, tiled=False, tile_size=64, tile_stride=32):
71 latents = self.vae_encoder(image, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)

Callers 1

from_model_managerMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected