(self, device, dtype, **kwargs)
| 25 | |
| 26 | class Model: |
| 27 | def __init__(self, device, dtype, **kwargs): |
| 28 | self.device = device |
| 29 | self.dtype = dtype |
| 30 | self.generator = torch.Generator(device=device) |
| 31 | self.pipe_dict = { |
| 32 | ModelType.Pix2Pix_Video: StableDiffusionInstructPix2PixPipeline, |
| 33 | ModelType.Text2Video: TextToVideoPipeline, |
| 34 | ModelType.ControlNetCanny: StableDiffusionControlNetPipeline, |
| 35 | ModelType.ControlNetCannyDB: StableDiffusionControlNetPipeline, |
| 36 | ModelType.ControlNetPose: StableDiffusionControlNetPipeline, |
| 37 | ModelType.ControlNetDepth: StableDiffusionControlNetPipeline, |
| 38 | } |
| 39 | self.controlnet_attn_proc = utils.CrossFrameAttnProcessor( |
| 40 | unet_chunk_size=2) |
| 41 | self.pix2pix_attn_proc = utils.CrossFrameAttnProcessor( |
| 42 | unet_chunk_size=3) |
| 43 | self.text2video_attn_proc = utils.CrossFrameAttnProcessor( |
| 44 | unet_chunk_size=2) |
| 45 | |
| 46 | self.pipe = None |
| 47 | self.model_type = None |
| 48 | |
| 49 | self.states = {} |
| 50 | self.model_name = "" |
| 51 | |
| 52 | def set_model(self, model_type: ModelType, model_id: str, **kwargs): |
| 53 | if hasattr(self, "pipe") and self.pipe is not None: |
nothing calls this directly
no outgoing calls
no test coverage detected