Coverage for wrapper/wrapper_usp.py: 71%

90 statements  

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

1import os 

2import sys 

3import datetime 

4import random 

5import logging 

6 

7import torch 

8import torch.distributed as dist 

9 

10from typing import Dict 

11from typing import Any 

12from typing import Optional 

13from typing import Tuple 

14 

15from wrapper_model import ModelGeneration 

16 

17if not torch.cuda.is_available(): 

18 raise RuntimeError("This module requires CUDA support.") 

19 

20from xfuser.config import EngineConfig 

21from xfuser.core.distributed import initialize_model_parallel 

22from xfuser.core.distributed import init_distributed_environment 

23 

24 

25class USPGeneration(ModelGeneration): 

26 """ 

27 Base class for generation using USP. 

28 This models support Unified Sequence Parallelism (USP). 

29 """ 

30 

31 def __init__( 

32 self, 

33 model_name: str, 

34 engine_config: Optional[EngineConfig] = None, 

35 param_dtype: torch.dtype = torch.bfloat16, 

36 ) -> None: 

37 super().__init__(model_name) 

38 

39 # Model components 

40 self.engine_config = engine_config 

41 self.ulysses_size = -1 

42 self.ring_size = -1 

43 if self.engine_config and getattr(self.engine_config, "parallel_config", None): 

44 self.ulysses_size = self.engine_config.parallel_config.sp_config.ulysses_degree 

45 self.ring_size = self.engine_config.parallel_config.sp_config.ring_degree 

46 

47 if self.engine_config and getattr(self.engine_config, "runtime_config", None): 

48 self.torch_compile = self.engine_config.runtime_config.use_torch_compile 

49 

50 self.param_dtype = param_dtype 

51 

52 # Parallelism 

53 self.gpu: Optional[str] = None 

54 if torch.cuda.is_available(): 

55 self.gpu = torch.cuda.get_device_name(0) 

56 

57 self.base_seed = random.randint(0, sys.maxsize) 

58 

59 # Model features 

60 self.num_heads = -1 

61 self.vae_stride: Optional[Tuple[int, int, int]] = None # time, height, width 

62 

63 def __del__(self) -> None: 

64 if dist.is_initialized(): 

65 dist.destroy_process_group() 

66 super().__del__() 

67 

68 def set_seed( 

69 self, 

70 seed: int 

71 ) -> None: 

72 """Set the seed for random number generation.""" 

73 self.base_seed = seed 

74 

75 if dist.is_initialized(): 

76 global_base_seed: list[int | None] = [self.base_seed] if self.rank == 0 else [None] 

77 dist.broadcast_object_list(global_base_seed, src=0) 

78 seed_value = global_base_seed[0] 

79 assert seed_value is not None 

80 self.base_seed = seed_value 

81 logging.info(f"[{self.rank}] Using base seed: {self.base_seed}") 

82 

83 def reset_seed(self) -> None: 

84 """Reset the seed for random number generation.""" 

85 rnd_seed = random.randint(0, sys.maxsize) 

86 self.set_seed(rnd_seed) 

87 

88 def init_parallelism(self) -> None: 

89 self.load_timer.start("torch_dist") 

90 

91 self.rank = int(os.getenv("RANK", 0)) 

92 self.local_rank = int(os.getenv("LOCAL_RANK", 0)) 

93 self.world_size = int(os.getenv("WORLD_SIZE", 1)) 

94 

95 self.device_id = self.local_rank 

96 

97 if not torch.cuda.is_available(): 

98 self.device_id = 0 

99 self.device = torch.device("cpu") 

100 logging.warning("CUDA is not available. Running on CPU.") 

101 self.load_timer.end("torch_dist") 

102 return # Single GPU mode, no parallelism needed 

103 

104 self.device = torch.device(f"cuda:{self.device_id}") 

105 

106 torch.cuda.set_device(self.local_rank) 

107 

108 if self.world_size <= 1: 

109 self.load_timer.end("torch_dist") 

110 return # Single GPU mode, no parallelism needed 

111 

112 if not dist.is_initialized(): 

113 dist.init_process_group( 

114 backend="nccl", 

115 init_method="env://", 

116 rank=self.rank, 

117 world_size=self.world_size, 

118 timeout=datetime.timedelta(hours=24), # Prevent NCCL timeout 

119 ) 

120 

121 # Unified Sequence Parallelism (USP) 

122 if self.ulysses_size * self.ring_size != self.world_size: 

123 raise ValueError( 

124 f"ulysses_size {self.ulysses_size} x ring_size {self.ring_size} != world size {self.world_size}.") 

125 

126 if self.ulysses_size > 1 and self.num_heads > 0: 

127 if self.num_heads % self.ulysses_size != 0: 

128 raise ValueError(f"`{self.num_heads}` cannot be divided evenly by `{self.ulysses_size}`.") 

129 

130 if dist.is_initialized(): 

131 self.reset_seed() 

132 

133 self.load_timer.end("torch_dist") 

134 

135 self.load_timer.start("init_distributed") 

136 init_distributed_environment( 

137 rank=dist.get_rank(), 

138 world_size=dist.get_world_size() 

139 ) 

140 self.load_timer.end("init_distributed") 

141 

142 self.load_timer.start("model_parallel") 

143 initialize_model_parallel( 

144 sequence_parallel_degree=dist.get_world_size(), 

145 ulysses_degree=self.ulysses_size, 

146 ring_degree=self.ring_size, 

147 ) 

148 self.load_timer.end("model_parallel") 

149 

150 if not dist.is_initialized(): 

151 raise RuntimeError("Distributed process group not initialized") 

152 

153 def get_health(self) -> Dict[str, Any]: 

154 ret = super().get_health() 

155 ret.update({ 

156 "gpu": self.gpu, 

157 "rank": self.rank, 

158 "world_size": self.world_size, 

159 "ulysses_size": self.ulysses_size, 

160 "ring_size": self.ring_size, 

161 "torch_compile": self.torch_compile, 

162 "dtype": str(self.param_dtype), 

163 "vae_stride": self.vae_stride, 

164 "num_heads": self.num_heads, 

165 }) 

166 return ret