@@ -92,21 +92,18 @@ void GlobalEnv::Init(int nthread_per_process, int tensor_parallel_size, bool seq
9292
9393 CHECK (!initialized_) << " Repeated initialization of GlobalEnv!" ;
9494
95- const int proc_world_size = GetEnvAsInt (" WORLD_SIZE" , GetEnvAsInt (" PROC_WORLD_SIZE" , 1 ));
96- nproc_per_node_ = GetEnvAsInt (" LOCAL_WORLD_SIZE" , GetEnvAsInt (" NPROC_PER_NODE" , 1 ));
97- CHECK_GT (nproc_per_node_, 0 ) << " NPROC_PER_NODE/LOCAL_WORLD_SIZE must be positive" ;
98- CHECK_GT (proc_world_size, 0 ) << " PROC_WORLD_SIZE/WORLD_SIZE must be positive" ;
99- CHECK_EQ (proc_world_size % nproc_per_node_, 0 )
100- << " PROC_WORLD_SIZE/WORLD_SIZE must be divisible by NPROC_PER_NODE/LOCAL_WORLD_SIZE" ;
95+ const int proc_world_size = GetEnvAsInt (" WORLD_SIZE" , 1 );
96+ nproc_per_node_ = GetEnvAsInt (" LOCAL_WORLD_SIZE" , 1 );
97+ CHECK_GT (nproc_per_node_, 0 ) << " LOCAL_WORLD_SIZE must be positive" ;
98+ CHECK_GT (proc_world_size, 0 ) << " WORLD_SIZE must be positive" ;
99+ CHECK_EQ (proc_world_size % nproc_per_node_, 0 ) << " WORLD_SIZE must be divisible by LOCAL_WORLD_SIZE" ;
101100 nnodes_ = proc_world_size / nproc_per_node_;
102- global_proc_rank_ = GetEnvAsInt (" RANK" , GetEnvAsInt (" GLOBAL_PROC_RANK" , 0 ));
103- local_proc_rank_ = GetEnvAsInt (" LOCAL_RANK" , GetEnvAsInt (" LOCAL_PROC_RANK" , 0 ));
104- CHECK_GE (global_proc_rank_, 0 ) << " GLOBAL_PROC_RANK/RANK must be non-negative" ;
105- CHECK_LT (global_proc_rank_, proc_world_size)
106- << " GLOBAL_PROC_RANK/RANK must be less than PROC_WORLD_SIZE/WORLD_SIZE" ;
107- CHECK_GE (local_proc_rank_, 0 ) << " LOCAL_PROC_RANK/LOCAL_RANK must be non-negative" ;
108- CHECK_LT (local_proc_rank_, nproc_per_node_)
109- << " LOCAL_PROC_RANK/LOCAL_RANK must be less than NPROC_PER_NODE/LOCAL_WORLD_SIZE" ;
101+ global_proc_rank_ = GetEnvAsInt (" RANK" , 0 );
102+ local_proc_rank_ = GetEnvAsInt (" LOCAL_RANK" , 0 );
103+ CHECK_GE (global_proc_rank_, 0 ) << " RANK must be non-negative" ;
104+ CHECK_LT (global_proc_rank_, proc_world_size) << " RANK must be less than WORLD_SIZE" ;
105+ CHECK_GE (local_proc_rank_, 0 ) << " LOCAL_RANK must be non-negative" ;
106+ CHECK_LT (local_proc_rank_, nproc_per_node_) << " LOCAL_RANK must be less than LOCAL_WORLD_SIZE" ;
110107
111108 nthread_per_process_ = nthread_per_process;
112109 world_size_ = proc_world_size * nthread_per_process;
0 commit comments