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
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1#!/usr/bin/env python3
3import sys
4import pytest
6from unittest.mock import patch
7from unittest.mock import MagicMock
8from tests.torch_mock import TorchMock
9from tests.diffusers_mock import DiffusersMock
11from PIL import Image
13mock_torch = TorchMock()
14mock_diffusers = DiffusersMock()
16sys.path.append("wrapper")
17sys.path.append("wrapper/hidream")
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())
36with patch.dict(sys.modules, mock_modules):
37 from hidream.wrapper_hidream import HiDreamGeneration
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"
47 with pytest.raises(ValueError):
48 await model.generate(64, 48, "test prompt")
50 model.init()
51 assert model.status == "ok"
53 health = model.get_health()
54 assert health is not None
55 timestamps = model.get_timestamps()
56 assert timestamps is not None
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 })
70 await model.warmup()
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)
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")
88 del model