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
« 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
7import torch
8import torch.distributed as dist
10from typing import Dict
11from typing import Any
12from typing import Optional
13from typing import Tuple
15from wrapper_model import ModelGeneration
17if not torch.cuda.is_available():
18 raise RuntimeError("This module requires CUDA support.")
20from xfuser.config import EngineConfig
21from xfuser.core.distributed import initialize_model_parallel
22from xfuser.core.distributed import init_distributed_environment
25class USPGeneration(ModelGeneration):
26 """
27 Base class for generation using USP.
28 This models support Unified Sequence Parallelism (USP).
29 """
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)
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
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
50 self.param_dtype = param_dtype
52 # Parallelism
53 self.gpu: Optional[str] = None
54 if torch.cuda.is_available():
55 self.gpu = torch.cuda.get_device_name(0)
57 self.base_seed = random.randint(0, sys.maxsize)
59 # Model features
60 self.num_heads = -1
61 self.vae_stride: Optional[Tuple[int, int, int]] = None # time, height, width
63 def __del__(self) -> None:
64 if dist.is_initialized():
65 dist.destroy_process_group()
66 super().__del__()
68 def set_seed(
69 self,
70 seed: int
71 ) -> None:
72 """Set the seed for random number generation."""
73 self.base_seed = seed
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}")
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)
88 def init_parallelism(self) -> None:
89 self.load_timer.start("torch_dist")
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))
95 self.device_id = self.local_rank
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
104 self.device = torch.device(f"cuda:{self.device_id}")
106 torch.cuda.set_device(self.local_rank)
108 if self.world_size <= 1:
109 self.load_timer.end("torch_dist")
110 return # Single GPU mode, no parallelism needed
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 )
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}.")
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}`.")
130 if dist.is_initialized():
131 self.reset_seed()
133 self.load_timer.end("torch_dist")
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")
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")
150 if not dist.is_initialized():
151 raise RuntimeError("Distributed process group not initialized")
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