Skip to content

Commit 0575ab6

Browse files
Remove duplicated groups for EP and DP_EP
This addresses the problem described in #222
1 parent cee5b9f commit 0575ab6

1 file changed

Lines changed: 33 additions & 37 deletions

File tree

‎astra-sim-alibabacloud/astra-sim/system/MockNcclGroup.cc‎

Lines changed: 33 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)