Coverage for tests/test_wrapper_hidream.py: 100%

44 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-09 04:47 +0000

1#!/usr/bin/env python3 

2 

3import sys 

4import pytest 

5 

6from unittest.mock import patch 

7from unittest.mock import MagicMock 

8from tests.torch_mock import TorchMock 

9from tests.diffusers_mock import DiffusersMock 

10 

11from PIL import Image 

12 

13mock_torch = TorchMock() 

14mock_diffusers = DiffusersMock() 

15 

16sys.path.append("wrapper") 

17sys.path.append("wrapper/hidream") 

18 

19mock_modules = { 

20 'nvidia_smi': MagicMock(), 

21 'imageio': MagicMock(), 

22 'cv2': MagicMock(), 

23 'torch': mock_torch, 

24 'xfuser': MagicMock(), 

25 'xfuser.config': MagicMock(), 

26 'xfuser.core': MagicMock(), 

27 'xfuser.core.distributed': MagicMock(), 

28 'xfuser.model_executor': MagicMock(), 

29 'xfuser.model_executor.layers': MagicMock(), 

30 'xfuser.model_executor.layers.attention_processor': MagicMock(), 

31 'transformers': MagicMock(), 

32} 

33mock_modules.update(mock_torch.get_sub_modules()) 

34mock_modules.update(mock_diffusers.get_sub_modules()) 

35 

36with patch.dict(sys.modules, mock_modules): 

37 from hidream.wrapper_hidream import HiDreamGeneration 

38 

39 

40@pytest.mark.asyncio 

41async def test_basic() -> None: 

42 model = HiDreamGeneration() 

43 assert model is not None 

44 assert model.model_name == "hidream" 

45 assert model.status == "initializing" 

46 

47 with pytest.raises(ValueError): 

48 await model.generate(64, 48, "test prompt") 

49 

50 model.init() 

51 assert model.status == "ok" 

52 

53 health = model.get_health() 

54 assert health is not None 

55 timestamps = model.get_timestamps() 

56 assert timestamps is not None 

57 

58 with pytest.raises(ValueError): 

59 await model.get_rest_args(None) 

60 with pytest.raises(ValueError): 

61 await model.get_rest_args({}) 

62 await model.get_rest_args({ 

63 "job_id": "unittest", 

64 "prompt": "Test prompt", 

65 "width": 80, 

66 "height": 60, 

67 "seed": 7, 

68 }) 

69 

70 await model.warmup() 

71 

72 image = await model.generate( 

73 width=1280, 

74 height=800, 

75 prompt="Test prompt") 

76 assert image is not None 

77 assert isinstance(image, Image.Image) 

78 assert image.size == (1280, 800) 

79 

80 # 48x48 not supported for 2 GPUs (latent shape 9, odd). 

81 model.world_size = 2 

82 with pytest.raises(ValueError, match="48x48 not supported for 2 GPUs"): 

83 await model.generate( 

84 width=48, 

85 height=48, 

86 prompt="Test prompt") 

87 

88 del model