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

1from unittest.mock import MagicMock 

2 

3import importlib.util 

4 

5from typing import Dict 

6from typing import Any 

7from typing import Tuple 

8from typing import Optional 

9 

10from PIL import Image 

11 

12 

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 

20 

21 

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 

37 

38 

39class DiffusersMock(MagicMock): 

40 def __init__( 

41 self, 

42 *args: Tuple, 

43 **kwargs: Dict 

44 ) -> None: 

45 super().__init__(*args, **kwargs) 

46 

47 self.__spec__ = importlib.util.spec_from_loader("diffusers", loader=None) 

48 

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") 

56 

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") 

67 

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 

74 

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 }