Coverage for tests/diffusers_mock.py: 100%
50 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1from unittest.mock import MagicMock
3import importlib.util
5from typing import Dict
6from typing import Any
7from typing import Tuple
8from typing import Optional
10from PIL import Image
13def _pipeline_img_call(*args: Any, **kwargs: Any) -> MagicMock:
14 """Mock pipeline __call__ that returns an image matching the requested width/height."""
15 width = kwargs.get("width", 64)
16 height = kwargs.get("height", 64)
17 output = MagicMock()
18 output.images = [Image.new("RGB", (width, height))]
19 return output
22def _make_pipeline_class(
23 name: str,
24 ret_type: Optional[str] = None
25) -> type:
26 """Create a real class that inherits from MagicMock so it can be used with isinstance()."""
27 cls = type(name, (MagicMock,), {})
28 instance = cls()
29 cls.__bool__ = lambda self: True # type: ignore[attr-defined]
30 type(instance).__bool__ = lambda self: True
31 instance.to = MagicMock(return_value=instance)
32 instance.vae_scale_factor = 8
33 if ret_type == "pil":
34 instance.side_effect = _pipeline_img_call
35 cls.from_pretrained = MagicMock(return_value=instance) # type: ignore[attr-defined]
36 return cls
39class DiffusersMock(MagicMock):
40 def __init__(
41 self,
42 *args: Tuple,
43 **kwargs: Dict
44 ) -> None:
45 super().__init__(*args, **kwargs)
47 self.__spec__ = importlib.util.spec_from_loader("diffusers", loader=None)
49 # Pipeline / model classes as real types (usable with isinstance).
50 self.FluxPipeline = _make_pipeline_class("FluxPipeline", ret_type="pil")
51 self.Flux2Pipeline = _make_pipeline_class("Flux2Pipeline", ret_type="pil")
52 self.Flux2KleinPipeline = _make_pipeline_class("Flux2KleinPipeline", ret_type="pil")
53 self.FluxKontextPipeline = _make_pipeline_class("FluxKontextPipeline", ret_type="pil")
54 self.FluxControlNetModel = _make_pipeline_class("FluxControlNetModel", ret_type="pil")
55 self.FluxControlNetPipeline = _make_pipeline_class("FluxControlNetPipeline", ret_type="pil")
57 self.DiffusionPipeline = _make_pipeline_class("DiffusionPipeline")
58 self.HunyuanVideoFramepackPipeline = _make_pipeline_class("HunyuanVideoFramepackPipeline")
59 self.AutoencoderKLHunyuanVideo = _make_pipeline_class("AutoencoderKLHunyuanVideo")
60 self.FlowMatchEulerDiscreteScheduler = _make_pipeline_class("FlowMatchEulerDiscreteScheduler")
61 self.LTXConditionPipeline = _make_pipeline_class("LTXConditionPipeline")
62 self.LTXLatentUpsamplePipeline = _make_pipeline_class("LTXLatentUpsamplePipeline")
63 self.QwenImagePipeline = _make_pipeline_class("QwenImagePipeline", ret_type="pil")
64 self.QwenImageEditPipeline = _make_pipeline_class("QwenImageEditPipeline", ret_type="pil")
65 self.CogView4Pipeline = _make_pipeline_class("CogView4Pipeline", ret_type="pil")
66 self.HiDreamImagePipeline = _make_pipeline_class("HiDreamImagePipeline", ret_type="pil")
68 def get_sub_modules(self) -> Dict[str, Any]:
69 pipelines_mock = MagicMock()
70 # Expose pipeline classes that are imported from diffusers.pipelines
71 pipelines_mock.FluxControlNetPipeline = self.FluxControlNetPipeline
72 pipelines_mock.pipeline_utils = MagicMock()
73 pipelines_mock.pipeline_utils.DiffusionPipeline = self.DiffusionPipeline
75 return {
76 "diffusers": self,
77 "diffusers.models": MagicMock(),
78 "diffusers.models.autoencoders": MagicMock(),
79 "diffusers.models.autoencoders.autoencoder_kl": MagicMock(),
80 "diffusers.models.activations": MagicMock(),
81 "diffusers.models.attention": MagicMock(),
82 "diffusers.models.attention_processor": MagicMock(),
83 "diffusers.models.embeddings": MagicMock(),
84 "diffusers.models.unets": MagicMock(),
85 "diffusers.models.unets.unet_2d_blocks": MagicMock(),
86 "diffusers.models.normalization": MagicMock(),
87 "diffusers.models.transformers": MagicMock(),
88 "diffusers.models.transformers.dual_transformer_2d": MagicMock(),
89 'diffusers.models.transformers.transformer_2d': MagicMock(),
90 "diffusers.models.transformers.transformer_flux2": MagicMock(),
91 "diffusers.models.transformers.transformer_hunyuan_video": MagicMock(),
92 "diffusers.models.downsampling": MagicMock(),
93 "diffusers.models.upsampling": MagicMock(),
94 "diffusers.models.resnet": MagicMock(),
95 "diffusers.pipelines": pipelines_mock,
96 "diffusers.pipelines.pipeline_utils": pipelines_mock.pipeline_utils,
97 "diffusers.pipelines.ltx": MagicMock(),
98 "diffusers.pipelines.ltx.pipeline_ltx_condition": MagicMock(),
99 "diffusers.configuration_utils": MagicMock(),
100 "diffusers.schedulers": MagicMock(),
101 "diffusers.schedulers.scheduling_utils": MagicMock(),
102 "diffusers.utils": MagicMock(),
103 "diffusers.utils.torch_utils": MagicMock(),
104 }