MindSpore分布式并行报错The strategy is XXX, shape XXX cannot be divisible by strategy value XXX

1. 系统环境

硬件环境(Ascend/GPU/CPU): Ascend910Pro
MindSpore版本:不限
执行模式(PyNative/ Graph):Graph
Python版本:不限
操作系统平台:Linux

2. 报错信息

2.1 报错信息

[ERROR] PARALLEL(371,ffff9fe22bf0,python):2023-11-10-09:29:03.578.320 [mindspore/ccsrc/frontend/parallel/ops info/operator info.cc:180] CheckStrategyValue] AddInfo2323: The strategy is ((128, 1, 1), (128, 1, 1)), shape 32 cannot be divisible by strategy value 128  
[ERROR] PARALLEL(371,ffff9fe22bf0,python):2023-11-10-09:29:03.578.372 [mindspore/ccsrc/frontend/parallel/ops info/operator info.cc:916] InitForCostModelWithAutoRepeatCalc] AddInfo2323: CheckStrategy failed.  
[ERROR] PARALLEL(371,ffff9fe22bf0,python):2023-11-10-09:29:03.578.384 [mindspore/ccsrc/frontend/parallel/ops info/operator info.cc:880] Init] AddInfo2323 : Init failed.   
[CRITICAL] PARALLEL(371,ffff9fe22bf0,python):2023-11-10-09:29:03.580.559 [mindspore/ccsrc/frontend/parallel/step parallel.cc:1953] ExtractStrategyAndInit] Failure:operator Add init failed  
The function call stack:  
In file /opt/huawei/schedule-train/algorithm/src/pangu_alpha.py(664)/ output_states = output_states + embedding_random/  
In file /opt/huawei/schedule-train/algorithm/src/pangu_alpha.py(750)/ logits = self.network(tokens,/

2.2 脚本信息

def __init__(self, config):  
    super(PanguAlphaModel, self).__init__()  
    # Network head to get logits over vocabulary  
    copied_parallel_config = copy.deepcopy(config.parallel_config)  
    if copied_parallel_config.pipeline_stage > 1:  
        copied_parallel_config.vocab_emb_dp = False  
    self.head = PanGuHead(hidden_size=config.hidden_size,  
                            parallel_config=copied_parallel_config)  
    self.head.pipeline_stage = config.parallel_config.pipeline_stage - 1  
    self.backbone = PanguAlpha_Model(config)  
    self.backbone.embedding.word_embedding.embedding_table.add_pipeline_stage(self.head.pipeline_stage)  

def construct(self, input_ids, input_position, attention_mask,  
                init_reset=True, batch_valid_length=None):  
    output_states, word_table = self.backbone(input_ids, input_position, attention_mask,  
                                                init_reset, batch_valid_length)  
    # 省略embedding_random的定义  
    output_states = output_states + embedding_random  
    logits = self.head(output_states, word_table)  
    return logits

3. 根因分析

模型在16节点128卡上跑,model_parallel=8, data_parallel=16
根据调用栈,报错行为

output_states = output_states + embedding_random

output_states和embedding_random的shape均为[32, 4096, 5120]
这个加法没有配置切分策略,默认的切分策略是((128, 1, 1), (128, 1, 1)),输入的第一维是32,不能被切分成128份。

4. 解决方案

此处不涉及模型权重,应该按数据并行维度data_parallel=16进行切分。定义Add算子并配置正确的切分策略((dp, 1, 1), (dp, 1, 1))

def __init__(self, config):  
    super(PanguAlphaModel, self).__init__()  
    # Network head to get logits over vocabulary  
    copied_parallel_config = copy.deepcopy(config.parallel_config)  
    if copied_parallel_config.pipeline_stage > 1:  
        copied_parallel_config.vocab_emb_dp = False  
    self.head = PanGuHead(hidden_size=config.hidden_size,  
                            parallel_config=copied_parallel_config)  
    self.head.pipeline_stage = config.parallel_config.pipeline_stage - 1  
    self.backbone = PanguAlpha_Model(config)  
    self.backbone.embedding.word_embedding.embedding_table.add_pipeline_stage(self.head.pipeline_stage)  
    self.add = P.Add().shard((config.parallel_config.data_parallel, 1, 1), (config.parallel_config.data_parallel, 1, 1))  
    
def construct(self, input_ids, input_position, attention_mask,  
                init_reset=True, batch_valid_length=None):  
    output_states, word_table = self.backbone(input_ids, input_position, attention_mask,  
                                                init_reset, batch_valid_length)  
    # 省略embedding_random的定义  
    output_states = self.add(output_states, embedding_random)  
    logits = self.head(output_states, word_table)  
    return logits