已合并
[Feature]:使能LayerwiseDisaggregated边云协同推理特性前提下, 支持配置使能多机推理特性和2p下发的计算图适配 #372
K先生创建于 2月4日
[Feature]:使能LayerwiseDisaggregated边云协同推理特性前提下, 支持配置使能多机推理特性和2p下发的计算图适配 #372
已合并
K先生创建于 2月4日
从已删除 :dev合入到Ascend/MindIE-LLMdev
共 20 个文件变更+349-213
@@ -296,13 +296,13 @@ std::vector<std::string> ConstructIntertensorList(const DecoderLayerParam &param
296 AddAttentionTensor(param, intermediateTensorList, deepseekV2IntermediateCandidates);296 AddAttentionTensor(param, intermediateTensorList, deepseekV2IntermediateCandidates);
297 if (param.ffnAllreduce || param.hasFfnComm) {297 if (param.ffnAllreduce || param.hasFfnComm) {
298 // 大ep场景下,开启h3p qkvdown dp时moe层不需要intermediate_mlp_out,最后一层除外298 // 大ep场景下,开启h3p qkvdown dp时moe层不需要intermediate_mlp_out,最后一层除外
299- if (!(param.enableQkvdownDp && !param.isLastLayer && param.ffnStreamNum > 1)) {299+ if (!(param.enableQkvdownDp && !param.isLastLayer && !param.isCloudLastLayer && param.ffnStreamNum > 1)) {
300 atb_speed::common::AddTensorToList(deepseekV2IntermediateCandidates,300 atb_speed::common::AddTensorToList(deepseekV2IntermediateCandidates,
301 "ffn_need_padding", intermediateTensorList);301 "ffn_need_padding", intermediateTensorList);
302 }302 }
303 }303 }
304 if (param.ffnAllGather) {304 if (param.ffnAllGather) {
305- if (!param.enableQkvdownDp || param.isLastLayer) {305+ if (!param.enableQkvdownDp || param.isLastLayer || param.isCloudLastLayer) {
306 atb_speed::common::AddTensorToList(deepseekV2IntermediateCandidates,306 atb_speed::common::AddTensorToList(deepseekV2IntermediateCandidates,
307 "ffn_allgather", intermediateTensorList);307 "ffn_allgather", intermediateTensorList);
308 }308 }
@@ -1268,8 +1268,8 @@ atb::Status SetMlpResidualAdd(atb::GraphParam &opGraph, const DecoderLayerParam
1268 ((param.hasAttnComm) && (param.hasFfnComm) ?1268 ((param.hasAttnComm) && (param.hasFfnComm) ?
1269 "intermediate_moe_out_with_shared_with_padding" : "intermediate_moe_out_with_shared")};1269 "intermediate_moe_out_with_shared_with_padding" : "intermediate_moe_out_with_shared")};
1270 mlpResidualAddOutTensorNames = {param.ffnAllGather || param.ffnReduceScatter ?1270 mlpResidualAddOutTensorNames = {param.ffnAllGather || param.ffnReduceScatter ?
1271- ((param.enableQkvdownDp && !param.isLastLayer) ? "out_decoder_layer" : "intermediate_mlp_out") :1271+ ((param.enableQkvdownDp && !param.isLastLayer && !param.isCloudLastLayer) ?
1272- "out_decoder_layer"};1272+ "out_decoder_layer" : "intermediate_mlp_out") : "out_decoder_layer"};
1273 }1273 }
1274 mlpResidualAddNode.inTensorIds = atb_speed::common::GetTensorIdxList(tensorMap, mlpResidualAddInTensorNames);1274 mlpResidualAddNode.inTensorIds = atb_speed::common::GetTensorIdxList(tensorMap, mlpResidualAddInTensorNames);
1275 mlpResidualAddNode.outTensorIds = atb_speed::common::GetTensorIdxList(tensorMap, mlpResidualAddOutTensorNames);1275 mlpResidualAddNode.outTensorIds = atb_speed::common::GetTensorIdxList(tensorMap, mlpResidualAddOutTensorNames);
@@ -1360,8 +1360,8 @@ int64_t SetMlpResidualAddNanToNum(atb::GraphParam &opGraph, const DecoderLayerPa
1360 "intermediate_dense_tp_rs_addout" : "out_decoder_layer"};1360 "intermediate_dense_tp_rs_addout" : "out_decoder_layer"};
1361 } else {1361 } else {
1362 mlpResidualAddOutTensorNames = {param.ffnAllGather || param.ffnReduceScatter ? \1362 mlpResidualAddOutTensorNames = {param.ffnAllGather || param.ffnReduceScatter ? \
1363- ((param.enableQkvdownDp && !param.isLastLayer) ? "out_decoder_layer" : \1363+ ((param.enableQkvdownDp && !param.isLastLayer && !param.isCloudLastLayer) ? \
1364- "intermediate_mlp_out") : "out_decoder_layer"};1364+ "out_decoder_layer" : "intermediate_mlp_out") : "out_decoder_layer"};
1365 }1365 }
1366 1366 
1367 nanToNumNode.inTensorIds = atb_speed::common::GetTensorIdxList(tensorMap, mlpResidualAddOutTensorNames);1367 nanToNumNode.inTensorIds = atb_speed::common::GetTensorIdxList(tensorMap, mlpResidualAddOutTensorNames);
@@ -1875,13 +1875,13 @@ atb::Status SetPostMoeProcess(std::map<std::string, uint32_t> &tensorMap,
1875 CHECK_OPERATION_STATUS_RETURN(SetFFNPadding(opGraph, param, tensorMap));1875 CHECK_OPERATION_STATUS_RETURN(SetFFNPadding(opGraph, param, tensorMap));
1876 }1876 }
1877 // h3p qkvdown dp move moe allgather to mla, without last moe1877 // h3p qkvdown dp move moe allgather to mla, without last moe
1878- if (!param.enableQkvdownDp || param.isLastLayer) {1878+ if (!param.enableQkvdownDp || param.isLastLayer || param.isCloudLastLayer) {
1879 CHECK_OPERATION_STATUS_RETURN(SetTPAllGatherNode(opGraph, param, tensorMap));1879 CHECK_OPERATION_STATUS_RETURN(SetTPAllGatherNode(opGraph, param, tensorMap));
1880 }1880 }
1881 }1881 }
1882 if (param.hasFfnComm) {1882 if (param.hasFfnComm) {
1883 // h3p qkvdown dp move moe gather to mla, without last moe1883 // h3p qkvdown dp move moe gather to mla, without last moe
1884- if (!param.enableQkvdownDp || param.isLastLayer) {1884+ if (!param.enableQkvdownDp || param.isLastLayer || param.isCloudLastLayer) {
1885 CHECK_OPERATION_STATUS_RETURN(SetFFNUnPadding(opGraph, param, tensorMap));1885 CHECK_OPERATION_STATUS_RETURN(SetFFNUnPadding(opGraph, param, tensorMap));
1886 }1886 }
1887 }1887 }
@@ -1906,8 +1906,8 @@ atb::Status DecoderLayer(DecoderLayerParam &param, atb::Operation **operation)
1906 CHECK_OPERATION_STATUS_RETURN(SetPostMoeProcess(tensorMap, param, opGraph));1906 CHECK_OPERATION_STATUS_RETURN(SetPostMoeProcess(tensorMap, param, opGraph));
1907 opGraph.inferShapeFunc = [=] (const atb::SVector<atb::TensorDesc> &inTensorDescs,1907 opGraph.inferShapeFunc = [=] (const atb::SVector<atb::TensorDesc> &inTensorDescs,
1908 atb::SVector<atb::TensorDesc> &outTensorDescs) {1908 atb::SVector<atb::TensorDesc> &outTensorDescs) {
1909- if ((param.mapping.Get(base::ATTN_DP).IsEnabled() || param.mapping.Get(base::ATTN_CP).IsEnabled()) && \1909+ if (((param.mapping.Get(base::ATTN_DP).IsEnabled() || param.mapping.Get(base::ATTN_CP).IsEnabled()) && \
1910- param.isLastLayer && !param.enableDpOut) {1910+ param.isLastLayer && !param.enableDpOut) || (param.enableQkvdownDp && param.isCloudLastLayer)) {
1911 outTensorDescs.at(0) = inTensorDescs.at(atb_speed::common::GetTensorIdx(tensorMap, "in_final_state"));1911 outTensorDescs.at(0) = inTensorDescs.at(atb_speed::common::GetTensorIdx(tensorMap, "in_final_state"));
1912 } else if (param.mapping.Get(base::ATTN_DP).IsEnabled() && param.isLastLayer && \1912 } else if (param.mapping.Get(base::ATTN_DP).IsEnabled() && param.isLastLayer && \
1913 param.enableDpOut && param.lmHeadLocalTp) {1913 param.enableDpOut && param.lmHeadLocalTp) {
@@ -63,6 +63,8 @@ public:
63 bool enableMlaPreprocess = false;63 bool enableMlaPreprocess = false;
64 bool isNzCache = false;64 bool isNzCache = false;
65 bool enablePrefixCache = false;65 bool enablePrefixCache = false;
66+ /// The following variables will only be used when layerwiseDisaggregated is enabled.
67+ bool isCloudLastLayer = false;
66 // 混合并行数据流68 // 混合并行数据流
67 int attnStreamNum = 1;69 int attnStreamNum = 1;
68 int ffnStreamNum = 1;70 int ffnStreamNum = 1;
@@ -327,6 +327,9 @@ void DeepseekV2ModelParam::ParseLayerwiseDisaggregatedParam(nlohmann::json &para
327 if (paramJson.contains("endLayerId")) {327 if (paramJson.contains("endLayerId")) {
328 endLayerId = atb_speed::base::FetchJsonParam<int32_t>(paramJson, "endLayerId");328 endLayerId = atb_speed::base::FetchJsonParam<int32_t>(paramJson, "endLayerId");
329 }329 }
330+ if (paramJson.contains("cloudLastLayerId")) {
331+ cloudLastLayerId = atb_speed::base::FetchJsonParam<int32_t>(paramJson, "cloudLastLayerId");
332+ }
330 333
331 this->isInternalLayer = this->layerwiseDisaggregated334 this->isInternalLayer = this->layerwiseDisaggregated
332 && (this->layerwiseMode == LWD_EDGE_FIRST || this->layerwiseMode == LWD_CLOUD_MIDDLE);335 && (this->layerwiseMode == LWD_EDGE_FIRST || this->layerwiseMode == LWD_CLOUD_MIDDLE);
@@ -504,20 +507,39 @@ void DecoderModel::ConstructInternalTensorMap()
504 "internal_tensor_cos_emb", "internal_tensor_sin_emb"507 "internal_tensor_cos_emb", "internal_tensor_sin_emb"
505 };508 };
506 }509 }
510+ if (param.mapping.Get(base::ATTN_DP).IsEnabled() || param.mapping.Get(base::ATTN_CP).IsEnabled()) {
511+ if (this->param.skipWordEmbedding && this->param.numHiddenLayers == 1 &&
512+ this->param.layerwiseMode == LWD_EDGE_LAST) {
513+ deepseekV2ModelInternalTensorCandidates["default"] = {
514+ "internal_tensor_cos_emb", "internal_tensor_sin_emb"
515+ };
516+ }
517+ }
507 }518 }
508 atb_speed::common::AssignTensorIdx(519 atb_speed::common::AssignTensorIdx(
509 deepseekV2ModelInternalTensorCandidates, "default", this->internalTensorMap);520 deepseekV2ModelInternalTensorCandidates, "default", this->internalTensorMap);
510 if (param.mapping.Get(base::ATTN_DP).IsEnabled() || param.mapping.Get(base::ATTN_CP).IsEnabled()) {521 if (param.mapping.Get(base::ATTN_DP).IsEnabled() || param.mapping.Get(base::ATTN_CP).IsEnabled()) {
511- atb_speed::common::AssignTensorIdx(522+ if (!this->param.layerwiseDisaggregated || this->param.layerwiseMode == LWD_EDGE_LAST) {
512- deepseekV2ModelInternalTensorCandidates, "last_layer", this->internalTensorMap);523+ atb_speed::common::AssignTensorIdx(
524+ deepseekV2ModelInternalTensorCandidates, "last_layer", this->internalTensorMap);
525+ }
513 }526 }
514 if (param.enableDpOut && param.lmHeadLocalTp) {527 if (param.enableDpOut && param.lmHeadLocalTp) {
515 atb_speed::common::AssignTensorIdx(528 atb_speed::common::AssignTensorIdx(
516 deepseekV2ModelInternalTensorCandidates, "enable_lm_head_local_tp_out", this->internalTensorMap);529 deepseekV2ModelInternalTensorCandidates, "enable_lm_head_local_tp_out", this->internalTensorMap);
517 }530 }
518 if (param.enableQkvdownDp) {531 if (param.enableQkvdownDp) {
519- atb_speed::common::AssignTensorIdx(532+ if (!this->param.layerwiseDisaggregated) {
520- deepseekV2ModelInternalTensorCandidates, "qkvdown_dp", this->internalTensorMap);533+ atb_speed::common::AssignTensorIdx(
534+ deepseekV2ModelInternalTensorCandidates, "qkvdown_dp", this->internalTensorMap);
535+ } else {
536+ if (this->param.layerwiseMode == LWD_EDGE_LAST ||
537+ (this->param.startLayerId >= this->param.firstKDenseReplace &&
538+ this->param.endLayerId - this->param.startLayerId >= 2)) { // moe层且中间层数大于2
539+ atb_speed::common::AssignTensorIdx(deepseekV2ModelInternalTensorCandidates,
540+ "qkvdown_dp", this->internalTensorMap);
541+ }
542+ }
521 }543 }
522}544}
523 545 
@@ -619,6 +641,38 @@ atb::TensorDesc DecoderModel::GetLogitsDesc(
619 return logitsDesc;641 return logitsDesc;
620}642}
621 643 
644+atb::TensorDesc DecoderModel::GetLWDLogitsDesc(
645+ const std::vector<atb::TensorDesc> &inTensorDescs)
646+{
647+ atb::TensorDesc logitsDesc;
648+ logitsDesc.dtype = this->param.isBF16 ? \
649+ aclDataType::ACL_BF16 : aclDataType::ACL_FLOAT16;
650+ logitsDesc.format = graph_.weightTensors.at(0).desc.format;
651+ if (this->param.layerwiseMode == LWD_EDGE_FIRST) {
652+ logitsDesc.shape.dimNum = inTensorDescs.at(0).shape.dimNum + 1;
653+ } else {
654+ logitsDesc.shape.dimNum = inTensorDescs.at(0).shape.dimNum;
655+ }
656+ logitsDesc.shape.dims[0] = inTensorDescs.at(0).shape.dims[0];
657+ logitsDesc.shape.dims[1] = this->param.hiddenSize;
658+ if (param.enableQkvdownDp && param.endLayerId > param.firstKDenseReplace &&
659+ param.startLayerId <= param.firstKDenseReplace) {
660+ logitsDesc.shape.dims[0] = inTensorDescs.at(
661+ atb_speed::common::GetTensorIdx(this->inTensorMap, "in_ffn_padding_idx_model")
662+ ).shape.dims[0];
663+ // 2: dynamic ep level
664+ if (param.expertParallelDegree != 2 && param.mapping.Get(base::MLP_TP).rankIds.size() != 0) {
665+ logitsDesc.shape.dims[0] /= param.mapping.Get(base::MLP_TP).rankIds.size();
666+ }
667+ }
668+ if (param.enableQkvdownDp && param.endLayerId == param.cloudLastLayerId + 1) {
669+ logitsDesc = inTensorDescs.at(
670+ atb_speed::common::GetTensorIdx(this->inTensorMap, "in_final_state_model"));
671+ }
672+ 
673+ return logitsDesc;
674+}
675+ 
622atb::Status DecoderModel::InferShape(676atb::Status DecoderModel::InferShape(
623 const std::vector<atb::TensorDesc> &inTensorDescs,677 const std::vector<atb::TensorDesc> &inTensorDescs,
624 std::vector<atb::TensorDesc> &outTensorDescs678 std::vector<atb::TensorDesc> &outTensorDescs
@@ -640,16 +694,7 @@ atb::Status DecoderModel::InferShape(
640 outTensorDescs.at(outTensorIdx) = GetLogitsDesc(inTensorDescs, logitsIndicesIdx);694 outTensorDescs.at(outTensorIdx) = GetLogitsDesc(inTensorDescs, logitsIndicesIdx);
641 } else {695 } else {
642 if (this->param.layerwiseMode == LWD_EDGE_FIRST || this->param.layerwiseMode == LWD_CLOUD_MIDDLE) {696 if (this->param.layerwiseMode == LWD_EDGE_FIRST || this->param.layerwiseMode == LWD_CLOUD_MIDDLE) {
643- outTensorDescs.at(outTensorIdx).dtype = this->param.isBF16 ? \697+ outTensorDescs.at(outTensorIdx) = GetLWDLogitsDesc(inTensorDescs);
644- aclDataType::ACL_BF16 : aclDataType::ACL_FLOAT16;
645- outTensorDescs.at(outTensorIdx).format = graph_.weightTensors.at(0).desc.format;
646- if (this->param.layerwiseMode == LWD_EDGE_FIRST) {
647- outTensorDescs.at(outTensorIdx).shape.dimNum = inTensorDescs.at(0).shape.dimNum + 1;
648- } else {
649- outTensorDescs.at(outTensorIdx).shape.dimNum = inTensorDescs.at(0).shape.dimNum;
650- }
651- outTensorDescs.at(outTensorIdx).shape.dims[0] = inTensorDescs.at(0).shape.dims[0];
652- outTensorDescs.at(outTensorIdx).shape.dims[1] = this->param.hiddenSize;
653 } else {698 } else {
654 outTensorDescs.at(outTensorIdx) = GetLogitsDesc(inTensorDescs, logitsIndicesIdx);699 outTensorDescs.at(outTensorIdx) = GetLogitsDesc(inTensorDescs, logitsIndicesIdx);
655 }700 }
@@ -930,7 +975,7 @@ void SetMoeParam(DecoderLayerParam &layerParam, const DeepseekV2ModelParam &para
930 975 
931void SetLayerwiseDisaggregatedParam(DecoderLayerParam &layerParam, const DeepseekV2ModelParam &param, int64_t layerId)976void SetLayerwiseDisaggregatedParam(DecoderLayerParam &layerParam, const DeepseekV2ModelParam &param, int64_t layerId)
932{977{
933- if (!param.layerwiseDisaggregated) {978+ if (!param.layerwiseDisaggregated) {
934 if (layerId == param.numHiddenLayers - 1) {979 if (layerId == param.numHiddenLayers - 1) {
935 layerParam.isLastLayer = true;980 layerParam.isLastLayer = true;
936 }981 }
@@ -950,6 +995,7 @@ void DecoderModel::SetLaywiseDisaggregatedQuantParam(DecoderLayerParam &layerPar
950 layerParam.attnLinearTransposeType = param.attnLinearTransposeType[layerId - param.startLayerId];995 layerParam.attnLinearTransposeType = param.attnLinearTransposeType[layerId - param.startLayerId];
951 layerParam.mlpLinearTransposeType = param.mlpLinearTransposeType[layerId - param.startLayerId];996 layerParam.mlpLinearTransposeType = param.mlpLinearTransposeType[layerId - param.startLayerId];
952 layerParam.moeLinearTransposeType = param.moeLinearTransposeType[layerId - param.startLayerId];997 layerParam.moeLinearTransposeType = param.moeLinearTransposeType[layerId - param.startLayerId];
998+ layerParam.isCloudLastLayer = param.enableQkvdownDp && layerId == param.cloudLastLayerId;
953}999}
954 1000 
955void DecoderModel::SetLayerParam(DecoderLayerParam &layerParam, int64_t layerId)1001void DecoderModel::SetLayerParam(DecoderLayerParam &layerParam, int64_t layerId)
@@ -1045,6 +1091,9 @@ atb::Status DecoderModel::AddSingleLayer(uint32_t layerId)
1045 }1091 }
1046 ATB_SPEED_LOG_DEBUG("start create Decoderlayer");1092 ATB_SPEED_LOG_DEBUG("start create Decoderlayer");
1047 CHECK_OPERATION_STATUS_RETURN(DecoderLayer(layerParam, &op));1093 CHECK_OPERATION_STATUS_RETURN(DecoderLayer(layerParam, &op));
1094+ if (this->param.layerwiseDisaggregated) { // DecoderLayer 中可能修改了 enableQkvdownDp
1095+ param.enableQkvdownDp = layerParam.enableQkvdownDp;
1096+ }
1048 ATB_SPEED_LOG_DEBUG("Decoderlayer create success");1097 ATB_SPEED_LOG_DEBUG("Decoderlayer create success");
1049 layerNode.operation.reset(op);1098 layerNode.operation.reset(op);
1050 ATB_SPEED_LOG_DEBUG("Decoderlayer inTensor number: " << layerNode.operation->GetInputNum());1099 ATB_SPEED_LOG_DEBUG("Decoderlayer inTensor number: " << layerNode.operation->GetInputNum());
@@ -129,6 +129,7 @@ private:
129 void ConstructOutTensorMap() override;129 void ConstructOutTensorMap() override;
130 atb::Status BindParamHostTensor(uint32_t nodeId) override;130 atb::Status BindParamHostTensor(uint32_t nodeId) override;
131 atb::TensorDesc GetLogitsDesc(const std::vector<atb::TensorDesc> &inTensorDescs, uint32_t logitsIndicesIdx);131 atb::TensorDesc GetLogitsDesc(const std::vector<atb::TensorDesc> &inTensorDescs, uint32_t logitsIndicesIdx);
132+ atb::TensorDesc GetLWDLogitsDesc(const std::vector<atb::TensorDesc> &inTensorDescs);
132 std::string GetLayerOutName(uint32_t layerId);133 std::string GetLayerOutName(uint32_t layerId);
133};134};
134REGISTER_MODEL(deepseekV2, DecoderModel);135REGISTER_MODEL(deepseekV2, DecoderModel);
@@ -131,10 +131,12 @@ public:
131 std::string moeEpRankTableFile = "";131 std::string moeEpRankTableFile = "";
132 std::string moeEpBackend = "";132 std::string moeEpBackend = "";
133 133 
134+ /// The following variables will only be used when layerwiseDisaggregated is enabled.
134 int32_t layerwiseMode = -1;135 int32_t layerwiseMode = -1;
135 int32_t hiddenSize = 0;136 int32_t hiddenSize = 0;
136 int32_t startLayerId = 0;137 int32_t startLayerId = 0;
137 int32_t endLayerId = 0;138 int32_t endLayerId = 0;
139+ int32_t cloudLastLayerId = 60;
138 bool layerwiseDisaggregated = false;140 bool layerwiseDisaggregated = false;
139 bool skipWordEmbedding = false;141 bool skipWordEmbedding = false;
140 bool isInternalLayer = false;142 bool isInternalLayer = false;
@@ -46,14 +46,14 @@ class LwdLayerStatus(int, Enum):
46 46 
47class LayerWiseAttr:47class LayerWiseAttr:
48 48 
49- __slot__ = ["start_num", "end_num", "split_type", "load_list", "ascend_weight_head",49+ __slot__ = ["edge_start_layer_count", "edge_end_layer_count", "split_type", "load_list", "ascend_weight_head",
50 "ascend_weight_tail", "ascend_weight_internal", "acl_inputs_prefill", 50 "ascend_weight_tail", "ascend_weight_internal", "acl_inputs_prefill",
51 "acl_inputs_decode", "acl_param_prefill", "acl_param_decode", "p_out_hidden",51 "acl_inputs_decode", "acl_param_prefill", "acl_param_decode", "p_out_hidden",
52- "weight_wrappers", "num_hidden_layers"]52+ "weight_wrappers", "num_hidden_layers", "acl_inputs_prefill_queue", "acl_param_prefill_queue"]
53 53 
54- def __init__(self, start_num, end_num, split_type):54+ def __init__(self, edge_start_layer_count, edge_end_layer_count, split_type):
55- self.start_num = start_num55+ self.edge_start_layer_count = edge_start_layer_count
56- self.end_num = end_num56+ self.edge_end_layer_count = edge_end_layer_count
57 self.split_type = split_type57 self.split_type = split_type
58 58 
59 59
@@ -82,15 +82,15 @@ class FlashForCausalLM(BaseModel):
82 layerwise_disaggregated = kwargs.get("layerwise_disaggregated", False)82 layerwise_disaggregated = kwargs.get("layerwise_disaggregated", False)
83 if layerwise_disaggregated:83 if layerwise_disaggregated:
84 split_type = None84 split_type = None
85- start_num = 185+ edge_start_layer_count = 1
86- end_num = 186+ edge_end_layer_count = 1
87 layerwise_disaggregated_role_type = kwargs.get("layerwise_disaggregated_role_type", "")87 layerwise_disaggregated_role_type = kwargs.get("layerwise_disaggregated_role_type", "")
88 self.layerwise_disaggregated = True88 self.layerwise_disaggregated = True
89 if layerwise_disaggregated_role_type == "slave":89 if layerwise_disaggregated_role_type == "slave":
90 split_type = DistributedType.CLOUD90 split_type = DistributedType.CLOUD
91 else:91 else:
92 split_type = DistributedType.EDGE92 split_type = DistributedType.EDGE
93- self.layerwise = LayerWiseAttr(start_num, end_num, split_type)93+ self.layerwise = LayerWiseAttr(edge_start_layer_count, edge_end_layer_count, split_type)
94 94
95 self.inference_mode = kwargs.get("inference_mode")95 self.inference_mode = kwargs.get("inference_mode")
96 96 
@@ -620,10 +620,3 @@ class FlashForCausalLM(BaseModel):
620 logits = self.execute_dap_ascend_operator(620 logits = self.execute_dap_ascend_operator(
621 all_inputs, json.dumps(acl_param_dict), is_prefill[0])621 all_inputs, json.dumps(acl_param_dict), is_prefill[0])
622 return logits622 return logits
623- 
624- def copy_input(self, output_buf, input_buf):
625- """save model input params when enable layerwise disaggregated"""
626- output_buf = []
627- for i in input_buf:
628- output_buf.append(i)
629- return output_buf
@@ -94,8 +94,8 @@ class LayerwiseCloudPrefillGraphWrapper(LayerwisePrefillGraphWrapper):
94 self.cosine_embed_tbl = None94 self.cosine_embed_tbl = None
95 95 
96 def set_param(self, model_type: str, params: Dict):96 def set_param(self, model_type: str, params: Dict):
97- self.graph_list = [torch.classes.ModelTorch.ModelTorch(model_type) 97+ layers_num = self.attr.num_hidden_layers - self.attr.edge_start_layer_count - self.attr.edge_end_layer_count
98- for _ in range(self.attr.num_hidden_layers - self.attr.start_num - self.attr.end_num)]98+ self.graph_list = [torch.classes.ModelTorch.ModelTorch(model_type) for _ in range(layers_num)]
99 params_list = params['layers']99 params_list = params['layers']
100 for graph, layer_params in zip(self.graph_list, params_list):100 for graph, layer_params in zip(self.graph_list, params_list):
101 graph.set_param(json.dumps({**layer_params, **self.feature_params}))101 graph.set_param(json.dumps({**layer_params, **self.feature_params}))
@@ -7,6 +7,7 @@
7# EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT,7# EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT,
8# MERCHANTABILITY OR FIT FOR A PARTICULAR PURPOSE.8# MERCHANTABILITY OR FIT FOR A PARTICULAR PURPOSE.
9# See the Mulan PSL v2 for more details.9# See the Mulan PSL v2 for more details.
10+import queue
10from typing import List11from typing import List
11from atb_llm.models.base.flash_causal_lm import LayerWiseAttr, LwdLayerStatus, DistributedType12from atb_llm.models.base.flash_causal_lm import LayerWiseAttr, LwdLayerStatus, DistributedType
12 13 
@@ -19,8 +20,12 @@ class LayerwiseModifier:
19 """20 """
20 def __init__(self, attr: LayerWiseAttr):21 def __init__(self, attr: LayerWiseAttr):
21 self.attr = attr22 self.attr = attr
22- self.acl_edge_inputs = [None, None]23+ self.acl_edge_decode_input = None
23- self.acl_edge_params = [None, None]24+ self.acl_edge_prefill_input = None
25+ self.acl_edge_prefill_input_queue = queue.Queue()
26+ self.acl_edge_decode_param = None
27+ self.acl_edge_prefill_param = None
28+ self.acl_edge_prefill_param_queue = queue.Queue()
24 self.acl_cloud_inputs = None29 self.acl_cloud_inputs = None
25 self.acl_cloud_params = None30 self.acl_cloud_params = None
26 self.acl_cloud_inner_hidden = None 31 self.acl_cloud_inner_hidden = None
@@ -35,6 +40,34 @@ class LayerwiseModifier:
35 @staticmethod 40 @staticmethod
36 def to_index(is_prefill):41 def to_index(is_prefill):
37 return 1 if is_prefill else 042 return 1 if is_prefill else 0
43+
44+ def get_input_param(self, is_prefill, is_end_layer):
45+ if is_prefill:
46+ if self.acl_edge_prefill_input is None:
47+ self.acl_edge_prefill_input = self.acl_edge_prefill_input_queue.get(timeout=900)
48+ self.acl_edge_prefill_param = self.acl_edge_prefill_param_queue.get(timeout=900)
49+ prefill_input = self.acl_edge_prefill_input
50+ prefill_param = self.acl_edge_prefill_param
51+ if is_end_layer:
52+ self.acl_edge_prefill_input = None
53+ self.acl_edge_prefill_param = None
54+ return prefill_input, prefill_param
55+ else:
56+ decode_input = self.acl_edge_decode_input
57+ decode_param = self.acl_edge_decode_param
58+ if is_end_layer:
59+ self.acl_edge_decode_input = None
60+ self.acl_edge_decode_param = None
61+ return decode_input, decode_param
62+ 
63+ def save_input_param(self, inputs, runtime_param, is_prefill):
64+ if is_prefill:
65+ self.acl_edge_prefill_input_queue.put([None] + inputs[1:])
66+ self.acl_edge_prefill_param_queue.put(runtime_param)
67+ else:
68+ # input[0] is hidden and needs to be replaced each time; no caching is required.
K先生
K先生K先生2月4日

1)加注释,说明为什么inputs的第0个需要替换成None 2)说明为什么需要copy

likedislike
K先生
K先生
2月4日 评论:
69+ self.acl_edge_decode_input = [None] + inputs[1:]
70+ self.acl_edge_decode_param = runtime_param
38 71 
39 def modify_inputs(72 def modify_inputs(
40 self,73 self,
@@ -52,31 +85,16 @@ class LayerwiseModifier:
52 if self.attr.split_type == DistributedType.EDGE:85 if self.attr.split_type == DistributedType.EDGE:
53 if exe_stage is None:86 if exe_stage is None:
54 return87 return
55- index = LayerwiseModifier.to_index(is_prefill)
56 if exe_stage.start_exec_layer == 0:88 if exe_stage.start_exec_layer == 0:
57- # 缓存输入在长序列场景下实现prefill阶段chunk穿插
58- if exe_stage.is_long_seq and is_prefill:
59- self.acl_edge_inputs_prefill_pre = self.acl_edge_inputs[index]
60- self.acl_edge_params_prefill_pre = self.acl_edge_params[index]
61 # 首层需要缓存输入89 # 首层需要缓存输入
62- self.acl_edge_inputs[index] = [None] + inputs[1:]90+ self.save_input_param(inputs, runtime_param, is_prefill)
63- self.acl_edge_params[index] = runtime_param.copy()
64 if exe_stage.end_exec_layer == 1:91 if exe_stage.end_exec_layer == 1:
65 # 尾层需要替换输入92 # 尾层需要替换输入
66- if exe_stage.is_long_seq and is_prefill and not exe_stage.end_of_generate_token:93+ last_input, last_param = self.get_input_param(is_prefill, True)
67- self.acl_edge_inputs_prefill_pre[0] = out_hidden94+ last_input[0] = out_hidden
68- inputs[:] = self.acl_edge_inputs_prefill_pre95+ inputs[:] = last_input
69- runtime_param.clear()96+ runtime_param.clear()
70- runtime_param.update(self.acl_edge_params_prefill_pre)97+ runtime_param.update(last_param)
71- else:
72- self.acl_edge_inputs[index][0] = out_hidden
73- inputs[:] = self.acl_edge_inputs[index]
74- runtime_param.clear()
75- runtime_param.update(self.acl_edge_params[index])
76- # 原则上需要清理,防止请求打完后内存泄漏
77- if exe_stage.end_of_generate_token:
78- self.acl_edge_inputs[index] = None
79- self.acl_edge_params[index] = None
80 else:98 else:
81 if exe_stage is None or not is_prefill:99 if exe_stage is None or not is_prefill:
82 inputs[0] = out_hidden100 inputs[0] = out_hidden
@@ -13,6 +13,7 @@ import os
13import json13import json
14import math14import math
15from enum import Enum15from enum import Enum
16+import queue
16from typing import List, Optional, Tuple17from typing import List, Optional, Tuple
17from dataclasses import asdict18from dataclasses import asdict
18 19 
@@ -188,12 +189,12 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
188 self.prefix_cache_enable = True189 self.prefix_cache_enable = True
189 self.layerwise.load_list = []190 self.layerwise.load_list = []
190 if self.layerwise.split_type == DistributedType.CLOUD:191 if self.layerwise.split_type == DistributedType.CLOUD:
191- start_layer = self.layerwise.start_num192+ start_layer = self.layerwise.edge_start_layer_count
192- end_layer = self.config.num_hidden_layers - self.layerwise.end_num193+ end_layer = self.config.num_hidden_layers - self.layerwise.edge_end_layer_count
193 self.layerwise.load_list = list(range(start_layer, end_layer))194 self.layerwise.load_list = list(range(start_layer, end_layer))
194 else:195 else:
195- self.layerwise.load_list = [i for i in range(0, self.layerwise.start_num)]196+ self.layerwise.load_list = [i for i in range(0, self.layerwise.edge_start_layer_count)]
196- start_layers = self.config.num_hidden_layers - self.layerwise.end_num197+ start_layers = self.config.num_hidden_layers - self.layerwise.edge_end_layer_count
197 other_load_list = [i for i in range(start_layers, self.config.num_hidden_layers)]198 other_load_list = [i for i in range(start_layers, self.config.num_hidden_layers)]
198 self.layerwise.load_list.extend(other_load_list)199 self.layerwise.load_list.extend(other_load_list)
199 self.model = FlashDeepseekV2Model(200 self.model = FlashDeepseekV2Model(
@@ -207,8 +208,10 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
207 self.layerwise.ascend_weight_internal = []208 self.layerwise.ascend_weight_internal = []
208 self.layerwise.acl_inputs_prefill = None209 self.layerwise.acl_inputs_prefill = None
209 self.layerwise.acl_inputs_decode = None210 self.layerwise.acl_inputs_decode = None
211+ self.layerwise.acl_inputs_prefill_queue = queue.Queue()
210 self.layerwise.acl_param_prefill = None212 self.layerwise.acl_param_prefill = None
211 self.layerwise.acl_param_decode = None213 self.layerwise.acl_param_decode = None
214+ self.layerwise.acl_param_prefill_queue = queue.Queue()
212 self.layerwise.p_out_hidden = None215 self.layerwise.p_out_hidden = None
213 self.layerwise.acl_inputs_prefill_pre = None216 self.layerwise.acl_inputs_prefill_pre = None
214 self.layerwise.acl_param_prefill_pre = None217 self.layerwise.acl_param_prefill_pre = None
@@ -550,6 +553,12 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
550 logger.warning(msg)553 logger.warning(msg)
551 config.models.deepseekv2.h3p.enable_shared_expert_overlap = False554 config.models.deepseekv2.h3p.enable_shared_expert_overlap = False
552 555 
556+ if self.layerwise_disaggregated:
557+ if self.layerwise.split_type != DistributedType.CLOUD:
558+ msg = "H3P moe dp optimization only takes effect on cloud side."
559+ logger.warning(msg)
560+ config.models.deepseekv2.h3p.enable_qkvdown_dp = False
561+ 
553 self.enable_qkvdown_dp = config.models.deepseekv2.h3p.enable_qkvdown_dp562 self.enable_qkvdown_dp = config.models.deepseekv2.h3p.enable_qkvdown_dp
554 self.enable_gating_dp = config.models.deepseekv2.h3p.enable_gating_dp563 self.enable_gating_dp = config.models.deepseekv2.h3p.enable_gating_dp
555 self.enable_shared_expert_dp = config.models.deepseekv2.h3p.enable_shared_expert_dp564 self.enable_shared_expert_dp = config.models.deepseekv2.h3p.enable_shared_expert_dp
@@ -622,10 +631,12 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
622 else:631 else:
623 self.acl_internal_decoder_operation = torch.classes.ModelTorch.ModelTorch(632 self.acl_internal_decoder_operation = torch.classes.ModelTorch.ModelTorch(
624 CPP_DEEPSEEKV2_MODEL_CLASS_NAME)633 CPP_DEEPSEEKV2_MODEL_CLASS_NAME)
634+ layers_num = self.config.num_hidden_layers - \
635+ self.layerwise.edge_start_layer_count - self.layerwise.edge_end_layer_count
625 self.encode_op_list = [torch.classes.ModelTorch.ModelTorch(CPP_DEEPSEEKV2_MODEL_CLASS_NAME) 636 self.encode_op_list = [torch.classes.ModelTorch.ModelTorch(CPP_DEEPSEEKV2_MODEL_CLASS_NAME)
626- for _ in range(self.config.num_hidden_layers - self.layerwise.start_num - self.layerwise.end_num)]637+ for _ in range(layers_num)]
627 self.encode_op_prefix_cache_list = [torch.classes.ModelTorch.ModelTorch(CPP_DEEPSEEKV2_MODEL_CLASS_NAME)638 self.encode_op_prefix_cache_list = [torch.classes.ModelTorch.ModelTorch(CPP_DEEPSEEKV2_MODEL_CLASS_NAME)
628- for _ in range(self.config.num_hidden_layers - self.layerwise.start_num - self.layerwise.end_num)]639+ for _ in range(layers_num)]
629 640 
630 def init_padding_idx(self):641 def init_padding_idx(self):
631 self.attn_padding_idx = self.placeholder642 self.attn_padding_idx = self.placeholder
@@ -777,16 +788,16 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
777 modify_ascend_params["layerwiseMode"] = mode788 modify_ascend_params["layerwiseMode"] = mode
778 if mode == 0:789 if mode == 0:
779 modify_ascend_params[START_ID] = 0790 modify_ascend_params[START_ID] = 0
780- modify_ascend_params[END_ID] = self.layerwise.start_num791+ modify_ascend_params[END_ID] = self.layerwise.edge_start_layer_count
781 modify_ascend_params[KVCACHE_QUANT_LAYERS] = \792 modify_ascend_params[KVCACHE_QUANT_LAYERS] = \
782- [self.kvcache_quant_layers[i] for i in range(self.layerwise.start_num)]793+ [self.kvcache_quant_layers[i] for i in range(self.layerwise.edge_start_layer_count)]
783 modify_ascend_params[MOE_PACK_QUANT_TYPE] = wrapper.moe_pack_type if wrapper.moe_pack_type else 0794 modify_ascend_params[MOE_PACK_QUANT_TYPE] = wrapper.moe_pack_type if wrapper.moe_pack_type else 0
784 elif mode == 1:795 elif mode == 1:
785- modify_ascend_params[START_ID] = self.layerwise.start_num796+ modify_ascend_params[START_ID] = self.layerwise.edge_start_layer_count
786- modify_ascend_params[END_ID] = self.config.num_hidden_layers - self.layerwise.end_num797+ modify_ascend_params[END_ID] = self.config.num_hidden_layers - self.layerwise.edge_end_layer_count
787- start_layer = self.layerwise.load_list.index(self.layerwise.start_num)798+ start_layer = self.layerwise.load_list.index(self.layerwise.edge_start_layer_count)
788 end_layer = self.layerwise.load_list.index(self.config.num_hidden_layers - \799 end_layer = self.layerwise.load_list.index(self.config.num_hidden_layers - \
789- self.layerwise.end_num - 1) + 1800+ self.layerwise.edge_end_layer_count - 1) + 1
790 modify_ascend_params[KVCACHE_QUANT_LAYERS] = [self.kvcache_quant_layers[i] \801 modify_ascend_params[KVCACHE_QUANT_LAYERS] = [self.kvcache_quant_layers[i] \
791 for i in range(start_layer, end_layer)]802 for i in range(start_layer, end_layer)]
792 if wrapper:803 if wrapper:
@@ -794,9 +805,10 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
794 else:805 else:
795 modify_ascend_params[MOE_PACK_QUANT_TYPE] = wrapper_list[-1].moe_pack_type806 modify_ascend_params[MOE_PACK_QUANT_TYPE] = wrapper_list[-1].moe_pack_type
796 elif mode == 2:807 elif mode == 2:
797- modify_ascend_params[START_ID] = self.config.num_hidden_layers - self.layerwise.end_num808+ modify_ascend_params[START_ID] = self.config.num_hidden_layers - self.layerwise.edge_end_layer_count
798 modify_ascend_params[END_ID] = self.config.num_hidden_layers809 modify_ascend_params[END_ID] = self.config.num_hidden_layers
799- start_layer = self.layerwise.load_list.index(self.config.num_hidden_layers - self.layerwise.end_num)810+ start_layer = self.layerwise.load_list.index(self.config.num_hidden_layers -
811+ self.layerwise.edge_end_layer_count)
800 end_layer = self.layerwise.load_list.index(self.config.num_hidden_layers - 1) + 1812 end_layer = self.layerwise.load_list.index(self.config.num_hidden_layers - 1) + 1
801 modify_ascend_params[KVCACHE_QUANT_LAYERS] = [self.kvcache_quant_layers[i] \813 modify_ascend_params[KVCACHE_QUANT_LAYERS] = [self.kvcache_quant_layers[i] \
802 for i in range(start_layer, end_layer)]814 for i in range(start_layer, end_layer)]
@@ -812,17 +824,18 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
812 end_layer = 0824 end_layer = 0
813 if mode == 0:825 if mode == 0:
814 start_layer = 0826 start_layer = 0
815- end_layer = self.layerwise.start_num827+ end_layer = self.layerwise.edge_start_layer_count
816 elif mode == 1:828 elif mode == 1:
817 if is_prefill:829 if is_prefill:
818 start_layer = self.layerwise.load_list.index(layer_no)830 start_layer = self.layerwise.load_list.index(layer_no)
819 end_layer = self.layerwise.load_list.index(layer_no) + 1831 end_layer = self.layerwise.load_list.index(layer_no) + 1
820 else:832 else:
821- start_layer = self.layerwise.load_list.index(self.layerwise.start_num)833+ start_layer = self.layerwise.load_list.index(self.layerwise.edge_start_layer_count)
822 end_layer = self.layerwise.load_list.index(834 end_layer = self.layerwise.load_list.index(
823- self.config.num_hidden_layers - self.layerwise.end_num - 1) + 1835+ self.config.num_hidden_layers - self.layerwise.edge_end_layer_count - 1) + 1
824 else:836 else:
825- start_layer = self.layerwise.load_list.index(self.config.num_hidden_layers - self.layerwise.end_num)837+ start_layer = self.layerwise.load_list.index(self.config.num_hidden_layers -
838+ self.layerwise.edge_end_layer_count)
826 end_layer = self.layerwise.load_list.index(self.config.num_hidden_layers - 1) + 1839 end_layer = self.layerwise.load_list.index(self.config.num_hidden_layers - 1) + 1
827 for i in range(start_layer, end_layer):840 for i in range(start_layer, end_layer):
828 layer = self.model.layers[i]841 layer = self.model.layers[i]
@@ -840,7 +853,8 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
840 else:853 else:
841 if self.layerwise.split_type == DistributedType.CLOUD:854 if self.layerwise.split_type == DistributedType.CLOUD:
842 self.layerwise.weight_wrappers = []855 self.layerwise.weight_wrappers = []
843- for i in range(self.layerwise.start_num, self.config.num_hidden_layers - self.layerwise.end_num):856+ for i in range(self.layerwise.edge_start_layer_count, self.config.num_hidden_layers -
857+ self.layerwise.edge_end_layer_count):
844 self.layerwise.weight_wrappers.append(858 self.layerwise.weight_wrappers.append(
845 self.get_layerwise_weights(mode=1, layer_no=i, is_prefill=True))859 self.get_layerwise_weights(mode=1, layer_no=i, is_prefill=True))
846 else:860 else:
@@ -1157,16 +1171,19 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
1157 self.acl_head_encoder_operation_prefixcache.set_weight(weight_wrapper_head.weights[:])1171 self.acl_head_encoder_operation_prefixcache.set_weight(weight_wrapper_head.weights[:])
1158 self.acl_tail_encoder_operation_prefixcache.set_weight(weight_wrapper_tail.weights[:])1172 self.acl_tail_encoder_operation_prefixcache.set_weight(weight_wrapper_tail.weights[:])
1159 else:1173 else:
1160- for layer in range(0, self.config.num_hidden_layers - self.layerwise.end_num - \1174+ for layer in range(0, self.config.num_hidden_layers - self.layerwise.edge_end_layer_count - \
1161- self.layerwise.start_num):1175+ self.layerwise.edge_start_layer_count):
1162 encoder_internal_param = self.get_layerwsie_ascend_param(1176 encoder_internal_param = self.get_layerwsie_ascend_param(
1163 encoder_param, 1, self.layerwise.weight_wrappers[layer]1177 encoder_param, 1, self.layerwise.weight_wrappers[layer]
1164 )1178 )
1165- encoder_internal_param[START_ID] = self.layerwise.start_num + layer1179+ encoder_internal_param[START_ID] = self.layerwise.edge_start_layer_count + layer
1166- encoder_internal_param[END_ID] = self.layerwise.start_num + layer + 11180+ encoder_internal_param[END_ID] = self.layerwise.edge_start_layer_count + layer + 1
1181+ encoder_internal_param["cloudLastLayerId"] = \
1182+ self.config.num_hidden_layers - self.layerwise.edge_end_layer_count - 1
1167 encoder_internal_param["numHiddenLayers"] = 11183 encoder_internal_param["numHiddenLayers"] = 1
1168 encoder_internal_param[KVCACHE_QUANT_LAYERS] = [1184 encoder_internal_param[KVCACHE_QUANT_LAYERS] = [
1169- self.kvcache_quant_layers[self.layerwise.load_list.index(self.layerwise.start_num + layer)]]1185+ self.kvcache_quant_layers[self.layerwise.load_list.index(
1186+ self.layerwise.edge_start_layer_count + layer)]]
1170 self.encode_op_list[layer].set_param(json.dumps({**encoder_internal_param}))1187 self.encode_op_list[layer].set_param(json.dumps({**encoder_internal_param}))
1171 self.encode_op_list[layer].set_weight(self.layerwise.weight_wrappers[layer].weights[:])1188 self.encode_op_list[layer].set_weight(self.layerwise.weight_wrappers[layer].weights[:])
1172 1189
@@ -1178,6 +1195,8 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
1178 decoder_internal_param = self.get_layerwsie_ascend_param(1195 decoder_internal_param = self.get_layerwsie_ascend_param(
1179 decoder_param, 1, wrapper_list=self.layerwise.weight_wrappers1196 decoder_param, 1, wrapper_list=self.layerwise.weight_wrappers
1180 )1197 )
1198+ decoder_internal_param["cloudLastLayerId"] = \
1199+ self.config.num_hidden_layers - self.layerwise.edge_end_layer_count - 1
1181 self.acl_internal_decoder_operation.set_param(json.dumps({**decoder_internal_param}))1200 self.acl_internal_decoder_operation.set_param(json.dumps({**decoder_internal_param}))
1182 self.acl_internal_decoder_operation.set_weight([weight_tensor \1201 self.acl_internal_decoder_operation.set_weight([weight_tensor \
1183 for wrapper in self.layerwise.weight_wrappers \1202 for wrapper in self.layerwise.weight_wrappers \
@@ -1633,9 +1652,14 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
1633 torch.npu.synchronize()1652 torch.npu.synchronize()
1634 perf_time_start = time.time()1653 perf_time_start = time.time()
1635 self.expert_array = self.placeholder1654 self.expert_array = self.placeholder
1636- final_hidden_states = torch.empty([self.token_size, self.config.hidden_size],1655+ 
1637- dtype=self.dtype,1656+ final_hidden_states_token_size = self.token_size
1638- device=input_ids.device)1657+ if self.layerwise_disaggregated and self.layerwise.split_type == DistributedType.CLOUD:
1658+ final_hidden_states_token_size = local_token_size
1659+ 
1660+ final_hidden_states = torch.empty([final_hidden_states_token_size, self.config.hidden_size],
1661+ dtype=self.dtype,
1662+ device=input_ids.device)
1639 1663 
1640 is_ep = (self.ep_level == ExpertParallelDegree.DYNAMIC_EP or \1664 is_ep = (self.ep_level == ExpertParallelDegree.DYNAMIC_EP or \
1641 (self.ep_level == ExpertParallelDegree.MIX_EP and is_prefill))1665 (self.ep_level == ExpertParallelDegree.MIX_EP and is_prefill))
@@ -2177,6 +2201,35 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
2177 acl_inputs_mtp.append(self.dense_allgather_unpad_idx)2201 acl_inputs_mtp.append(self.dense_allgather_unpad_idx)
2178 return acl_inputs_mtp2202 return acl_inputs_mtp
2179 2203 
2204+ def layerwise_get_input_param(self, is_prefill, is_end_layer):
2205+ if is_prefill:
2206+ if self.layerwise.acl_inputs_prefill is None:
2207+ self.layerwise.acl_inputs_prefill = self.layerwise.acl_inputs_prefill_queue.get(timeout=900)
2208+ self.layerwise.acl_param_prefill = self.layerwise.acl_param_prefill_queue.get(timeout=900)
2209+ prefill_input = self.layerwise.acl_inputs_prefill
2210+ prefill_param = self.layerwise.acl_param_prefill
2211+ if is_end_layer:
2212+ self.layerwise.acl_inputs_prefill = None
2213+ self.layerwise.acl_param_prefill = None
2214+ return prefill_input, prefill_param
2215+ else:
2216+ decode_input = self.layerwise.acl_inputs_decode
2217+ decode_param = self.layerwise.acl_param_decode
2218+ if is_end_layer:
2219+ self.layerwise.acl_inputs_decode = None
2220+ self.layerwise.acl_param_decode = None
2221+ return decode_input, decode_param
2222+ 
2223+ def layerwise_save_input_param(self, inputs, runtime_param, is_prefill):
2224+ # input[0] is hidden and needs to be replaced each time; no caching is required.
2225+ inputs_copy = [None] + inputs[1:]
2226+ if is_prefill:
2227+ self.layerwise.acl_inputs_prefill_queue.put(inputs_copy)
2228+ self.layerwise.acl_param_prefill_queue.put(runtime_param)
2229+ else:
2230+ self.layerwise.acl_inputs_decode = inputs_copy
2231+ self.layerwise.acl_param_decode = runtime_param
2232+
2180 def forward_layerwise_disaggregated_edge(2233 def forward_layerwise_disaggregated_edge(
2181 self,2234 self,
2182 input_ids: torch.Tensor,2235 input_ids: torch.Tensor,
@@ -2207,45 +2260,29 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
2207 if not is_prefill:2260 if not is_prefill:
2208 out_hidden = self.execute_ascend_operator(acl_inputs, acl_param, is_prefill,2261 out_hidden = self.execute_ascend_operator(acl_inputs, acl_param, is_prefill,
2209 split_part=LwdLayerStatus.EDGE_START_LAYER)2262 split_part=LwdLayerStatus.EDGE_START_LAYER)
2210- self.layerwise.acl_inputs_decode = self.copy_input(self.layerwise.acl_inputs_decode, acl_inputs)2263+ self.layerwise_save_input_param(acl_inputs, acl_param, is_prefill)
2211- self.layerwise.acl_param_decode = acl_param
2212- self.layerwise.acl_inputs_decode[0] = None
2213 acl_inputs = []2264 acl_inputs = []
2214 else:2265 else:
2215 out_hidden = self.execute_ascend_operator(acl_inputs, acl_param, is_prefill,2266 out_hidden = self.execute_ascend_operator(acl_inputs, acl_param, is_prefill,
2216 split_part=LwdLayerStatus.EDGE_START_LAYER)2267 split_part=LwdLayerStatus.EDGE_START_LAYER)
2217- if layerwise_disaggregated_exe_stage.is_long_seq and self.layerwise.acl_inputs_prefill is not None:2268+ self.layerwise_save_input_param(acl_inputs, acl_param, is_prefill)
2218- self.layerwise.acl_inputs_prefill_pre = self.copy_input(self.layerwise.acl_inputs_prefill_pre, \
2219- self.layerwise.acl_inputs_prefill)
2220- self.layerwise.acl_inputs_prefill_pre[0] = None
2221- self.layerwise.acl_param_prefill_pre = self.layerwise.acl_param_prefill
2222- self.layerwise.acl_inputs_prefill = self.copy_input(self.layerwise.acl_inputs_prefill, acl_inputs)
2223- self.layerwise.acl_param_prefill = acl_param
2224- self.layerwise.acl_inputs_prefill[0] = None
2225 acl_inputs = []2269 acl_inputs = []
2226 return out_hidden2270 return out_hidden
2227 if layerwise_disaggregated_exe_stage.end_exec_layer == 1:2271 if layerwise_disaggregated_exe_stage.end_exec_layer == 1:
2228 if not is_prefill:2272 if not is_prefill:
2229- self.layerwise.acl_inputs_decode[0] = out_hidden2273+ last_input, last_param = self.layerwise_get_input_param(is_prefill, True)
2230- logits = self.execute_ascend_operator(self.layerwise.acl_inputs_decode,2274+ last_input[0] = out_hidden
2231- self.layerwise.acl_param_decode, is_prefill,2275+ logits = self.execute_ascend_operator(last_input, last_param, is_prefill,
2232- split_part=LwdLayerStatus.EDGE_END_LAYER)2276+ split_part=LwdLayerStatus.EDGE_END_LAYER)
2233- self.layerwise.acl_inputs_decode = []
2234 else:2277 else:
2235 if layerwise_disaggregated_exe_stage.is_long_seq and \2278 if layerwise_disaggregated_exe_stage.is_long_seq and \
2236- layerwise_disaggregated_exe_stage.long_seq_start_idx != 0:2279+ layerwise_disaggregated_exe_stage.long_seq_start_idx != 0 and \
2280+ not layerwise_disaggregated_exe_stage.request_dp_empty:
2237 self.has_prefixcache = True2281 self.has_prefixcache = True
2238- if layerwise_disaggregated_exe_stage.is_long_seq and \2282+ last_input, last_param = self.layerwise_get_input_param(is_prefill, True)
2239- not layerwise_disaggregated_exe_stage.end_of_generate_token:2283+ last_input[0] = out_hidden
2240- self.layerwise.acl_inputs_prefill_pre[0] = out_hidden2284+ logits = self.execute_ascend_operator(last_input, last_param, is_prefill,
2241- logits = self.execute_ascend_operator(self.layerwise.acl_inputs_prefill_pre,2285+ split_part=LwdLayerStatus.EDGE_END_LAYER)
2242- self.layerwise.acl_param_prefill_pre, is_prefill,
2243- split_part=LwdLayerStatus.EDGE_END_LAYER)
2244- else:
2245- self.layerwise.acl_inputs_prefill[0] = out_hidden
2246- logits = self.execute_ascend_operator(self.layerwise.acl_inputs_prefill,
2247- self.layerwise.acl_param_prefill, is_prefill,
2248- split_part=LwdLayerStatus.EDGE_END_LAYER)
2249 return logits2286 return logits
2250 2287
2251 return out_hidden2288 return out_hidden
@@ -2302,13 +2339,14 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
2302 if i > layerwise_disaggregated_exe_stage.start_exec_layer:2339 if i > layerwise_disaggregated_exe_stage.start_exec_layer:
2303 acl_inputs[0] = out_hidden2340 acl_inputs[0] = out_hidden
2304 if layerwise_disaggregated_exe_stage.is_long_seq and \2341 if layerwise_disaggregated_exe_stage.is_long_seq and \
2305- layerwise_disaggregated_exe_stage.long_seq_start_idx != 0:2342+ layerwise_disaggregated_exe_stage.long_seq_start_idx != 0 and \
2343+ not layerwise_disaggregated_exe_stage.request_dp_empty:
2306 self.has_prefixcache = True2344 self.has_prefixcache = True
2307 out_hidden = self.execute_ascend_operator(acl_inputs, acl_param, is_prefill,2345 out_hidden = self.execute_ascend_operator(acl_inputs, acl_param, is_prefill,
2308 split_part=LwdLayerStatus.CLOUD_MIDDLE_LAYER, layer_index=i)2346 split_part=LwdLayerStatus.CLOUD_MIDDLE_LAYER, layer_index=i)
2309 self.layerwise.p_out_hidden = out_hidden2347 self.layerwise.p_out_hidden = out_hidden
2310- self.layerwise.acl_inputs_prefill = self.copy_input(self.layerwise.acl_inputs_prefill, acl_inputs)2348+ # acl_inputs[0] is hidden and needs to be replaced each time; no caching is required.
2311- self.layerwise.acl_inputs_prefill[0] = None2349+ self.layerwise.acl_inputs_prefill = [None] + acl_inputs[1:]
2312 self.layerwise.acl_param_prefill = acl_param2350 self.layerwise.acl_param_prefill = acl_param
2313 2351
2314 return out_hidden 2352 return out_hidden
@@ -2354,7 +2392,8 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
2354 self.free_operation_inputs(is_prefill)2392 self.free_operation_inputs(is_prefill)
2355 else:2393 else:
2356 if self.layerwise.split_type == DistributedType.CLOUD:2394 if self.layerwise.split_type == DistributedType.CLOUD:
2357- for i in range(self.config.num_hidden_layers - self.layerwise.start_num - self.layerwise.end_num):2395+ for i in range(self.config.num_hidden_layers - self.layerwise.edge_start_layer_count -
2396+ self.layerwise.edge_end_layer_count):
2358 acl_inputs[0] = out_hidden2397 acl_inputs[0] = out_hidden
2359 out_hidden = self.execute_ascend_operator(2398 out_hidden = self.execute_ascend_operator(
2360 acl_inputs, acl_param, is_prefill,2399 acl_inputs, acl_param, is_prefill,
@@ -2803,16 +2842,16 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
2803 if self.layerwise_disaggregated:2842 if self.layerwise_disaggregated:
2804 if self.layerwise.split_type == DistributedType.EDGE:2843 if self.layerwise.split_type == DistributedType.EDGE:
2805 start_cache_num = self.layerwise.load_list.index(2844 start_cache_num = self.layerwise.load_list.index(
2806- self.config.num_hidden_layers - self.layerwise.end_num)2845+ self.config.num_hidden_layers - self.layerwise.edge_end_layer_count)
2807 end_cache_num = self.layerwise.load_list.index(self.config.num_hidden_layers - 1) + 12846 end_cache_num = self.layerwise.load_list.index(self.config.num_hidden_layers - 1) + 1
2808- k_caches_sp1, k_caches_sp3 = [k_caches[i] for i in range(self.layerwise.start_num)], \2847+ k_caches_sp1, k_caches_sp3 = [k_caches[i] for i in range(self.layerwise.edge_start_layer_count)], \
2809 [k_caches[i] for i in range(start_cache_num, end_cache_num)]2848 [k_caches[i] for i in range(start_cache_num, end_cache_num)]
2810- v_caches_sp1, v_caches_sp3 = [v_caches[i] for i in range(self.layerwise.start_num)], \2849+ v_caches_sp1, v_caches_sp3 = [v_caches[i] for i in range(self.layerwise.edge_start_layer_count)], \
2811 [v_caches[i] for i in range(start_cache_num, end_cache_num)]2850 [v_caches[i] for i in range(start_cache_num, end_cache_num)]
2812 else:2851 else:
2813- start_cache_num = self.layerwise.load_list.index(self.layerwise.start_num)2852+ start_cache_num = self.layerwise.load_list.index(self.layerwise.edge_start_layer_count)
2814 end_cache_num = self.layerwise.load_list.index(2853 end_cache_num = self.layerwise.load_list.index(
2815- self.config.num_hidden_layers - self.layerwise.end_num - 1) + 12854+ self.config.num_hidden_layers - self.layerwise.edge_end_layer_count - 1) + 1
2816 k_caches_sp2 = [k_caches[i] for i in range(start_cache_num, end_cache_num)]2855 k_caches_sp2 = [k_caches[i] for i in range(start_cache_num, end_cache_num)]
2817 v_caches_sp2 = [v_caches[i] for i in range(start_cache_num, end_cache_num)]2856 v_caches_sp2 = [v_caches[i] for i in range(start_cache_num, end_cache_num)]
2818 if not self.layerwise_disaggregated:2857 if not self.layerwise_disaggregated:
@@ -2850,8 +2889,8 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM):
2850 self.acl_tail_encoder_operation_prefixcache.set_kv_cache(k_caches_sp3, v_caches_sp3)2889 self.acl_tail_encoder_operation_prefixcache.set_kv_cache(k_caches_sp3, v_caches_sp3)
2851 else:2890 else:
2852 self.acl_internal_decoder_operation.set_kv_cache(k_caches_sp2, v_caches_sp2)2891 self.acl_internal_decoder_operation.set_kv_cache(k_caches_sp2, v_caches_sp2)
2853- for layer in range(self.config.num_hidden_layers - self.layerwise.start_num - \2892+ for layer in range(self.config.num_hidden_layers - self.layerwise.edge_start_layer_count - \
2854- self.layerwise.end_num):2893+ self.layerwise.edge_end_layer_count):
2855 self.encode_op_list[layer].set_kv_cache([k_caches[layer]], [v_caches[layer]])2894 self.encode_op_list[layer].set_kv_cache([k_caches[layer]], [v_caches[layer]])
2856 self.encode_op_prefix_cache_list[layer].set_kv_cache([k_caches[layer]], [v_caches[layer]])2895 self.encode_op_prefix_cache_list[layer].set_kv_cache([k_caches[layer]], [v_caches[layer]])
2857 2896 
@@ -75,13 +75,14 @@ class FlashQwen2ForCausalLM(FlashForCausalLM):
75 self.inference_mode.enable_prefill_pa = True75 self.inference_mode.enable_prefill_pa = True
76 self.layerwise.load_list = []76 self.layerwise.load_list = []
77 if self.layerwise.split_type == DistributedType.CLOUD:77 if self.layerwise.split_type == DistributedType.CLOUD:
78- start_layer = self.layerwise.start_num78+ start_layer = self.layerwise.edge_start_layer_count
79- end_layer = self.config.num_hidden_layers - self.layerwise.end_num79+ end_layer = self.config.num_hidden_layers - self.layerwise.edge_end_layer_count
80 self.layerwise.load_list = list(range(start_layer, end_layer))80 self.layerwise.load_list = list(range(start_layer, end_layer))
81 else:81 else:
82- self.layerwise.load_list = [i for i in range(0, self.layerwise.start_num)]82+ self.layerwise.load_list = [i for i in range(0, self.layerwise.edge_start_layer_count)]
83- start_num = self.config.num_hidden_layers - self.layerwise.end_num83+ edge_start_layer_count = self.config.num_hidden_layers - self.layerwise.edge_end_layer_count
84- self.layerwise.load_list.extend([i for i in range(start_num, self.config.num_hidden_layers)])84+ self.layerwise.load_list.extend([i
85+ for i in range(edge_start_layer_count, self.config.num_hidden_layers)])
85 86 
86 self.transformer = FlashQwenModel(87 self.transformer = FlashQwenModel(
87 config, weights, model_prefix=model_prefix, lmhead_prefix=lmhead_prefix,88 config, weights, model_prefix=model_prefix, lmhead_prefix=lmhead_prefix,
@@ -307,17 +308,18 @@ class FlashQwen2ForCausalLM(FlashForCausalLM):
307 modify_ascend_params["layerwiseMode"] = mode308 modify_ascend_params["layerwiseMode"] = mode
308 modify_ascend_params["linearDescs"] = wrapper.linear_descs309 modify_ascend_params["linearDescs"] = wrapper.linear_descs
309 if mode == LwdLayerStatus.EDGE_START_LAYER:310 if mode == LwdLayerStatus.EDGE_START_LAYER:
310- modify_ascend_params[LINEAR_HAS_BIAS] = linear_has_bias * self.layerwise.start_num311+ modify_ascend_params[LINEAR_HAS_BIAS] = linear_has_bias * self.layerwise.edge_start_layer_count
311 modify_ascend_params[LWD_START_ID] = 0312 modify_ascend_params[LWD_START_ID] = 0
312- modify_ascend_params[LWD_END_ID] = self.layerwise.start_num313+ modify_ascend_params[LWD_END_ID] = self.layerwise.edge_start_layer_count
313 elif mode == LwdLayerStatus.CLOUD_MIDDLE_LAYER:314 elif mode == LwdLayerStatus.CLOUD_MIDDLE_LAYER:
314- modify_ascend_params[LINEAR_HAS_BIAS] = linear_has_bias * (self.config.num_hidden_layers - 315+ modify_ascend_params[LINEAR_HAS_BIAS] = linear_has_bias * (
315- self.layerwise.start_num - self.layerwise.end_num)316+ self.config.num_hidden_layers - self.layerwise.edge_start_layer_count -
316- modify_ascend_params[LWD_START_ID] = self.layerwise.start_num317+ self.layerwise.edge_end_layer_count)
317- modify_ascend_params[LWD_END_ID] = self.config.num_hidden_layers - self.layerwise.end_num318+ modify_ascend_params[LWD_START_ID] = self.layerwise.edge_start_layer_count
319+ modify_ascend_params[LWD_END_ID] = self.config.num_hidden_layers - self.layerwise.edge_end_layer_count
318 elif mode == LwdLayerStatus.EDGE_END_LAYER:320 elif mode == LwdLayerStatus.EDGE_END_LAYER:
319- modify_ascend_params[LINEAR_HAS_BIAS] = linear_has_bias * self.layerwise.end_num321+ modify_ascend_params[LINEAR_HAS_BIAS] = linear_has_bias * self.layerwise.edge_end_layer_count
320- modify_ascend_params[LWD_START_ID] = self.config.num_hidden_layers - self.layerwise.end_num322+ modify_ascend_params[LWD_START_ID] = self.config.num_hidden_layers - self.layerwise.edge_end_layer_count
321 modify_ascend_params[LWD_END_ID] = self.config.num_hidden_layers323 modify_ascend_params[LWD_END_ID] = self.config.num_hidden_layers
322 modify_ascend_params["numHiddenLayers"] = modify_ascend_params[LWD_END_ID] - modify_ascend_params[LWD_START_ID]324 modify_ascend_params["numHiddenLayers"] = modify_ascend_params[LWD_END_ID] - modify_ascend_params[LWD_START_ID]
323 325
@@ -352,17 +354,18 @@ class FlashQwen2ForCausalLM(FlashForCausalLM):
352 end_layer = 0354 end_layer = 0
353 if mode == LwdLayerStatus.EDGE_START_LAYER:355 if mode == LwdLayerStatus.EDGE_START_LAYER:
354 start_layer = 0356 start_layer = 0
355- end_layer = self.layerwise.start_num357+ end_layer = self.layerwise.edge_start_layer_count
356 elif mode == LwdLayerStatus.CLOUD_MIDDLE_LAYER:358 elif mode == LwdLayerStatus.CLOUD_MIDDLE_LAYER:
357 if is_prefill:359 if is_prefill:
358 start_layer = self.layerwise.load_list.index(layer_no)360 start_layer = self.layerwise.load_list.index(layer_no)
359 end_layer = self.layerwise.load_list.index(layer_no) + 1361 end_layer = self.layerwise.load_list.index(layer_no) + 1
360 else:362 else:
361- start_layer = self.layerwise.load_list.index(self.layerwise.start_num)363+ start_layer = self.layerwise.load_list.index(self.layerwise.edge_start_layer_count)
362 end_layer = self.layerwise.load_list.index(364 end_layer = self.layerwise.load_list.index(
363- self.config.num_hidden_layers - self.layerwise.end_num - 1) + 1365+ self.config.num_hidden_layers - self.layerwise.edge_end_layer_count - 1) + 1
364 else:366 else:
365- start_layer = self.layerwise.load_list.index(self.config.num_hidden_layers - self.layerwise.end_num)367+ start_layer = self.layerwise.load_list.index(self.config.num_hidden_layers -
368+ self.layerwise.edge_end_layer_count)
366 end_layer = self.layerwise.load_list.index(self.config.num_hidden_layers - 1) + 1369 end_layer = self.layerwise.load_list.index(self.config.num_hidden_layers - 1) + 1
367 for i in range(start_layer, end_layer):370 for i in range(start_layer, end_layer):
368 layer = self.transformer.h[i]371 layer = self.transformer.h[i]
@@ -403,7 +406,8 @@ class FlashQwen2ForCausalLM(FlashForCausalLM):
403 else:406 else:
404 if self.layerwise.split_type == DistributedType.CLOUD:407 if self.layerwise.split_type == DistributedType.CLOUD:
405 self.layerwise.weight_wrappers = []408 self.layerwise.weight_wrappers = []
406- for i in range(self.layerwise.start_num, self.config.num_hidden_layers - self.layerwise.end_num):409+ for i in range(self.layerwise.edge_start_layer_count,
410+ self.config.num_hidden_layers - self.layerwise.edge_end_layer_count):
407 self.layerwise.weight_wrappers.append(self.get_layerwise_weights(411 self.layerwise.weight_wrappers.append(self.get_layerwise_weights(
408 mode=LwdLayerStatus.CLOUD_MIDDLE_LAYER, layer_no=i, is_prefill=True))412 mode=LwdLayerStatus.CLOUD_MIDDLE_LAYER, layer_no=i, is_prefill=True))
409 decode_weight_wapper = self.get_layerwise_weights(mode=LwdLayerStatus.CLOUD_MIDDLE_LAYER,413 decode_weight_wapper = self.get_layerwise_weights(mode=LwdLayerStatus.CLOUD_MIDDLE_LAYER,
@@ -566,11 +570,11 @@ class FlashQwen2ForCausalLM(FlashForCausalLM):
566 params_list = []570 params_list = []
567 weights_list = []571 weights_list = []
568 for layer in range(0, self.config.num_hidden_layers - \572 for layer in range(0, self.config.num_hidden_layers - \
569- self.layerwise.end_num - self.layerwise.start_num):573+ self.layerwise.edge_end_layer_count - self.layerwise.edge_start_layer_count):
570 encoder_internal_param = self.get_layerwsie_ascend_param(encoder_param, 1,574 encoder_internal_param = self.get_layerwsie_ascend_param(encoder_param, 1,
571 linear_has_bias, self.layerwise.weight_wrappers[layer])575 linear_has_bias, self.layerwise.weight_wrappers[layer])
572- encoder_internal_param[LWD_START_ID] = self.layerwise.start_num + layer576+ encoder_internal_param[LWD_START_ID] = self.layerwise.edge_start_layer_count + layer
573- encoder_internal_param[LWD_END_ID] = self.layerwise.start_num + layer + 1577+ encoder_internal_param[LWD_END_ID] = self.layerwise.edge_start_layer_count + layer + 1
574 encoder_internal_param["numHiddenLayers"] = 1578 encoder_internal_param["numHiddenLayers"] = 1
575 encoder_internal_param[LINEAR_HAS_BIAS] = linear_has_bias579 encoder_internal_param[LINEAR_HAS_BIAS] = linear_has_bias
576 encoder_internal_param["reuseEmbedTable"] = self.long_seq_enable and layer != 0580 encoder_internal_param["reuseEmbedTable"] = self.long_seq_enable and layer != 0
@@ -767,11 +771,11 @@ class FlashQwen2ForCausalLM(FlashForCausalLM):
767 if self.layerwise_disaggregated:771 if self.layerwise_disaggregated:
768 if self.layerwise.split_type == DistributedType.EDGE:772 if self.layerwise.split_type == DistributedType.EDGE:
769 start_cache_num = self.layerwise.load_list.index(773 start_cache_num = self.layerwise.load_list.index(
770- self.config.num_hidden_layers - self.layerwise.end_num)774+ self.config.num_hidden_layers - self.layerwise.edge_end_layer_count)
771 end_cache_num = self.layerwise.load_list.index(self.config.num_hidden_layers - 1) + 1775 end_cache_num = self.layerwise.load_list.index(self.config.num_hidden_layers - 1) + 1
772- k_caches_sp1, k_caches_sp3 = [k_caches[i] for i in range(self.layerwise.start_num)], \776+ k_caches_sp1, k_caches_sp3 = [k_caches[i] for i in range(self.layerwise.edge_start_layer_count)], \
773 [k_caches[i] for i in range(start_cache_num, end_cache_num)]777 [k_caches[i] for i in range(start_cache_num, end_cache_num)]
774- v_caches_sp1, v_caches_sp3 = [v_caches[i] for i in range(self.layerwise.start_num)], \778+ v_caches_sp1, v_caches_sp3 = [v_caches[i] for i in range(self.layerwise.edge_start_layer_count)], \
775 [v_caches[i] for i in range(start_cache_num, end_cache_num)]779 [v_caches[i] for i in range(start_cache_num, end_cache_num)]
776 layerwise_k_caches = {780 layerwise_k_caches = {
777 LWD_HEAD: k_caches_sp1,781 LWD_HEAD: k_caches_sp1,
@@ -783,9 +787,9 @@ class FlashQwen2ForCausalLM(FlashForCausalLM):
783 }787 }
784 self.graph_manager.set_kv_cache(layerwise_k_caches, layerwise_v_caches)788 self.graph_manager.set_kv_cache(layerwise_k_caches, layerwise_v_caches)
785 else:789 else:
786- start_cache_num = self.layerwise.load_list.index(self.layerwise.start_num)790+ start_cache_num = self.layerwise.load_list.index(self.layerwise.edge_start_layer_count)
787 end_cache_num = self.layerwise.load_list.index(791 end_cache_num = self.layerwise.load_list.index(
788- self.config.num_hidden_layers - self.layerwise.end_num - 1) + 1792+ self.config.num_hidden_layers - self.layerwise.edge_end_layer_count - 1) + 1
789 k_caches_sp2 = [k_caches[i] 793 k_caches_sp2 = [k_caches[i]
790 for i in range(start_cache_num, end_cache_num)]794 for i in range(start_cache_num, end_cache_num)]
791 v_caches_sp2 = [v_caches[i] 795 v_caches_sp2 = [v_caches[i]
@@ -798,8 +802,8 @@ class FlashQwen2ForCausalLM(FlashForCausalLM):
798 LWD_LAYERS: [],802 LWD_LAYERS: [],
799 },803 },
800 }804 }
801- for layer in range(self.config.num_hidden_layers - self.layerwise.start_num - \805+ for layer in range(self.config.num_hidden_layers - self.layerwise.edge_start_layer_count - \
802- self.layerwise.end_num):806+ self.layerwise.edge_end_layer_count):
803 encoder_kv_caches['k']['layers'].append([k_caches[layer]])807 encoder_kv_caches['k']['layers'].append([k_caches[layer]])
804 encoder_kv_caches['v']['layers'].append([v_caches[layer]])808 encoder_kv_caches['v']['layers'].append([v_caches[layer]])
805 specified_kv_caches = {809 specified_kv_caches = {
@@ -273,10 +273,11 @@ class ModelRunner:
273 TLS_CRL_PATH: kwargs.get(TLS_CRL_PATH, ''),273 TLS_CRL_PATH: kwargs.get(TLS_CRL_PATH, ''),
274 TLS_CRL_FILES: kwargs.get(TLS_CRL_FILES, ''),274 TLS_CRL_FILES: kwargs.get(TLS_CRL_FILES, ''),
275 }275 }
276- self.data_comm = EdgeCloudDataComm(self.dtype)276+ batch_p_num = kwargs.get('batch_p_num', 1)
277+ self.data_comm = EdgeCloudDataComm(self.dtype, batch_p_num)
277 self.ctrl_comm = EdgeCloudCtrlComm(tls_config)278 self.ctrl_comm = EdgeCloudCtrlComm(tls_config)
278- self.time_counter = CloudCutPolicy(self.layerwise_disaggregated_role_type, model_name_or_path)279+ self.time_counter = CloudCutPolicy(self.layerwise_disaggregated_role_type, model_name_or_path, batch_p_num)
279- self.chunk_prefill_manager = ChunkPrefilPolicy(model_name_or_path)280+ self.chunk_prefill_manager = ChunkPrefilPolicy(model_name_or_path, batch_p_num)
280 self.prefill_input_lengths = None281 self.prefill_input_lengths = None
281 self.edge_pre_chunk_length = 0282 self.edge_pre_chunk_length = 0
282 self.prefill_total_seq_len = 0283 self.prefill_total_seq_len = 0
@@ -590,7 +591,7 @@ class ModelRunner:
590 self.data_comm.prefill_seq_len_queue.put(self.edge_pre_chunk_length)591 self.data_comm.prefill_seq_len_queue.put(self.edge_pre_chunk_length)
591 logger.info(f"[layerwiseDisaggregated] edge rank {self.rank}, put {self.edge_pre_chunk_length}")592 logger.info(f"[layerwiseDisaggregated] edge rank {self.rank}, put {self.edge_pre_chunk_length}")
592 self.edge_pre_chunk_length = int(hidden.shape[0])593 self.edge_pre_chunk_length = int(hidden.shape[0])
593- if layerwise_disaggregated_exe_stage.long_seq_end_idx == self.prefill_total_seq_len:594+ if layerwise_disaggregated_exe_stage.is_last_chunk:
594 self.data_comm.prefill_seq_len_queue.put(self.edge_pre_chunk_length)595 self.data_comm.prefill_seq_len_queue.put(self.edge_pre_chunk_length)
595 logger.info(f"[layerwiseDisaggregated] edge rank {self.rank}, "596 logger.info(f"[layerwiseDisaggregated] edge rank {self.rank}, "
596 f"end put {self.edge_pre_chunk_length}")597 f"end put {self.edge_pre_chunk_length}")
@@ -604,6 +605,8 @@ class ModelRunner:
604 layerwise_disaggregated_exe_stage.long_seq_start_idx == 0):605 layerwise_disaggregated_exe_stage.long_seq_start_idx == 0):
605 self.data_comm.p_shape[self.data_comm.recv_index] = self.data_comm.prefill_seq_len_queue.get()606 self.data_comm.p_shape[self.data_comm.recv_index] = self.data_comm.prefill_seq_len_queue.get()
606 self.data_comm.recv_hidden('p', self.data_comm.p_shape)607 self.data_comm.recv_hidden('p', self.data_comm.p_shape)
608+ logger.info(f"[layerwiseDisaggregated] edge rank {self.rank} prefill recv start first part, "
609+ f"the data length is: {self.data_comm.p_shape}")
607 return hidden610 return hidden
608 else:611 else:
609 tmp = self.data_comm.data_wait_after_recv('p')612 tmp = self.data_comm.data_wait_after_recv('p')
@@ -617,7 +620,7 @@ class ModelRunner:
617 if not self.data_comm.prefill_seq_len_queue.empty():620 if not self.data_comm.prefill_seq_len_queue.empty():
618 self.data_comm.p_shape[self.data_comm.recv_index] = self.data_comm.prefill_seq_len_queue.get()621 self.data_comm.p_shape[self.data_comm.recv_index] = self.data_comm.prefill_seq_len_queue.get()
619 self.data_comm.recv_hidden('p', self.data_comm.p_shape)622 self.data_comm.recv_hidden('p', self.data_comm.p_shape)
620- logger.info(f"[layerwiseDisaggregated] edge rank {self.rank} prefill recv start, "623+ logger.info(f"[layerwiseDisaggregated] edge rank {self.rank} prefill recv start post part, "
621 f"self.data_comm.p_shape: {self.data_comm.p_shape}")624 f"self.data_comm.p_shape: {self.data_comm.p_shape}")
622 625
623 return res626 return res
@@ -773,9 +776,10 @@ class ModelRunner:
773 hidden = self.data_comm.broadcast_hidden(tmp, self.data_comm.p_shape, 'p')776 hidden = self.data_comm.broadcast_hidden(tmp, self.data_comm.p_shape, 'p')
774 logger.info(f"[layerwiseDisaggregated] cloud rank {self.rank} prefill recv {hidden.shape}")777 logger.info(f"[layerwiseDisaggregated] cloud rank {self.rank} prefill recv {hidden.shape}")
775 if layerwise_disaggregated_exe_stage.is_long_seq and \778 if layerwise_disaggregated_exe_stage.is_long_seq and \
776- layerwise_disaggregated_exe_stage.long_seq_end_idx != self.prefill_total_seq_len:779+ not layerwise_disaggregated_exe_stage.is_last_chunk: # 不是最后一个序列
777 prefill_seq_len = layerwise_disaggregated_exe_stage.long_seq_next_end_idx - \780 prefill_seq_len = layerwise_disaggregated_exe_stage.long_seq_next_end_idx - \
778 layerwise_disaggregated_exe_stage.long_seq_end_idx781 layerwise_disaggregated_exe_stage.long_seq_end_idx
782+ prefill_seq_len = prefill_seq_len if prefill_seq_len > 0 else 1 # 至少长度应为1
779 self.data_comm.prefill_seq_len_queue.put(prefill_seq_len)783 self.data_comm.prefill_seq_len_queue.put(prefill_seq_len)
780 logger.info(f"[layerwiseDisaggregated] cloud rank {self.rank}, "784 logger.info(f"[layerwiseDisaggregated] cloud rank {self.rank}, "
781 f"queue input is {prefill_seq_len}")785 f"queue input is {prefill_seq_len}")
@@ -827,9 +831,9 @@ class ModelRunner:
827 input_ids = kwargs.get("input_ids")831 input_ids = kwargs.get("input_ids")
828 is_prefill = kwargs.get("is_prefill")832 is_prefill = kwargs.get("is_prefill")
829 if is_prefill:833 if is_prefill:
830- batch_size = len(input_lengths)834+ batch_size = len(input_lengths) * self.mapping.attn_dp.group_size
831 else:835 else:
832- batch_size = len(input_ids)836+ batch_size = len(input_ids) * self.mapping.attn_dp.group_size
833 out_dict = {OUT_HIDDEN: torch.ones([len(input_ids), self.model.hidden_size],837 out_dict = {OUT_HIDDEN: torch.ones([len(input_ids), self.model.hidden_size],
834 dtype=self.dtype, device=self.device)}838 dtype=self.dtype, device=self.device)}
835 kwargs.update(out_dict)839 kwargs.update(out_dict)
@@ -22,10 +22,11 @@ class ChunkPrefilPolicy():
22 cls._instance = super(ChunkPrefilPolicy, cls).__new__(cls)22 cls._instance = super(ChunkPrefilPolicy, cls).__new__(cls)
23 return cls._instance23 return cls._instance
24 24 
25- def __init__(self, model_name_or_path='qwen'):25+ def __init__(self, model_name_or_path='qwen', batch_p_num=1):
26 self.soc_name = acl.get_soc_name()26 self.soc_name = acl.get_soc_name()
27 if not hasattr(self, 'initialized'):27 if not hasattr(self, 'initialized'):
28 self.model_type = self.__get_model_name(model_name_or_path)28 self.model_type = self.__get_model_name(model_name_or_path)
29+ self.batch_p_num = batch_p_num
29 # For NPU Soc is Ascend910B2 or other models, use the following default prefill_chunk_map30 # For NPU Soc is Ascend910B2 or other models, use the following default prefill_chunk_map
30 self.prefill_chunk_map = {128: 33, 64: 20, 32: 10, 16: 6, 8: 2}31 self.prefill_chunk_map = {128: 33, 64: 20, 32: 10, 16: 6, 8: 2}
31 self.__ajust_prefill_chunk_map_for_diff_npu_soc()32 self.__ajust_prefill_chunk_map_for_diff_npu_soc()
@@ -54,8 +55,14 @@ class ChunkPrefilPolicy():
54 if self.soc_name == 'Ascend910B2':55 if self.soc_name == 'Ascend910B2':
55 return56 return
56 if self.soc_name == 'Ascend910B3':57 if self.soc_name == 'Ascend910B3':
57- self.prefill_chunk_map = {128: 33, 64: 20, 32: 10, 16: 6, 8: 2}58+ if self.batch_p_num == 1:
59+ self.prefill_chunk_map = {128: 33, 64: 20, 32: 10, 16: 6, 8: 2}
60+ else:
61+ self.prefill_chunk_map = {128: 33, 64: 20, 32: 10, 16: 6, 8: 2}
58 return62 return
59 if self.soc_name == 'Ascend910B4':63 if self.soc_name == 'Ascend910B4':
60- self.prefill_chunk_map = {128: 33, 64: 20, 32: 10, 16: 6, 8: 2}64+ if self.batch_p_num == 1:
65+ self.prefill_chunk_map = {128: 33, 64: 20, 32: 10, 16: 6, 8: 2}
66+ else:
67+ self.prefill_chunk_map = {128: 33, 64: 20, 32: 10, 16: 6, 8: 2}
61 return68 return
@@ -47,13 +47,15 @@ class CloudCutPolicy():
47 cls._instance = super(CloudCutPolicy, cls).__new__(cls)47 cls._instance = super(CloudCutPolicy, cls).__new__(cls)
48 return cls._instance48 return cls._instance
49 49 
50- def __init__(self, name="", model_name_or_path='qwen'):50+ def __init__(self, name="", model_name_or_path='qwen', batch_p_num=1):
51 if not hasattr(self, 'initialized'):51 if not hasattr(self, 'initialized'):
52 self.name = name52 self.name = name
53 self.model_type = self.__get_model_name(model_name_or_path)53 self.model_type = self.__get_model_name(model_name_or_path)
54 self.role_type = CloudCutClassType.CLOUD if name == "slave" else CloudCutClassType.OTHER54 self.role_type = CloudCutClassType.CLOUD if name == "slave" else CloudCutClassType.OTHER
55 self.rank_id = None55 self.rank_id = None
56 self.soc_name = acl.get_soc_name()56 self.soc_name = acl.get_soc_name()
57+ self.batch_p_num = batch_p_num
58+ self.multi_nodes_enable = False
57 self.initialized = False59 self.initialized = False
58 60 
59 # Predict the number of chunks, hardcoded based on empirical values: n(K): [cut_num, cut_num_max];61 # Predict the number of chunks, hardcoded based on empirical values: n(K): [cut_num, cut_num_max];
@@ -115,11 +117,14 @@ class CloudCutPolicy():
115 return CloudCutModelType.DEEP_SEEK117 return CloudCutModelType.DEEP_SEEK
116 return CloudCutModelType.QWEN118 return CloudCutModelType.QWEN
117 119 
118- def initialize(self, name, rank_id, max_cut_num, min_cut_num):120+ def initialize(self, name, rank_id, max_cut_num, min_cut_num, multi_nodes_enable):
119 self.role_type = CloudCutClassType.CLOUD if name == "slave" else CloudCutClassType.OTHER121 self.role_type = CloudCutClassType.CLOUD if name == "slave" else CloudCutClassType.OTHER
120 self.rank_id = rank_id122 self.rank_id = rank_id
121 self.max_cut_num = max_cut_num123 self.max_cut_num = max_cut_num
122 self.min_cut_num = min_cut_num124 self.min_cut_num = min_cut_num
125+ self.multi_nodes_enable = multi_nodes_enable
126+ if self.model_type == CloudCutModelType.DEEP_SEEK and self.multi_nodes_enable:
127+ self.__ajust_prefill_cut_num_for_multi_nodes()
123 self.initialized = True128 self.initialized = True
124 129 
125 def get_cut_num(self, input_data: CloudCutInputData):130 def get_cut_num(self, input_data: CloudCutInputData):
@@ -265,16 +270,37 @@ class CloudCutPolicy():
265 def __ajust_prefill_cut_num_for_diff_npu_soc(self):270 def __ajust_prefill_cut_num_for_diff_npu_soc(self):
266 if self.soc_name == 'Ascend910B2':271 if self.soc_name == 'Ascend910B2':
267 logger.info(f"[layerwiseDisaggregated] npu soc is Ascend910B2, ajust prefill cut num.")272 logger.info(f"[layerwiseDisaggregated] npu soc is Ascend910B2, ajust prefill cut num.")
273+ if self.batch_p_num != 1:
274+ self.prefill_default_cut_map = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 6, 0: 8}
275+ self.prefill_cut_num_max = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 5, 0: 8}
276+ self.prefill_cut_num_min = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 5, 3: 6, 2: 5, 1: 5, 0: 8}
268 return277 return
269 if self.soc_name == 'Ascend910B3':278 if self.soc_name == 'Ascend910B3':
270 logger.info(f"[layerwiseDisaggregated] npu soc is Ascend910B3, ajust prefill cut num.")279 logger.info(f"[layerwiseDisaggregated] npu soc is Ascend910B3, ajust prefill cut num.")
271- self.prefill_default_cut_map = {128: 330, 64: 120, 32: 70, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 6, 0: 8}280+ if self.batch_p_num == 1:
272- self.prefill_cut_num_max = {128: 330, 64: 120, 32: 70, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 5, 0: 8}281+ self.prefill_default_cut_map = {128: 330, 64: 120, 32: 70, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 6, 0: 8}
273- self.prefill_cut_num_min = {128: 330, 64: 120, 32: 70, 16: 24, 8: 10, 4: 4, 3: 6, 2: 5, 1: 5, 0: 8}282+ self.prefill_cut_num_max = {128: 330, 64: 120, 32: 70, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 5, 0: 8}
283+ self.prefill_cut_num_min = {128: 330, 64: 120, 32: 70, 16: 24, 8: 10, 4: 4, 3: 6, 2: 5, 1: 5, 0: 8}
284+ else:
285+ self.prefill_default_cut_map = {128: 330, 64: 120, 32: 70, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 6, 0: 8}
286+ self.prefill_cut_num_max = {128: 330, 64: 120, 32: 70, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 5, 0: 8}
287+ self.prefill_cut_num_min = {128: 330, 64: 120, 32: 70, 16: 24, 8: 10, 4: 5, 3: 6, 2: 5, 1: 5, 0: 8}
274 return288 return
275 if self.soc_name == 'Ascend910B4':289 if self.soc_name == 'Ascend910B4':
276 logger.info(f"[layerwiseDisaggregated] npu soc is Ascend910B4, ajust prefill cut num.")290 logger.info(f"[layerwiseDisaggregated] npu soc is Ascend910B4, ajust prefill cut num.")
277- self.prefill_default_cut_map = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 6, 0: 8}291+ if self.batch_p_num == 1:
278- self.prefill_cut_num_max = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 5, 0: 8}292+ self.prefill_default_cut_map = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 6, 0: 8}
279- self.prefill_cut_num_min = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 4, 3: 6, 2: 5, 1: 5, 0: 8}293+ self.prefill_cut_num_max = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 5, 0: 8}
294+ self.prefill_cut_num_min = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 4, 3: 6, 2: 5, 1: 5, 0: 8}
295+ else:
296+ self.prefill_default_cut_map = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 6, 0: 8}
297+ self.prefill_cut_num_max = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 6, 3: 8, 2: 9, 1: 5, 0: 8}
298+ self.prefill_cut_num_min = {128: 330, 64: 120, 32: 100, 16: 24, 8: 10, 4: 5, 3: 6, 2: 5, 1: 5, 0: 8}
280 return299 return
300+ 
301+ def __ajust_prefill_cut_num_for_multi_nodes(self):
302+ self.prefill_default_cut_map = {31.5: 80, 15.5: 59, 7.5: 45, 3.8: 19, 3.3: 20, 1.8: 17, 0.8: 21, 0: 21}
303+ self.prefill_cut_num_max = {31.5: 80, 15.5: 59, 7.5: 45, 3.8: 19, 3.3: 20, 1.8: 17, 0.8: 21, 0: 21}
304+ self.prefill_cut_num_min = {31.5: 80, 15.5: 59, 7.5: 45, 3.8: 19, 3.3: 20, 1.8: 17, 0.8: 21, 0: 21}
305+ logger.info(f"[layerwiseDisaggregated] cut policy init multi nodes success, model_type: {self.model_type} "
306+ f"role_type: {self.role_type} default_cut_map: {self.prefill_default_cut_map}")
@@ -294,14 +294,11 @@ class EdgeCloudCtrlComm:
294 self.decode_comm_finish = False294 self.decode_comm_finish = False
295 self.prefill_comm_finish = False295 self.prefill_comm_finish = False
296 self.prefill_comm_finish_tcp_count = 0296 self.prefill_comm_finish_tcp_count = 0
297- self.prefill_comm_finish_irecv = False
298 297 
299 self.prefill_recv_msg = ''298 self.prefill_recv_msg = ''
300 self.decode_recv_msg = ''299 self.decode_recv_msg = ''
301 self.prefill_send_msg = ''300 self.prefill_send_msg = ''
302 self.decode_send_msg = ''301 self.decode_send_msg = ''
303- self.parse_msg_cnt = 0
304- self.to_msg_cnt = 0
305 302
306 self.multi_nodes_infer_enabled = False303 self.multi_nodes_infer_enabled = False
307 self.multi_nodes_is_master = False304 self.multi_nodes_is_master = False
@@ -313,6 +310,13 @@ class EdgeCloudCtrlComm:
313 310 
314 self.tls_config = tls_config311 self.tls_config = tls_config
315 312 
313+ @staticmethod
314+ def shape_to_msg(shape):
315+ if shape is None or len(shape) != 2:
316+ return None
317+ msg = f"pull|{json.dumps(list(shape))}|0"
318+ return msg
319+ 
316 def init_role(self, role, server_ip, server_port):320 def init_role(self, role, server_ip, server_port):
317 self.role = role321 self.role = role
318 322 
@@ -454,18 +458,4 @@ class EdgeCloudCtrlComm:
454 decision = self.multi_nodes_ctrl_client.recv()458 decision = self.multi_nodes_ctrl_client.recv()
455 self.multi_nodes_ctrl_client.send("ok")459 self.multi_nodes_ctrl_client.send("ok")
456 logger.info(f"[layerwiseDisaggregated-{self.rank}] recv multi nodes decision {decision}.")460 logger.info(f"[layerwiseDisaggregated-{self.rank}] recv multi nodes decision {decision}.")
457- return decision461+ return decision
458- 
459- def parse_shape(self, data):
460- if not data.startswith("pull"):
461- return []
462- h_shape_d = list(json.loads(data.split('|')[1]))
463- self.parse_msg_cnt += 1
464- return h_shape_d
465- 
466- def shape_to_msg(self, shape):
467- if shape is None or len(shape) != 2:
468- return None
469- msg = f"pull|{json.dumps(list(shape))}|0"
470- self.to_msg_cnt += 1
471- return msg
@@ -83,8 +83,6 @@ class EdgeCloudDataComm:
83 self.lock = threading.Lock()83 self.lock = threading.Lock()
84 self.flag_pre_recv = True84 self.flag_pre_recv = True
85 85 
86- self.need_set_decode_device = False
87- self.need_set_prefill_device = False
88 self.set_decode_device_done = False86 self.set_decode_device_done = False
89 self.set_prefill_device_done = False87 self.set_prefill_device_done = False
90 88 
@@ -379,7 +377,7 @@ class EdgeCloudDataComm:
379 377
380 if self.rank == src_rank: # 对应的卡才进行收发378 if self.rank == src_rank: # 对应的卡才进行收发
381 if mode == 'p':379 if mode == 'p':
382- if self.role == CLOUD and not self.set_prefill_device_done and self.need_set_prefill_device:380+ if self.role == CLOUD and not self.set_prefill_device_done:
383 torch.npu.set_device(torch.device(f"npu:{self.rank}"))381 torch.npu.set_device(torch.device(f"npu:{self.rank}"))
384 self.set_prefill_device_done = True382 self.set_prefill_device_done = True
385 self.target_p[recv_index] = self.out_hidden_p[recv_index][:shape[recv_index], :]383 self.target_p[recv_index] = self.out_hidden_p[recv_index][:shape[recv_index], :]
@@ -391,7 +389,7 @@ class EdgeCloudDataComm:
391 self.ret_p[recv_index] = ret389 self.ret_p[recv_index] = ret
392 logger.info(f"[rank-{self.rank}] prefill start async recv, shape={shape[recv_index]}")390 logger.info(f"[rank-{self.rank}] prefill start async recv, shape={shape[recv_index]}")
393 else:391 else:
394- if self.role == CLOUD and not self.set_decode_device_done and self.need_set_decode_device:392+ if self.role == CLOUD and not self.set_decode_device_done:
395 torch.npu.set_device(torch.device(f"npu:{self.rank}"))393 torch.npu.set_device(torch.device(f"npu:{self.rank}"))
396 self.set_decode_device_done = True394 self.set_decode_device_done = True
397 self.target_d = self.out_hidden_d[:shape, :]395 self.target_d = self.out_hidden_d[:shape, :]
@@ -22,7 +22,7 @@ from tests.pythontest.atb_llm.models.base.mock_class import MockTorchClasses
22 22 
23class TestLayerwiseDecodeGraphWrapper(unittest.TestCase):23class TestLayerwiseDecodeGraphWrapper(unittest.TestCase):
24 def setUp(self):24 def setUp(self):
25- self.layerwise = LayerWiseAttr(start_num=1, end_num=1, split_type=DistributedType.CLOUD)25+ self.layerwise = LayerWiseAttr(edge_start_layer_count=1, edge_end_layer_count=1, split_type=DistributedType.CLOUD)
26 self.graph_wrapper = LayerwiseDecodeGraphWrapper(self.layerwise)26 self.graph_wrapper = LayerwiseDecodeGraphWrapper(self.layerwise)
27 self.model_type = "test_class"27 self.model_type = "test_class"
28 28 
@@ -38,7 +38,7 @@ class TestLayerwiseEdgeDecodeGraphWrapper(unittest.TestCase):
38 def setUp(self):38 def setUp(self):
39 self.mock_torch_classes = MockTorchClasses()39 self.mock_torch_classes = MockTorchClasses()
40 torch.classes = self.mock_torch_classes40 torch.classes = self.mock_torch_classes
41- self.layerwise = LayerWiseAttr(start_num=1, end_num=1, split_type=DistributedType.CLOUD)41+ self.layerwise = LayerWiseAttr(edge_start_layer_count=1, edge_end_layer_count=1, split_type=DistributedType.CLOUD)
42 self.graph_wrapper = LayerwiseEdgeDecodeGraphWrapper(self.layerwise)42 self.graph_wrapper = LayerwiseEdgeDecodeGraphWrapper(self.layerwise)
43 self.model_type = "test_class"43 self.model_type = "test_class"
44 44 
@@ -23,7 +23,7 @@ from tests.pythontest.atb_llm.models.base.mock_class import MockTorchClasses
23 23 
24class TestLayerwisePrefillGraphWrapper(unittest.TestCase):24class TestLayerwisePrefillGraphWrapper(unittest.TestCase):
25 def setUp(self):25 def setUp(self):
26- self.layerwise = LayerWiseAttr(start_num=1, end_num=1, split_type=DistributedType.CLOUD)26+ self.layerwise = LayerWiseAttr(edge_start_layer_count=1, edge_end_layer_count=1, split_type=DistributedType.CLOUD)
27 self.graph_wrapper = LayerwisePrefillGraphWrapper(self.layerwise)27 self.graph_wrapper = LayerwisePrefillGraphWrapper(self.layerwise)
28 self.model_type = "test_class"28 self.model_type = "test_class"
29 29 
@@ -38,7 +38,7 @@ class TestLayerwiseEdgePrefillGraphWrapper(unittest.TestCase):
38 def setUp(self):38 def setUp(self):
39 self.mock_torch_classes = MockTorchClasses()39 self.mock_torch_classes = MockTorchClasses()
40 torch.classes = self.mock_torch_classes40 torch.classes = self.mock_torch_classes
41- self.layerwise = LayerWiseAttr(start_num=1, end_num=1, split_type=DistributedType.CLOUD)41+ self.layerwise = LayerWiseAttr(edge_start_layer_count=1, edge_end_layer_count=1, split_type=DistributedType.CLOUD)
42 self.graph_wrapper = LayerwiseEdgePrefillGraphWrapper(self.layerwise)42 self.graph_wrapper = LayerwiseEdgePrefillGraphWrapper(self.layerwise)
43 self.model_type = "test_class"43 self.model_type = "test_class"
44 44 
@@ -145,7 +145,7 @@ class TestLayerwiseCloudPrefillGraphWrapper(unittest.TestCase):
145 def setUp(self):145 def setUp(self):
146 self.mock_torch_classes = MockTorchClasses()146 self.mock_torch_classes = MockTorchClasses()
147 torch.classes = self.mock_torch_classes147 torch.classes = self.mock_torch_classes
148- self.layerwise = LayerWiseAttr(start_num=1, end_num=1, split_type=DistributedType.CLOUD)148+ self.layerwise = LayerWiseAttr(edge_start_layer_count=1, edge_end_layer_count=1, split_type=DistributedType.CLOUD)
149 self.layerwise.num_hidden_layers = 4149 self.layerwise.num_hidden_layers = 4
150 self.config = Qwen2Config(150 self.config = Qwen2Config(
151 )151 )
@@ -16,7 +16,7 @@ from atb_llm.models.base.flash_causal_lm import FlashForCausalLM, LayerWiseAttr,
16 16 
17class TestLayerwiseModifier(unittest.TestCase):17class TestLayerwiseModifier(unittest.TestCase):
18 def setUp(self):18 def setUp(self):
19- self.layerwise = LayerWiseAttr(start_num=1, end_num=1, split_type=DistributedType.CLOUD)19+ self.layerwise = LayerWiseAttr(edge_start_layer_count=1, edge_end_layer_count=1, split_type=DistributedType.CLOUD)
20 self.layerwise_modifier = LayerwiseModifier(self.layerwise)20 self.layerwise_modifier = LayerwiseModifier(self.layerwise)
21 21 
22 22 
@@ -20,7 +20,7 @@ class TestCloudCutPolicy(unittest.TestCase):
20 mock_acl.get_soc_name = Mock()20 mock_acl.get_soc_name = Mock()
21 mock_acl.get_soc_name.return_value = 'Ascend910B4'21 mock_acl.get_soc_name.return_value = 'Ascend910B4'
22 self.cloud_cut_policy = CloudCutPolicy("slave")22 self.cloud_cut_policy = CloudCutPolicy("slave")
23- self.cloud_cut_policy.initialize("slave", 0, 62, 2)23+ self.cloud_cut_policy.initialize("slave", 0, 62, 2, False)
24 pass24 pass
25 25 
26 @classmethod26 @classmethod
@@ -73,7 +73,16 @@ class TestCloudCutPolicy(unittest.TestCase):
73 self.assertEqual(cut_num, 8)73 self.assertEqual(cut_num, 8)
74 74
75 def test_ajust_prefill_cut_num_for_diff_npu_soc(self):75 def test_ajust_prefill_cut_num_for_diff_npu_soc(self):
76+ self.cloud_cut_policy.soc_name = 'Ascend910B2'
77+ self.cloud_cut_policy.batch_p_num = 2
78+ self.cloud_cut_policy._CloudCutPolicy__ajust_prefill_cut_num_for_diff_npu_soc()
79+ self.assertEqual(self.cloud_cut_policy.prefill_default_cut_map.get(32), 100)
80+
76 self.cloud_cut_policy.soc_name = 'Ascend910B3'81 self.cloud_cut_policy.soc_name = 'Ascend910B3'
82+ self.cloud_cut_policy.batch_p_num = 1
83+ self.cloud_cut_policy._CloudCutPolicy__ajust_prefill_cut_num_for_diff_npu_soc()
84+ self.assertEqual(self.cloud_cut_policy.prefill_default_cut_map.get(32), 70)
85+ self.cloud_cut_policy.batch_p_num = 2
77 self.cloud_cut_policy._CloudCutPolicy__ajust_prefill_cut_num_for_diff_npu_soc()86 self.cloud_cut_policy._CloudCutPolicy__ajust_prefill_cut_num_for_diff_npu_soc()
78 self.assertEqual(self.cloud_cut_policy.prefill_default_cut_map.get(32), 70)87 self.assertEqual(self.cloud_cut_policy.prefill_default_cut_map.get(32), 70)
79 88
@@ -235,7 +235,6 @@ class TestEdgeCloudCtrlComm(unittest.TestCase):
235 EdgeCloudCtrlComm.decode_comm_finish = False235 EdgeCloudCtrlComm.decode_comm_finish = False
236 EdgeCloudCtrlComm.prefill_comm_finish = False236 EdgeCloudCtrlComm.prefill_comm_finish = False
237 EdgeCloudCtrlComm.prefill_comm_finish_tcp_count = 0237 EdgeCloudCtrlComm.prefill_comm_finish_tcp_count = 0
238- EdgeCloudCtrlComm.prefill_comm_finish_irecv = False
239 238 
240 EdgeCloudCtrlComm.prefill_recv_msg = ''239 EdgeCloudCtrlComm.prefill_recv_msg = ''
241 EdgeCloudCtrlComm.decode_recv_msg = ''240 EdgeCloudCtrlComm.decode_recv_msg = ''
@@ -338,11 +337,6 @@ class TestEdgeCloudCtrlComm(unittest.TestCase):
338 result = comm.is_edge_cloud_ctrl_comm_success()337 result = comm.is_edge_cloud_ctrl_comm_success()
339 self.assertFalse(result)338 self.assertFalse(result)
340 339 
341- def test_parse_shape(self):
342- comm = EdgeCloudCtrlComm({})
343- self.assertEqual(comm.parse_shape(" "), [])
344- self.assertEqual(comm.parse_shape("pull|[1,2,3,4]"), [1, 2, 3, 4])
345- 
346 def test_shape_to_msg(self):340 def test_shape_to_msg(self):
347 comm = EdgeCloudCtrlComm({})341 comm = EdgeCloudCtrlComm({})
348 self.assertIsNone(comm.shape_to_msg([]))342 self.assertIsNone(comm.shape_to_msg([]))