已合并
[Feature]:使能LayerwiseDisaggregated边云协同推理特性前提下, 支持配置使能多机推理特性和2p下发的计算图适配 #372
K先生创建于 2月4日
[Feature]:使能LayerwiseDisaggregated边云协同推理特性前提下, 支持配置使能多机推理特性和2p下发的计算图适配 #372
已合并
从已删除 :dev合入到Ascend/MindIE-LLMdev
共 20 个文件变更+349-213
| @@ -296,13 +296,13 @@ std::vector<std::string> ConstructIntertensorList(const DecoderLayerParam ¶m | |||
| 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 moe | 1877 | // 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 moe | 1883 | // 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 ¶m, 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 ¶ | |||
| 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->layerwiseDisaggregated | 334 | 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 | + | ||
| 622 | atb::Status DecoderModel::InferShape( | 676 | atb::Status DecoderModel::InferShape( |
| 623 | const std::vector<atb::TensorDesc> &inTensorDescs, | 677 | const std::vector<atb::TensorDesc> &inTensorDescs, |
| 624 | std::vector<atb::TensorDesc> &outTensorDescs | 678 | 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 ¶ | |||
| 930 | 975 | ||
| 931 | void SetLayerwiseDisaggregatedParam(DecoderLayerParam &layerParam, const DeepseekV2ModelParam ¶m, int64_t layerId) | 976 | void SetLayerwiseDisaggregatedParam(DecoderLayerParam &layerParam, const DeepseekV2ModelParam ¶m, 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 | ||
| 955 | void DecoderModel::SetLayerParam(DecoderLayerParam &layerParam, int64_t layerId) | 1001 | void 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 | }; |
| 134 | REGISTER_MODEL(deepseekV2, DecoderModel); | 135 | REGISTER_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 | ||
| 47 | class LayerWiseAttr: | 47 | class 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_num | 55 | + self.edge_start_layer_count = edge_start_layer_count |
| 56 | - self.end_num = end_num | 56 | + self.edge_end_layer_count = edge_end_layer_count |
| 57 | self.split_type = split_type | 57 | 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 = None | 84 | split_type = None |
| 85 | - start_num = 1 | 85 | + edge_start_layer_count = 1 |
| 86 | - end_num = 1 | 86 | + 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 = True | 88 | self.layerwise_disaggregated = True |
| 89 | if layerwise_disaggregated_role_type == "slave": | 89 | if layerwise_disaggregated_role_type == "slave": |
| 90 | split_type = DistributedType.CLOUD | 90 | split_type = DistributedType.CLOUD |
| 91 | else: | 91 | else: |
| 92 | split_type = DistributedType.EDGE | 92 | 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 logits | 622 | 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 = None | 94 | 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 | ||
| 10 | from typing import List | 11 | from typing import List |
| 11 | from atb_llm.models.base.flash_causal_lm import LayerWiseAttr, LwdLayerStatus, DistributedType | 12 | from 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 = attr | 22 | 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 = None | 29 | self.acl_cloud_inputs = None |
| 25 | self.acl_cloud_params = None | 30 | 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 | 40 | ||
| 36 | def to_index(is_prefill): | 41 | def to_index(is_prefill): |
| 37 | return 1 if is_prefill else 0 | 42 | 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. | ||
| 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 | return | 87 | 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_hidden | 94 | + last_input[0] = out_hidden |
| 68 | - inputs[:] = self.acl_edge_inputs_prefill_pre | 95 | + 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_hidden | 100 | inputs[0] = out_hidden |
| @@ -13,6 +13,7 @@ import os | |||
| 13 | import json | 13 | import json |
| 14 | import math | 14 | import math |
| 15 | from enum import Enum | 15 | from enum import Enum |
| 16 | +import queue | ||
| 16 | from typing import List, Optional, Tuple | 17 | from typing import List, Optional, Tuple |
| 17 | from dataclasses import asdict | 18 | from dataclasses import asdict |
| 18 | 19 | ||
| @@ -188,12 +189,12 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM): | |||
| 188 | self.prefix_cache_enable = True | 189 | 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_num | 192 | + start_layer = self.layerwise.edge_start_layer_count |
| 192 | - end_layer = self.config.num_hidden_layers - self.layerwise.end_num | 193 | + 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_num | 197 | + 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 = None | 209 | self.layerwise.acl_inputs_prefill = None |
| 209 | self.layerwise.acl_inputs_decode = None | 210 | self.layerwise.acl_inputs_decode = None |
| 211 | + self.layerwise.acl_inputs_prefill_queue = queue.Queue() | ||
| 210 | self.layerwise.acl_param_prefill = None | 212 | self.layerwise.acl_param_prefill = None |
| 211 | self.layerwise.acl_param_decode = None | 213 | self.layerwise.acl_param_decode = None |
| 214 | + self.layerwise.acl_param_prefill_queue = queue.Queue() | ||
| 212 | self.layerwise.p_out_hidden = None | 215 | self.layerwise.p_out_hidden = None |
| 213 | self.layerwise.acl_inputs_prefill_pre = None | 216 | self.layerwise.acl_inputs_prefill_pre = None |
| 214 | self.layerwise.acl_param_prefill_pre = None | 217 | 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 = False | 554 | 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_dp | 562 | self.enable_qkvdown_dp = config.models.deepseekv2.h3p.enable_qkvdown_dp |
| 554 | self.enable_gating_dp = config.models.deepseekv2.h3p.enable_gating_dp | 563 | self.enable_gating_dp = config.models.deepseekv2.h3p.enable_gating_dp |
| 555 | self.enable_shared_expert_dp = config.models.deepseekv2.h3p.enable_shared_expert_dp | 564 | 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.placeholder | 642 | self.attn_padding_idx = self.placeholder |
| @@ -777,16 +788,16 @@ class FlashDeepseekv2ForCausalLM(FlashForCausalLM): | |||
| 777 | modify_ascend_params["layerwiseMode"] = mode | 788 | modify_ascend_params["layerwiseMode"] = mode |
| 778 | if mode == 0: | 789 | if mode == 0: |
| 779 | modify_ascend_params[START_ID] = 0 | 790 | modify_ascend_params[START_ID] = 0 |
| 780 | - modify_ascend_params[END_ID] = self.layerwise.start_num | 791 | + 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 0 | 794 | 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_num | 796 | + 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_num | 797 | + 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) + 1 | 800 | + 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_type | 806 | 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_num | 808 | + 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_layers | 809 | 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) + 1 | 812 | 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 = 0 | 824 | end_layer = 0 |
| 813 | if mode == 0: | 825 | if mode == 0: |
| 814 | start_layer = 0 | 826 | start_layer = 0 |
| 815 | - end_layer = self.layerwise.start_num | 827 | + 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) + 1 | 831 | 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) + 1 | 835 | + 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) + 1 | 839 | 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 + layer | 1179 | + encoder_internal_param[START_ID] = self.layerwise.edge_start_layer_count + layer |
| 1166 | - encoder_internal_param[END_ID] = self.layerwise.start_num + layer + 1 | 1180 | + 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"] = 1 | 1183 | 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_wrappers | 1196 | 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.placeholder | 1654 | 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_mtp | 2202 | 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_hidden | 2270 | 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_hidden | 2273 | + 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 = True | 2281 | 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_hidden | 2284 | + 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 logits | 2286 | return logits |
| 2250 | 2287 | ||
| 2251 | return out_hidden | 2288 | 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_hidden | 2340 | 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 = True | 2344 | 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_hidden | 2347 | 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] = None | 2349 | + self.layerwise.acl_inputs_prefill = [None] + acl_inputs[1:] |
| 2312 | self.layerwise.acl_param_prefill = acl_param | 2350 | 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_hidden | 2397 | 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) + 1 | 2846 | 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) + 1 | 2854 | + 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 = True | 75 | 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_num | 78 | + start_layer = self.layerwise.edge_start_layer_count |
| 79 | - end_layer = self.config.num_hidden_layers - self.layerwise.end_num | 79 | + 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_num | 83 | + 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"] = mode | 308 | modify_ascend_params["layerwiseMode"] = mode |
| 308 | modify_ascend_params["linearDescs"] = wrapper.linear_descs | 309 | 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_num | 311 | + modify_ascend_params[LINEAR_HAS_BIAS] = linear_has_bias * self.layerwise.edge_start_layer_count |
| 311 | modify_ascend_params[LWD_START_ID] = 0 | 312 | modify_ascend_params[LWD_START_ID] = 0 |
| 312 | - modify_ascend_params[LWD_END_ID] = self.layerwise.start_num | 313 | + 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_num | 317 | + self.layerwise.edge_end_layer_count) |
| 317 | - modify_ascend_params[LWD_END_ID] = self.config.num_hidden_layers - self.layerwise.end_num | 318 | + 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_num | 321 | + 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_num | 322 | + 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_layers | 323 | 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 = 0 | 354 | end_layer = 0 |
| 353 | if mode == LwdLayerStatus.EDGE_START_LAYER: | 355 | if mode == LwdLayerStatus.EDGE_START_LAYER: |
| 354 | start_layer = 0 | 356 | start_layer = 0 |
| 355 | - end_layer = self.layerwise.start_num | 357 | + 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) + 1 | 361 | 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) + 1 | 365 | + 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) + 1 | 369 | 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 + layer | 576 | + 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 + 1 | 577 | + encoder_internal_param[LWD_END_ID] = self.layerwise.edge_start_layer_count + layer + 1 |
| 574 | encoder_internal_param["numHiddenLayers"] = 1 | 578 | encoder_internal_param["numHiddenLayers"] = 1 |
| 575 | encoder_internal_param[LINEAR_HAS_BIAS] = linear_has_bias | 579 | encoder_internal_param[LINEAR_HAS_BIAS] = linear_has_bias |
| 576 | encoder_internal_param["reuseEmbedTable"] = self.long_seq_enable and layer != 0 | 580 | 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) + 1 | 775 | 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) + 1 | 792 | + 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 = None | 281 | self.prefill_input_lengths = None |
| 281 | self.edge_pre_chunk_length = 0 | 282 | self.edge_pre_chunk_length = 0 |
| 282 | self.prefill_total_seq_len = 0 | 283 | 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 hidden | 610 | 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 res | 626 | 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_idx | 781 | 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._instance | 23 | 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_map | 30 | # 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 | return | 56 | 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 | return | 62 | 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 | return | 68 | 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._instance | 48 | 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 = name | 52 | 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.OTHER | 54 | self.role_type = CloudCutClassType.CLOUD if name == "slave" else CloudCutClassType.OTHER |
| 55 | self.rank_id = None | 55 | 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 = False | 59 | 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_SEEK | 117 | return CloudCutModelType.DEEP_SEEK |
| 116 | return CloudCutModelType.QWEN | 118 | 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.OTHER | 121 | self.role_type = CloudCutClassType.CLOUD if name == "slave" else CloudCutClassType.OTHER |
| 120 | self.rank_id = rank_id | 122 | self.rank_id = rank_id |
| 121 | self.max_cut_num = max_cut_num | 123 | self.max_cut_num = max_cut_num |
| 122 | self.min_cut_num = min_cut_num | 124 | 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 = True | 128 | 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 | return | 277 | 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 | return | 288 | 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 | return | 299 | 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 = False | 294 | self.decode_comm_finish = False |
| 295 | self.prefill_comm_finish = False | 295 | self.prefill_comm_finish = False |
| 296 | self.prefill_comm_finish_tcp_count = 0 | 296 | 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 = False | 303 | self.multi_nodes_infer_enabled = False |
| 307 | self.multi_nodes_is_master = False | 304 | self.multi_nodes_is_master = False |
| @@ -313,6 +310,13 @@ class EdgeCloudCtrlComm: | |||
| 313 | 310 | ||
| 314 | self.tls_config = tls_config | 311 | self.tls_config = tls_config |
| 315 | 312 | ||
| 313 | + | ||
| 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 = role | 321 | 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 decision | 461 | + 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 = True | 84 | 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 = False | 86 | self.set_decode_device_done = False |
| 89 | self.set_prefill_device_done = False | 87 | 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 = True | 382 | 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] = ret | 389 | 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 = True | 394 | 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 | ||
| 23 | class TestLayerwiseDecodeGraphWrapper(unittest.TestCase): | 23 | class 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_classes | 40 | 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 | ||
| 24 | class TestLayerwisePrefillGraphWrapper(unittest.TestCase): | 24 | class 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_classes | 40 | 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_classes | 147 | 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 = 4 | 149 | self.layerwise.num_hidden_layers = 4 |
| 150 | self.config = Qwen2Config( | 150 | self.config = Qwen2Config( |
| 151 | ) | 151 | ) |
Mexamples/atb_models/tests/pythontest/atb_llm/models/base/inputs_modifier/test_layerwise_modifier.py+1-1
| @@ -16,7 +16,7 @@ from atb_llm.models.base.flash_causal_lm import FlashForCausalLM, LayerWiseAttr, | |||
| 16 | 16 | ||
| 17 | class TestLayerwiseModifier(unittest.TestCase): | 17 | class 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 | ||
Mexamples/atb_models/tests/pythontest/atb_llm/utils/layerwise_disaggregated/test_cloud_cut_policy.py+10-1
| @@ -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 | pass | 24 | pass |
| 25 | 25 | ||
| 26 | 26 | ||
| @@ -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 = False | 235 | EdgeCloudCtrlComm.decode_comm_finish = False |
| 236 | EdgeCloudCtrlComm.prefill_comm_finish = False | 236 | EdgeCloudCtrlComm.prefill_comm_finish = False |
| 237 | EdgeCloudCtrlComm.prefill_comm_finish_tcp_count = 0 | 237 | 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([])) |
1)加注释,说明为什么inputs的第0个需要替换成None 2)说明为什么需要copy