MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / set_parallel_configure_for_layer

Function set_parallel_configure_for_layer

codegeex/mindspore/src/pangu_alpha.py:261–293  ·  view source on GitHub ↗

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
)

Source from the content-addressed store, hash-verified

259
260
261def 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
296class PanguAlpha_Model(Cell):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected