MCPcopy Create free account
hub / github.com/TencentARC/AnimeGamer / initialize_distributed

Function initialize_distributed

VDM_Decoder/arguments.py:185–259  ·  view source on GitHub ↗

Initialize torch.distributed.

(args)

Source from the content-addressed store, hash-verified

183
184
185def initialize_distributed(args):
186 """Initialize torch.distributed."""
187 if torch.distributed.is_initialized():
188 if mpu.model_parallel_is_initialized():
189 if args.model_parallel_size != mpu.get_model_parallel_world_size():
190 raise ValueError(
191 "model_parallel_size is inconsistent with prior configuration."
192 "We currently do not support changing model_parallel_size."
193 )
194 return False
195 else:
196 if args.model_parallel_size > 1:
197 warnings.warn(
198 "model_parallel_size > 1 but torch.distributed is not initialized via SAT."
199 "Please carefully make sure the correctness on your own."
200 )
201 mpu.initialize_model_parallel(args.model_parallel_size)
202 return True
203 # the automatic assignment of devices has been moved to arguments.py
204 if args.device == "cpu":
205 pass
206 else:
207 torch.cuda.set_device(args.device)
208 # Call the init process
209 init_method = "tcp://"
210 args.master_ip = os.getenv("MASTER_ADDR", "localhost")
211
212 if args.world_size == 1:
213 from sat.helpers import get_free_port
214
215 default_master_port = str(get_free_port())
216 else:
217 default_master_port = "6000"
218 args.master_port = os.getenv("MASTER_PORT", default_master_port)
219 init_method += args.master_ip + ":" + args.master_port
220 torch.distributed.init_process_group(
221 backend=args.distributed_backend, world_size=args.world_size, rank=args.rank, init_method=init_method
222 )
223
224 # Set the model-parallel / data-parallel communicators.
225 mpu.initialize_model_parallel(args.model_parallel_size)
226
227 # Set vae context parallel group equal to model parallel group
228 from .sgm.util import set_context_parallel_group, initialize_context_parallel
229
230 if args.model_parallel_size <= 2:
231 set_context_parallel_group(args.model_parallel_size, mpu.get_model_parallel_group())
232 else:
233 initialize_context_parallel(2)
234 # mpu.initialize_model_parallel(1)
235 # Optional DeepSpeed Activation Checkpointing Features
236 if args.deepspeed:
237 import deepspeed
238
239 deepspeed.init_distributed(
240 dist_backend=args.distributed_backend, world_size=args.world_size, rank=args.rank, init_method=init_method
241 )
242 # # It seems that it has no negative influence to configure it even without using checkpointing.

Callers 1

get_argsFunction · 0.85

Calls 2

Tested by

no test coverage detected