| 238 | |
| 239 | # Data partition according to the rank |
| 240 | def partition(global_rank, world_size, train_x, train_y, val_x, val_y): |
| 241 | # Partition training data |
| 242 | data_per_rank = train_x.shape[0] // world_size |
| 243 | idx_start = global_rank * data_per_rank |
| 244 | idx_end = (global_rank + 1) * data_per_rank |
| 245 | train_x = train_x[idx_start:idx_end] |
| 246 | train_y = train_y[idx_start:idx_end] |
| 247 | |
| 248 | # Partition evaluation data |
| 249 | data_per_rank = val_x.shape[0] // world_size |
| 250 | idx_start = global_rank * data_per_rank |
| 251 | idx_end = (global_rank + 1) * data_per_rank |
| 252 | val_x = val_x[idx_start:idx_end] |
| 253 | val_y = val_y[idx_start:idx_end] |
| 254 | return train_x, train_y, val_x, val_y |
| 255 | |
| 256 | |
| 257 | # Function to all reduce NUMPY accuracy and loss from multiple devices |