@@ -106,25 +106,23 @@ namespace MockNccl {
106106 std::set<int > EPnodes;
107107 for (int i = 0 ; i < TP_nums / _EP_size; i++){
108108 TP_idx = i*_EP_size;
109- for (int j =0 ;j<_EP_size;j++){
110- for (int k = 0 ;k<AllTPGroups[TP_idx].Ranks .size ();k++){
111- ranks.clear ();
112- EPnodes.clear ();
113- for (int l = TP_idx;l<TP_idx+_EP_size;l++){
114- int tmp_rank = AllTPGroups[l].Ranks [k];
115- int node_idx = tmp_rank/_gpus_per_nodes;
116- ranks.push_back (tmp_rank);
117- GroupIndex[std::make_pair (tmp_rank, EP )] = all_group_idx;
118- EPnodes.insert (node_idx);
119- }
120- NVSwitchs.clear ();
121- for (int idx:EPnodes){
122- NVSwitchs.push_back (_NVSwitch[idx]);
123- GroupIndex[std::make_pair (_NVSwitch[idx],EP )] = all_group_idx;
124- }
125- AllGroups[all_group_idx] = GroupInfo (all_group_idx,EP ,EPnodes.size (),_EP_size,ranks,NVSwitchs);
126- all_group_idx++;
109+ for (int k = 0 ;k<AllTPGroups[TP_idx].Ranks .size ();k++){
110+ ranks.clear ();
111+ EPnodes.clear ();
112+ for (int l = TP_idx;l<TP_idx+_EP_size;l++){
113+ int tmp_rank = AllTPGroups[l].Ranks [k];
114+ int node_idx = tmp_rank/_gpus_per_nodes;
115+ ranks.push_back (tmp_rank);
116+ GroupIndex[std::make_pair (tmp_rank, EP )] = all_group_idx;
117+ EPnodes.insert (node_idx);
118+ }
119+ NVSwitchs.clear ();
120+ for (int idx:EPnodes){
121+ NVSwitchs.push_back (_NVSwitch[idx]);
122+ GroupIndex[std::make_pair (_NVSwitch[idx],EP )] = all_group_idx;
127123 }
124+ AllGroups[all_group_idx] = GroupInfo (all_group_idx,EP ,EPnodes.size (),_EP_size,ranks,NVSwitchs);
125+ all_group_idx++;
128126 }
129127 }
130128 }
@@ -134,25 +132,23 @@ namespace MockNccl {
134132 std::set<int > DP_EP_nodes;
135133 for (int i = 0 ; i < TP_nums / _DP_EP_size; i++){
136134 TP_idx = i;
137- for (int j = 0 ; j < _DP_EP_size; j++){
138- for (int k = 0 ; k < AllTPGroups[TP_idx].Ranks .size (); k++){
139- ranks.clear ();
140- DP_EP_nodes.clear ();
141- for (int l = TP_idx; l < TP_idx + _DP_EP_size * _EP_size; l += _EP_size){
142- int tmp_rank = AllTPGroups[l].Ranks [k];
143- int node_idx = tmp_rank / _gpus_per_nodes;
144- ranks.push_back (tmp_rank);
145- GroupIndex[std::make_pair (tmp_rank, DP_EP )] = all_group_idx;
146- DP_EP_nodes.insert (node_idx);
147- }
148- NVSwitchs.clear ();
149- for (int idx : DP_EP_nodes){
150- NVSwitchs.push_back (_NVSwitch[idx]);
151- GroupIndex[std::make_pair (_NVSwitch[idx], DP_EP )] = all_group_idx;
152- }
153- AllGroups[all_group_idx] = GroupInfo (all_group_idx, DP_EP , DP_EP_nodes.size (), _DP_EP_size, ranks, NVSwitchs);
154- all_group_idx++;
135+ for (int k = 0 ; k < AllTPGroups[TP_idx].Ranks .size (); k++){
136+ ranks.clear ();
137+ DP_EP_nodes.clear ();
138+ for (int l = TP_idx; l < TP_idx + _DP_EP_size * _EP_size; l += _EP_size){
139+ int tmp_rank = AllTPGroups[l].Ranks [k];
140+ int node_idx = tmp_rank / _gpus_per_nodes;
141+ ranks.push_back (tmp_rank);
142+ GroupIndex[std::make_pair (tmp_rank, DP_EP )] = all_group_idx;
143+ DP_EP_nodes.insert (node_idx);
144+ }
145+ NVSwitchs.clear ();
146+ for (int idx : DP_EP_nodes){
147+ NVSwitchs.push_back (_NVSwitch[idx]);
148+ GroupIndex[std::make_pair (_NVSwitch[idx], DP_EP )] = all_group_idx;
155149 }
150+ AllGroups[all_group_idx] = GroupInfo (all_group_idx, DP_EP , DP_EP_nodes.size (), _DP_EP_size, ranks, NVSwitchs);
151+ all_group_idx++;
156152 }
157153 }
158154 }
@@ -2100,4 +2096,4 @@ namespace MockNccl {
21002096 return info;
21012097 }
21022098 }
2103- }
2099+ }
0 commit comments