r""" Default setting for the pipeline is: `(layer_id + offset) // (layers / pipeline_stage)`. Args: network(Cell) - Represents the transformer block layer_id(int) - Means the layer index for the current module, counts from zero. offset(int) - Mea
(
network, layer_id, offset, parallel_config, layers
)
| 259 | |
| 260 | |
| 261 | def set_parallel_configure_for_layer( |
| 262 | network, layer_id, offset, parallel_config, layers |
| 263 | ): |
| 264 | r""" |
| 265 | Default setting for the pipeline is: `(layer_id + offset) // (layers / pipeline_stage)`. |
| 266 | |
| 267 | |
| 268 | Args: |
| 269 | network(Cell) - Represents the transformer block |
| 270 | layer_id(int) - Means the layer index for the current module, counts from zero. |
| 271 | offset(int) - Means the layer_index needs a offset, if there are other modules in the net. |
| 272 | layers(int) - The total layers used for the model. |
| 273 | """ |
| 274 | # Used for the pipeline's stages setting |
| 275 | # As the final layer is not included here, so we need to manually add here. |
| 276 | # original: if set two stages, layers on two stages will be [15, 16+1] |
| 277 | # with 1 added, the layers on two stages will be [16, 15 +1] |
| 278 | pp_dis = max(int((layers + 1) / parallel_config.pipeline_stage), 1) |
| 279 | # the pipeline stage must be in [0, parallel_config.pipeline_stage - 1] |
| 280 | pp_id = min((layer_id + offset) // pp_dis, parallel_config.pipeline_stage - 1) |
| 281 | network.pipeline_stage = pp_id |
| 282 | print(f"pipeline stage id is {pp_id}", flush=True) |
| 283 | |
| 284 | # Used for optimizer's fusion tag |
| 285 | dis = max(int((layers + 1) / parallel_config.gradient_aggregation_group), 1) |
| 286 | if parallel_config.pipeline_stage > 1: |
| 287 | # we give the fusion in pipeline mode a fixed value, otherwise the performance may become worse. |
| 288 | network.set_comm_fusion(2) |
| 289 | else: |
| 290 | network.set_comm_fusion(int((layer_id + offset) / dis) + 1) |
| 291 | # Used for enabling recomputation of the block |
| 292 | if parallel_config.recompute: |
| 293 | network.recompute(recompute_slice_activation=True) |
| 294 | |
| 295 | |
| 296 | class PanguAlpha_Model(Cell): |
nothing calls this directly
no outgoing calls
no test coverage detected