已合并
【新需求】GLM4.1V适配lcoc, flash comm和权重预取 #251
凡银创建于 1月15日
【新需求】GLM4.1V适配lcoc, flash comm和权重预取 #251
已合并
从已删除 :glm4v合入到Ascend/MindIE-LLMdev
共 8 个文件变更+111-90
| @@ -16,11 +16,9 @@ namespace glm41v { | |||
| 16 | 16 | ||
| 17 | static const uint64_t NUM2 = 2; | 17 | static const uint64_t NUM2 = 2; |
| 18 | 18 | ||
| 19 | -Glm41vDecoderLayer::Glm41vDecoderLayer( | 19 | +DecoderLayer::DecoderLayer( |
| 20 | - const Glm41vLayerParam ¶m) : atb_speed::base::DecoderLayer<atb::infer::RmsNormParam>(param) | 20 | + const atb_speed::base::LayerParam ¶m) : atb_speed::base::DecoderLayer<atb::infer::RmsNormParam>(param) |
| 21 | { | 21 | { |
| 22 | - this->param = param; | ||
| 23 | - this->param.CheckParam(); | ||
| 24 | this->inTensorCandidates["post_self_attn_norm_weight"] = { | 22 | this->inTensorCandidates["post_self_attn_norm_weight"] = { |
| 25 | "in_post_self_attn_norm_weight" | 23 | "in_post_self_attn_norm_weight" |
| 26 | }; | 24 | }; |
| @@ -29,7 +27,7 @@ Glm41vDecoderLayer::Glm41vDecoderLayer( | |||
| 29 | }; | 27 | }; |
| 30 | }; | 28 | }; |
| 31 | 29 | ||
| 32 | -void Glm41vDecoderLayer::ConstructInTensorMap() | 30 | +void DecoderLayer::ConstructInTensorMap() |
| 33 | { | 31 | { |
| 34 | this->inTensorList.clear(); | 32 | this->inTensorList.clear(); |
| 35 | atb_speed::common::AddTensorToList(this->inTensorCandidates, "input_norm_weight", this->inTensorList); | 33 | atb_speed::common::AddTensorToList(this->inTensorCandidates, "input_norm_weight", this->inTensorList); |
| @@ -42,27 +40,20 @@ void Glm41vDecoderLayer::ConstructInTensorMap() | |||
| 42 | if (param.enableSpeculate) { | 40 | if (param.enableSpeculate) { |
| 43 | atb_speed::common::AddTensorToList(this->inTensorCandidates, "q_len", this->inTensorList); | 41 | atb_speed::common::AddTensorToList(this->inTensorCandidates, "q_len", this->inTensorList); |
| 44 | } | 42 | } |
| 45 | - atb_speed::common::AddTensorToList( | 43 | + if (param.enableFlashComm) { |
| 46 | - this->internalTensorCandidates, "default", this->intermediateTensorList); | 44 | + atb_speed::common::AddTensorToList(this->inTensorCandidates, "flash_comm", this->inTensorList); |
| 47 | - this->graph.inTensorNum = this->inTensorList.size(); | 45 | + } |
| 48 | } | 46 | } |
| 49 | 47 | ||
| 50 | -void Glm41vDecoderLayer::SetFusionAttentionParam( | 48 | +void DecoderLayer::SetFusionAttentionParam( |
| 51 | atb_speed::common::FusionAttentionParam<atb::infer::RmsNormParam> &fusionAttentionParam) | 49 | atb_speed::common::FusionAttentionParam<atb::infer::RmsNormParam> &fusionAttentionParam) |
| 52 | { | 50 | { |
| 53 | - DecoderLayer<atb::infer::RmsNormParam>::SetFusionAttentionParam(fusionAttentionParam); | 51 | + atb_speed::base::DecoderLayer<atb::infer::RmsNormParam>::SetFusionAttentionParam(fusionAttentionParam); |
| 54 | fusionAttentionParam.rotaryType = atb_speed::common::RotaryType::HALF_ROTARY; | 52 | fusionAttentionParam.rotaryType = atb_speed::common::RotaryType::HALF_ROTARY; |
| 55 | fusionAttentionParam.ropeParam.rotaryCoeff = this->param.hiddenSizePerAttentionHead / NUM2; | 53 | fusionAttentionParam.ropeParam.rotaryCoeff = this->param.hiddenSizePerAttentionHead / NUM2; |
| 56 | } | 54 | } |
| 57 | 55 | ||
| 58 | -void Glm41vDecoderLayer::SetFusionAttentionNormParam( | 56 | +atb::Status DecoderLayer::AddPostSelfAttentionRMSNorm() |
| 59 | - atb_speed::common::FusionAttentionParam<atb::infer::RmsNormParam> &fusionAttentionParam) | ||
| 60 | -{ | ||
| 61 | - DecoderLayer<atb::infer::RmsNormParam>::SetFusionAttentionNormParam(fusionAttentionParam); | ||
| 62 | - fusionAttentionParam.enableNormQuantOp = false; | ||
| 63 | -} | ||
| 64 | - | ||
| 65 | -atb::Status Glm41vDecoderLayer::AddPostSelfAttentionRMSNorm() | ||
| 66 | { | 57 | { |
| 67 | atb::infer::RmsNormParam normParam; | 58 | atb::infer::RmsNormParam normParam; |
| 68 | normParam.layerType = atb::infer::RmsNormParam::RmsNormType::RMS_NORM_NORM; | 59 | normParam.layerType = atb::infer::RmsNormParam::RmsNormType::RMS_NORM_NORM; |
| @@ -78,7 +69,7 @@ atb::Status Glm41vDecoderLayer::AddPostSelfAttentionRMSNorm() | |||
| 78 | return atb::NO_ERROR; | 69 | return atb::NO_ERROR; |
| 79 | } | 70 | } |
| 80 | 71 | ||
| 81 | -atb::Status Glm41vDecoderLayer::AddPostMlpRMSNorm() | 72 | +atb::Status DecoderLayer::AddPostMlpRMSNorm() |
| 82 | { | 73 | { |
| 83 | atb::infer::RmsNormParam normParam; | 74 | atb::infer::RmsNormParam normParam; |
| 84 | normParam.layerType = atb::infer::RmsNormParam::RmsNormType::RMS_NORM_NORM; | 75 | normParam.layerType = atb::infer::RmsNormParam::RmsNormType::RMS_NORM_NORM; |
| @@ -94,21 +85,14 @@ atb::Status Glm41vDecoderLayer::AddPostMlpRMSNorm() | |||
| 94 | return atb::NO_ERROR; | 85 | return atb::NO_ERROR; |
| 95 | } | 86 | } |
| 96 | 87 | ||
| 97 | -atb::Status Glm41vDecoderLayer::AddOperationToGraph() | 88 | +atb::Status DecoderLayer::AddOperationToGraph() |
| 98 | { | 89 | { |
| 99 | CHECK_OPERATION_STATUS_RETURN(this->AddFusionAttention()); | 90 | CHECK_OPERATION_STATUS_RETURN(this->AddFusionAttention()); |
| 100 | CHECK_OPERATION_STATUS_RETURN(this->AddPostSelfAttentionRMSNorm()); | 91 | CHECK_OPERATION_STATUS_RETURN(this->AddPostSelfAttentionRMSNorm()); |
| 101 | CHECK_OPERATION_STATUS_RETURN(this->AddFusionAttentionResidualAdd()); | 92 | CHECK_OPERATION_STATUS_RETURN(this->AddFusionAttentionResidualAdd()); |
| 102 | - if (param.hasAttnDp && param.hasMlpTp) { | ||
| 103 | - CHECK_OPERATION_STATUS_RETURN(this->AddFusedAllGather()); | ||
| 104 | - } | ||
| 105 | CHECK_OPERATION_STATUS_RETURN(this->AddMlp()); | 93 | CHECK_OPERATION_STATUS_RETURN(this->AddMlp()); |
| 106 | CHECK_OPERATION_STATUS_RETURN(this->AddPostMlpRMSNorm()); | 94 | CHECK_OPERATION_STATUS_RETURN(this->AddPostMlpRMSNorm()); |
| 107 | CHECK_OPERATION_STATUS_RETURN(this->AddMlpResidualAdd()); | 95 | CHECK_OPERATION_STATUS_RETURN(this->AddMlpResidualAdd()); |
| 108 | - if (param.hasAttnDp && param.hasMlpTp) { | ||
| 109 | - CHECK_OPERATION_STATUS_RETURN(this->AddRevertAllGather()); | ||
| 110 | - ATB_SPEED_LOG_DEBUG("Revert AllGather finished"); | ||
| 111 | - } | ||
| 112 | return atb::NO_ERROR; | 96 | return atb::NO_ERROR; |
| 113 | } | 97 | } |
| 114 | 98 | ||
| @@ -20,22 +20,18 @@ | |||
| 20 | namespace atb_speed { | 20 | namespace atb_speed { |
| 21 | namespace glm41v { | 21 | namespace glm41v { |
| 22 | 22 | ||
| 23 | -class Glm41vLayerParam : public atb_speed::base::LayerParam { | 23 | + |
| 24 | -}; | 24 | +class DecoderLayer : public atb_speed::base::DecoderLayer<atb::infer::RmsNormParam> { |
| 25 | -class Glm41vDecoderLayer : public atb_speed::base::DecoderLayer<atb::infer::RmsNormParam> { | ||
| 26 | public: | 25 | public: |
| 27 | - explicit Glm41vDecoderLayer(const Glm41vLayerParam ¶m); | 26 | + explicit DecoderLayer(const atb_speed::base::LayerParam ¶m); |
| 28 | - ~Glm41vDecoderLayer() override {}; | 27 | + ~DecoderLayer() override {}; |
| 29 | protected: | 28 | protected: |
| 30 | void ConstructInTensorMap() override; | 29 | void ConstructInTensorMap() override; |
| 31 | atb::Status AddOperationToGraph() override; | 30 | atb::Status AddOperationToGraph() override; |
| 32 | void SetFusionAttentionParam( | 31 | void SetFusionAttentionParam( |
| 33 | atb_speed::common::FusionAttentionParam<atb::infer::RmsNormParam> &fusionAttentionParam) override; | 32 | atb_speed::common::FusionAttentionParam<atb::infer::RmsNormParam> &fusionAttentionParam) override; |
| 34 | - void SetFusionAttentionNormParam( | ||
| 35 | - atb_speed::common::FusionAttentionParam<atb::infer::RmsNormParam> &fusionAttentionParam) override; | ||
| 36 | atb::Status AddPostSelfAttentionRMSNorm(); | 33 | atb::Status AddPostSelfAttentionRMSNorm(); |
| 37 | atb::Status AddPostMlpRMSNorm(); | 34 | atb::Status AddPostMlpRMSNorm(); |
| 38 | - Glm41vLayerParam param; | ||
| 39 | }; | 35 | }; |
| 40 | 36 | ||
| 41 | 37 | ||
| @@ -17,17 +17,16 @@ namespace glm41v { | |||
| 17 | // Weight count | 17 | // Weight count |
| 18 | const uint64_t GLM4_WEIGHT_COUNT_PER_LAYER = 52; | 18 | const uint64_t GLM4_WEIGHT_COUNT_PER_LAYER = 52; |
| 19 | 19 | ||
| 20 | -Glm41vDecoderModel::Glm41vDecoderModel(const std::string ¶m) : atb_speed::base::DecoderModel(param) | 20 | +DecoderModel::DecoderModel(const std::string ¶m) : atb_speed::base::DecoderModel(param) |
| 21 | { | 21 | { |
| 22 | - this->param.FromString(param); | ||
| 23 | this->weightCountPerLayer = GLM4_WEIGHT_COUNT_PER_LAYER; | 22 | this->weightCountPerLayer = GLM4_WEIGHT_COUNT_PER_LAYER; |
| 24 | } | 23 | } |
| 25 | 24 | ||
| 26 | -atb::Status Glm41vDecoderModel::CreateLayerOperation(atb::Operation **op, uint32_t layerId) | 25 | +atb::Status DecoderModel::CreateLayerOperation(atb::Operation **op, uint32_t layerId) |
| 27 | { | 26 | { |
| 28 | - Glm41vLayerParam layerParam; | 27 | + atb_speed::base::LayerParam layerParam; |
| 29 | this->SetLayerParam(layerParam, layerId); | 28 | this->SetLayerParam(layerParam, layerId); |
| 30 | - Glm41vDecoderLayer decoderLayer(layerParam); | 29 | + atb_speed::glm41v::DecoderLayer decoderLayer(layerParam); |
| 31 | CHECK_OPERATION_STATUS_RETURN(decoderLayer.BuildGraph(op)); | 30 | CHECK_OPERATION_STATUS_RETURN(decoderLayer.BuildGraph(op)); |
| 32 | return atb::NO_ERROR; | 31 | return atb::NO_ERROR; |
| 33 | } | 32 | } |
| @@ -22,18 +22,15 @@ | |||
| 22 | namespace atb_speed { | 22 | namespace atb_speed { |
| 23 | namespace glm41v { | 23 | namespace glm41v { |
| 24 | 24 | ||
| 25 | -class Glm41vModelParam : public atb_speed::base::ModelParam { | 25 | +class DecoderModel : public atb_speed::base::DecoderModel { |
| 26 | -}; | ||
| 27 | -class Glm41vDecoderModel : public atb_speed::base::DecoderModel { | ||
| 28 | public: | 26 | public: |
| 29 | - explicit Glm41vDecoderModel(const std::string ¶m); | 27 | + explicit DecoderModel(const std::string ¶m); |
| 30 | protected: | 28 | protected: |
| 31 | atb::Status CreateLayerOperation(atb::Operation **op, uint32_t layerId) override; | 29 | atb::Status CreateLayerOperation(atb::Operation **op, uint32_t layerId) override; |
| 32 | - Glm41vModelParam param; | ||
| 33 | }; | 30 | }; |
| 34 | 31 | ||
| 35 | 32 | ||
| 36 | -REGISTER_MODEL(glm41v, Glm41vDecoderModel); | 33 | +REGISTER_MODEL(glm41v, DecoderModel); |
| 37 | } // namespace glm41v | 34 | } // namespace glm41v |
| 38 | } // namespace atb_speed | 35 | } // namespace atb_speed |
| 39 | 36 | ||
| @@ -27,9 +27,10 @@ from torch import nn | |||
| 27 | import torch_npu | 27 | import torch_npu |
| 28 | from atb_llm.models.base.flash_causal_lm import FlashForCausalLM | 28 | from atb_llm.models.base.flash_causal_lm import FlashForCausalLM |
| 29 | from atb_llm.models.base.modeling import FlashAttention, FlashLayer, MLP | 29 | from atb_llm.models.base.modeling import FlashAttention, FlashLayer, MLP |
| 30 | -from atb_llm.models.base.graph_manager.graph_manager import ATBGraphManager | ||
| 31 | from atb_llm.models.base.inputs_modifier.qlen_modifier import QLenModifier | 30 | from atb_llm.models.base.inputs_modifier.qlen_modifier import QLenModifier |
| 32 | -from atb_llm.models.base.graph_manager import DapGraphWrapper, SpeculateGraphWrapper | 31 | +from atb_llm.models.base.inputs_modifier.flash_comm_modifier import FlashCommModifier |
| 32 | +from atb_llm.models.base.graph_manager import ATBGraphManager, DapGraphWrapper, SpeculateGraphWrapper, \ | ||
| 33 | + FlashCommGraphWrapper | ||
| 33 | from atb_llm.utils.initial import NPUSocInfo | 34 | from atb_llm.utils.initial import NPUSocInfo |
| 34 | from atb_llm.utils.layers import TensorParallelRowLinear, RMSNorm, TensorEmbedding, TensorHead, \ | 35 | from atb_llm.utils.layers import TensorParallelRowLinear, RMSNorm, TensorEmbedding, TensorHead, \ |
| 35 | load_column_multi, PositionRotaryEmbedding, AttentionMask | 36 | load_column_multi, PositionRotaryEmbedding, AttentionMask |
| @@ -41,6 +42,12 @@ from atb_llm.utils.log.error_code import ErrorCode | |||
| 41 | from atb_llm.utils.log import logger | 42 | from atb_llm.utils.log import logger |
| 42 | 43 | ||
| 43 | 44 | ||
| 45 | +_800_9000_SOCS = (100, 101, 102, 103, 104) | ||
| 46 | +DUO_SOCS = (200, 201, 202, 203, 204, 205) | ||
| 47 | +A2_SOCS = (220, 221, 222, 223, 224, 225) | ||
| 48 | +A3_SOCS = (250, 251, 252, 253, 254, 255) | ||
| 49 | + | ||
| 50 | + | ||
| 44 | class Glm41vTextAttention(FlashAttention): | 51 | class Glm41vTextAttention(FlashAttention): |
| 45 | def __init__( | 52 | def __init__( |
| 46 | self, | 53 | self, |
| @@ -119,6 +126,7 @@ class FlashGlm41vTextModelForCausalLM(FlashForCausalLM): | |||
| 119 | else: | 126 | else: |
| 120 | prefix = "model.language_model" | 127 | prefix = "model.language_model" |
| 121 | self.enable_rope_quant_kvcache = self.config.quantization_config.kv_quant_type is not None | 128 | self.enable_rope_quant_kvcache = self.config.quantization_config.kv_quant_type is not None |
| 129 | + self.hidden_size = config.hidden_size | ||
| 122 | self.multi_query_group_num = self.config.num_key_value_heads | 130 | self.multi_query_group_num = self.config.num_key_value_heads |
| 123 | 131 | ||
| 124 | self.embed_tokens = TensorEmbedding(prefix=f"{prefix}.embed_tokens", weights=weights) | 132 | self.embed_tokens = TensorEmbedding(prefix=f"{prefix}.embed_tokens", weights=weights) |
| @@ -169,6 +177,7 @@ class FlashGlm41vTextModelForCausalLM(FlashForCausalLM): | |||
| 169 | # Multi graph management | 177 | # Multi graph management |
| 170 | self.graph_manager = ATBGraphManager() | 178 | self.graph_manager = ATBGraphManager() |
| 171 | self.qlen_decorator = QLenModifier() | 179 | self.qlen_decorator = QLenModifier() |
| 180 | + self.flash_comm_modifier = FlashCommModifier(weights, self.hidden_size, self._flash_comm_gate()) | ||
| 172 | 181 | ||
| 173 | def init_ascend_operations(self, config): | 182 | def init_ascend_operations(self, config): |
| 174 | pass | 183 | pass |
| @@ -228,7 +237,7 @@ class FlashGlm41vTextModelForCausalLM(FlashForCausalLM): | |||
| 228 | "isUnpadInputs": True, | 237 | "isUnpadInputs": True, |
| 229 | "skipWordEmbedding": True, | 238 | "skipWordEmbedding": True, |
| 230 | "isLmHeadParallel": True, | 239 | "isLmHeadParallel": True, |
| 231 | - "enableSwiGLU": True if self.soc_info.soc_version != 240 else False, | 240 | + "enableSwiGLU": True, |
| 232 | "rank": self.tp_rank, | 241 | "rank": self.tp_rank, |
| 233 | "worldSize": self.tp_world_size, | 242 | "worldSize": self.tp_world_size, |
| 234 | "backend": self.soc_info.communication_backend, | 243 | "backend": self.soc_info.communication_backend, |
| @@ -238,20 +247,23 @@ class FlashGlm41vTextModelForCausalLM(FlashForCausalLM): | |||
| 238 | encoder_param = { | 247 | encoder_param = { |
| 239 | **coder_param, | 248 | **coder_param, |
| 240 | "isPrefill": True, | 249 | "isPrefill": True, |
| 241 | - "supportLcoc": False if self.soc_info.need_nz else True | 250 | + "enablePreFetchWeight": self.soc_info.soc_version in DUO_SOCS, # Negative performance gains in A2 |
| 251 | + "enableLcoc": self.lcoc_enable, | ||
| 242 | } | 252 | } |
| 243 | decoder_param = { | 253 | decoder_param = { |
| 244 | **coder_param, | 254 | **coder_param, |
| 245 | "isPrefill": False, | 255 | "isPrefill": False, |
| 246 | - "supportLcoc": False | 256 | + "enableLcoc": False |
| 247 | } | 257 | } |
| 248 | if self.speculate_enable: | 258 | if self.speculate_enable: |
| 249 | self.graph_manager.register_graph(SpeculateGraphWrapper()) | 259 | self.graph_manager.register_graph(SpeculateGraphWrapper()) |
| 250 | if self.enable_dap: | 260 | if self.enable_dap: |
| 251 | self.graph_manager.register_graph(DapGraphWrapper()) | 261 | self.graph_manager.register_graph(DapGraphWrapper()) |
| 262 | + if self.flash_comm_modifier.enable_flash_comm: | ||
| 263 | + self.graph_manager.register_graph(FlashCommGraphWrapper()) | ||
| 252 | 264 | ||
| 253 | specified_params = {"decode": decoder_param} | 265 | specified_params = {"decode": decoder_param} |
| 254 | - self.graph_manager.set_param("glm41v_Glm41vDecoderModel", encoder_param, specified_params) | 266 | + self.graph_manager.set_param("glm41v_DecoderModel", encoder_param, specified_params) |
| 255 | self.graph_manager.set_weight(self.ascend_weight) | 267 | self.graph_manager.set_weight(self.ascend_weight) |
| 256 | 268 | ||
| 257 | def init_kvcache(self, kv_cache): | 269 | def init_kvcache(self, kv_cache): |
| @@ -328,6 +340,11 @@ class FlashGlm41vTextModelForCausalLM(FlashForCausalLM): | |||
| 328 | enable_splitfuse_pa=not self.soc_info.is_300i(), | 340 | enable_splitfuse_pa=not self.soc_info.is_300i(), |
| 329 | **kwargs | 341 | **kwargs |
| 330 | ) | 342 | ) |
| 343 | + self.flash_comm_modifier.modify_inputs( | ||
| 344 | + self.acl_operation_inputs, | ||
| 345 | + is_prefill, | ||
| 346 | + acl_param | ||
| 347 | + ) | ||
| 331 | self.acl_param = json.dumps(acl_param) | 348 | self.acl_param = json.dumps(acl_param) |
| 332 | return self.acl_operation_inputs, self.acl_param | 349 | return self.acl_operation_inputs, self.acl_param |
| 333 | 350 | ||
| @@ -434,4 +451,15 @@ class FlashGlm41vTextModelForCausalLM(FlashForCausalLM): | |||
| 434 | err_msg = "Number of output tensors is not equal to the expected value." | 451 | err_msg = "Number of output tensors is not equal to the expected value." |
| 435 | logger.error(err_msg, ErrorCode.ATB_MODELS_PARAM_OUT_OF_RANGE) | 452 | logger.error(err_msg, ErrorCode.ATB_MODELS_PARAM_OUT_OF_RANGE) |
| 436 | raise RuntimeError(err_msg) | 453 | raise RuntimeError(err_msg) |
| 437 | - return acl_model_out | 454 | + return acl_model_out |
| 455 | + | ||
| 456 | + def _flash_comm_gate(self) -> bool: | ||
| 457 | + soc_version = self.soc_info.soc_version | ||
| 458 | + return not any([ | ||
| 459 | + self.enable_dap, | ||
| 460 | + self.tp_world_size == 1, | ||
| 461 | + soc_version in _800_9000_SOCS, | ||
| 462 | + soc_version in DUO_SOCS and self.tp_world_size > 4, | ||
| 463 | + soc_version in A2_SOCS + A3_SOCS and not self.soc_info.is_support_hccs(), | ||
| 464 | + self.lcoc_enable | ||
| 465 | + ]) | ||
| @@ -48,6 +48,11 @@ from atb_llm.utils.layers import ( | |||
| 48 | ) | 48 | ) |
| 49 | 49 | ||
| 50 | 50 | ||
| 51 | +A2_SOCS = (220, 221, 222, 223, 224, 225) | ||
| 52 | +A3_SOCS = (250, 251, 252, 253, 254, 255) | ||
| 53 | +INITIAL_MAX_GRID_SIZE = 8192 | ||
| 54 | + | ||
| 55 | + | ||
| 51 | class Glm41vVisionPatchEmbed(nn.Module): | 56 | class Glm41vVisionPatchEmbed(nn.Module): |
| 52 | def __init__(self, config, weights, prefix): | 57 | def __init__(self, config, weights, prefix): |
| 53 | super().__init__() | 58 | super().__init__() |
| @@ -100,11 +105,18 @@ class Glm41vVisionRotaryEmbedding(nn.Module): | |||
| 100 | super().__init__() | 105 | super().__init__() |
| 101 | inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float) / dim)) | 106 | inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float) / dim)) |
| 102 | self.register_buffer("inv_freq", inv_freq, persistent=False) | 107 | self.register_buffer("inv_freq", inv_freq, persistent=False) |
| 108 | + self.max_grid_size = INITIAL_MAX_GRID_SIZE | ||
| 109 | + self.freqs = None | ||
| 110 | + | ||
| 111 | + def build_freq_table(self): | ||
| 112 | + seq = torch.arange(self.max_grid_size, device=self.inv_freq.device, dtype=self.inv_freq.dtype) | ||
| 113 | + self.freqs = torch.outer(seq, self.inv_freq) | ||
| 103 | 114 | ||
| 104 | def forward(self, seqlen: int) -> torch.Tensor: | 115 | def forward(self, seqlen: int) -> torch.Tensor: |
| 105 | - seq = torch.arange(seqlen, device=self.inv_freq.device, dtype=self.inv_freq.dtype) | 116 | + if self.freqs is None or seqlen > self.max_grid_size: |
| 106 | - freqs = torch.outer(seq, self.inv_freq) | 117 | + self.max_grid_size = max(seqlen, self.max_grid_size) |
| 107 | - return freqs | 118 | + self.build_freq_table() |
| 119 | + return self.freqs | ||
| 108 | 120 | ||
| 109 | 121 | ||
| 110 | class Glm41vVisionEmbeddings(nn.Module): | 122 | class Glm41vVisionEmbeddings(nn.Module): |
| @@ -116,6 +128,7 @@ class Glm41vVisionEmbeddings(nn.Module): | |||
| 116 | self.patch_size = config.patch_size | 128 | self.patch_size = config.patch_size |
| 117 | self.num_patches = (self.image_size // self.patch_size) ** 2 | 129 | self.num_patches = (self.image_size // self.patch_size) ** 2 |
| 118 | self.num_positions = self.num_patches | 130 | self.num_positions = self.num_patches |
| 131 | + self.soc_info = NPUSocInfo() | ||
| 119 | self.position_embedding = nn.Embedding(self.num_positions, self.embed_dim) | 132 | self.position_embedding = nn.Embedding(self.num_positions, self.embed_dim) |
| 120 | self.position_embedding.weight.data = weights.get_tensor(f"{prefix}.position_embedding.weight") | 133 | self.position_embedding.weight.data = weights.get_tensor(f"{prefix}.position_embedding.weight") |
| 121 | self.register_buffer("position_ids", torch.arange(self.num_positions).expand((1, -1)), persistent=False) | 134 | self.register_buffer("position_ids", torch.arange(self.num_positions).expand((1, -1)), persistent=False) |
| @@ -126,21 +139,19 @@ class Glm41vVisionEmbeddings(nn.Module): | |||
| 126 | hidden_size = pos_embed_weight.shape[1] | 139 | hidden_size = pos_embed_weight.shape[1] |
| 127 | total_seq = h_coords.shape[0] | 140 | total_seq = h_coords.shape[0] |
| 128 | dtype = pos_embed_weight.dtype | 141 | dtype = pos_embed_weight.dtype |
| 129 | - device = torch.device("cpu") | 142 | + device = pos_embed_weight.device |
| 130 | 143 | ||
| 131 | # Move coordinates to correct device | 144 | # Move coordinates to correct device |
| 132 | - h_coords, w_coords = h_coords.cpu(), w_coords.cpu() | 145 | + h_coords, w_coords = h_coords.numpy(), w_coords.numpy() |
| 133 | - | ||
| 134 | # Handle empty sequence case | 146 | # Handle empty sequence case |
| 135 | if total_seq == 0: | 147 | if total_seq == 0: |
| 136 | adapted_pos_embed = torch.empty(0, hidden_size, device=device, dtype=dtype) | 148 | adapted_pos_embed = torch.empty(0, hidden_size, device=device, dtype=dtype) |
| 137 | else: | 149 | else: |
| 138 | # Convert inputs to tensors if needed | 150 | # Convert inputs to tensors if needed |
| 139 | if isinstance(lengths, list): | 151 | if isinstance(lengths, list): |
| 140 | - lengths = torch.tensor(lengths, device=device, dtype=torch.long) | 152 | + lengths = np.array(lengths, dtype=np.int64) |
| 141 | - if not isinstance(image_shapes, torch.Tensor): | 153 | + if not isinstance(image_shapes, np.ndarray): |
| 142 | - image_shapes = torch.tensor(image_shapes, device=device, dtype=torch.long) | 154 | + image_shapes = np.array(image_shapes, dtype=np.int64) |
| 143 | - | ||
| 144 | # Prepare 2D position embedding | 155 | # Prepare 2D position embedding |
| 145 | orig_size_sq = pos_embed_weight.shape[0] | 156 | orig_size_sq = pos_embed_weight.shape[0] |
| 146 | orig_size = int(orig_size_sq**0.5) | 157 | orig_size = int(orig_size_sq**0.5) |
| @@ -150,29 +161,32 @@ class Glm41vVisionEmbeddings(nn.Module): | |||
| 150 | .unsqueeze(0) | 161 | .unsqueeze(0) |
| 151 | .to(device=device, dtype=torch.float32) | 162 | .to(device=device, dtype=torch.float32) |
| 152 | ) | 163 | ) |
| 153 | - | ||
| 154 | # Calculate target dimensions for each patch | 164 | # Calculate target dimensions for each patch |
| 155 | - target_h = torch.cat([image_shapes[i, 1].repeat(lengths[i]) for i in range(len(lengths))]).to( | 165 | + target_h = np.concatenate( |
| 156 | - device=device, dtype=torch.float32 | 166 | + [np.repeat(image_shapes[i, 1], lengths[i]) for i in range(len(lengths))] |
| 157 | - ) | 167 | + ).astype(np.float32) |
| 158 | - target_w = torch.cat([image_shapes[i, 2].repeat(lengths[i]) for i in range(len(lengths))]).to( | 168 | + target_w = np.concatenate( |
| 159 | - device=device, dtype=torch.float32 | 169 | + [np.repeat(image_shapes[i, 2], lengths[i]) for i in range(len(lengths))] |
| 160 | - ) | 170 | + ).astype(np.float32) |
| 161 | - | ||
| 162 | # Normalize coordinates to [-1, 1] range for grid_sample | 171 | # Normalize coordinates to [-1, 1] range for grid_sample |
| 163 | - h_coords = h_coords.to(device=device, dtype=torch.float32) | 172 | + h_coords = h_coords.astype(np.float32) |
| 164 | - w_coords = w_coords.to(device=device, dtype=torch.float32) | 173 | + w_coords = w_coords.astype(np.float32) |
| 165 | norm_w = ((w_coords + 0.5) / target_w) * 2 - 1 | 174 | norm_w = ((w_coords + 0.5) / target_w) * 2 - 1 |
| 166 | norm_h = ((h_coords + 0.5) / target_h) * 2 - 1 | 175 | norm_h = ((h_coords + 0.5) / target_h) * 2 - 1 |
| 167 | - | ||
| 168 | # Create sampling grid | 176 | # Create sampling grid |
| 169 | - grid = torch.stack((norm_w, norm_h), dim=-1).unsqueeze(0).unsqueeze(2) | 177 | + grid = np.stack((norm_w, norm_h), axis=-1) |
| 170 | - | 178 | + grid = torch.tensor(grid, dtype=torch.float32, device=device).unsqueeze(0).unsqueeze(2) |
| 171 | # Perform bicubic interpolation | 179 | # Perform bicubic interpolation |
| 172 | - interpolated_embed_fp32 = F.grid_sample( | 180 | + if self.soc_info.soc_version in A2_SOCS + A3_SOCS: |
| 173 | - pos_embed_2d, grid, mode="bicubic", align_corners=False, padding_mode="border" | 181 | + interpolated_embed_fp32 = F.grid_sample( |
| 174 | - ) | 182 | + pos_embed_2d, grid, mode="bicubic", align_corners=False, padding_mode="border" |
| 175 | - | 183 | + ) |
| 184 | + else: | ||
| 185 | + # GridSample2D in bicubic mode is not supported in 300I DUO cards | ||
| 186 | + upsampled = F.interpolate(pos_embed_2d, scale_factor=2, mode='bilinear', align_corners=False) | ||
| 187 | + interpolated_embed_fp32 = F.grid_sample(upsampled, grid, | ||
| 188 | + mode='bilinear', align_corners=False, | ||
| 189 | + padding_mode='zeros') | ||
| 176 | # Reshape and convert back to original dtype | 190 | # Reshape and convert back to original dtype |
| 177 | adapted_pos_embed_fp32 = interpolated_embed_fp32.squeeze(0).squeeze(-1).permute(1, 0) | 191 | adapted_pos_embed_fp32 = interpolated_embed_fp32.squeeze(0).squeeze(-1).permute(1, 0) |
| 178 | adapted_pos_embed = adapted_pos_embed_fp32.to(dtype) | 192 | adapted_pos_embed = adapted_pos_embed_fp32.to(dtype) |
| @@ -847,6 +861,7 @@ class Glm41vVisionModel(nn.Module): | |||
| 847 | def get_adapted_pos_embed(self, grid_thw): | 861 | def get_adapted_pos_embed(self, grid_thw): |
| 848 | grid_thw = torch.tensor(np.array(json.loads(grid_thw))) | 862 | grid_thw = torch.tensor(np.array(json.loads(grid_thw))) |
| 849 | rotary_pos_emb, image_type_ids = self.rot_pos_emb(grid_thw) | 863 | rotary_pos_emb, image_type_ids = self.rot_pos_emb(grid_thw) |
| 864 | + rotary_pos_emb = rotary_pos_emb.npu() | ||
| 850 | emb = torch.cat((rotary_pos_emb, rotary_pos_emb), dim=-1) | 865 | emb = torch.cat((rotary_pos_emb, rotary_pos_emb), dim=-1) |
| 851 | position_embeddings = (emb.cos(), emb.sin()) | 866 | position_embeddings = (emb.cos(), emb.sin()) |
| 852 | 867 | ||
| @@ -33,21 +33,19 @@ bool CheckGlm41vFusionAttentionParam( | |||
| 33 | return true; | 33 | return true; |
| 34 | } | 34 | } |
| 35 | 35 | ||
| 36 | -TEST(Glm41vDecoderLayerTest, Glm41vDecoderLayer) | 36 | +TEST(Glm41vDecoderLayerTest, DecoderLayer) |
| 37 | { | 37 | { |
| 38 | GlobalMockObject::verify(); | 38 | GlobalMockObject::verify(); |
| 39 | 39 | ||
| 40 | - atb_speed::glm41v::Glm41vLayerParam layerParam; | 40 | + atb_speed::base::LayerParam layerParam; |
| 41 | layerParam.hiddenSizePerAttentionHead = NUM128; | 41 | layerParam.hiddenSizePerAttentionHead = NUM128; |
| 42 | 42 | ||
| 43 | - atb_speed::glm41v::Glm41vDecoderLayer decoderLayer(layerParam); | 43 | + atb_speed::glm41v::DecoderLayer decoderLayer(layerParam); |
| 44 | decoderLayer.ConstructInTensorMap(); | 44 | decoderLayer.ConstructInTensorMap(); |
| 45 | - EXPECT_TRUE(decoderLayer.graph.inTensorNum == NUM63); | 45 | + EXPECT_TRUE(decoderLayer.inTensorList.size() == NUM63); |
| 46 | atb_speed::common::FusionAttentionParam<atb::infer::RmsNormParam> fusionAttentionParam; | 46 | atb_speed::common::FusionAttentionParam<atb::infer::RmsNormParam> fusionAttentionParam; |
| 47 | decoderLayer.SetFusionAttentionParam(fusionAttentionParam); | 47 | decoderLayer.SetFusionAttentionParam(fusionAttentionParam); |
| 48 | EXPECT_TRUE(CheckGlm41vFusionAttentionParam(fusionAttentionParam)); | 48 | EXPECT_TRUE(CheckGlm41vFusionAttentionParam(fusionAttentionParam)); |
| 49 | - decoderLayer.SetFusionAttentionNormParam(fusionAttentionParam); | ||
| 50 | - EXPECT_FALSE(fusionAttentionParam.enableNormQuantOp); | ||
| 51 | MOCKER(atb::CreateOperation<atb::infer::RmsNormParam>).expects(atLeast(1)) | 49 | MOCKER(atb::CreateOperation<atb::infer::RmsNormParam>).expects(atLeast(1)) |
| 52 | .with(any(), any()).will(returnValue(0)); | 50 | .with(any(), any()).will(returnValue(0)); |
| 53 | atb::Status ret = decoderLayer.AddPostSelfAttentionRMSNorm(); | 51 | atb::Status ret = decoderLayer.AddPostSelfAttentionRMSNorm(); |
| @@ -17,7 +17,7 @@ | |||
| 17 | 17 | ||
| 18 | namespace atb_speed { | 18 | namespace atb_speed { |
| 19 | 19 | ||
| 20 | -bool CheckGlm41vWeightCountPerLayer(const atb_speed::glm41v::Glm41vDecoderModel &decoderModel) | 20 | +bool CheckGlm41vWeightCountPerLayer(const atb_speed::glm41v::DecoderModel &decoderModel) |
| 21 | { | 21 | { |
| 22 | constexpr int GLM4_WEIGHT_COUNT_PER_LAYER = 52; | 22 | constexpr int GLM4_WEIGHT_COUNT_PER_LAYER = 52; |
| 23 | if (decoderModel.weightCountPerLayer != GLM4_WEIGHT_COUNT_PER_LAYER) { | 23 | if (decoderModel.weightCountPerLayer != GLM4_WEIGHT_COUNT_PER_LAYER) { |
| @@ -26,7 +26,7 @@ bool CheckGlm41vWeightCountPerLayer(const atb_speed::glm41v::Glm41vDecoderModel | |||
| 26 | return true; | 26 | return true; |
| 27 | } | 27 | } |
| 28 | 28 | ||
| 29 | -TEST(Glm41vDecoderModelTest, Glm41vDecoderModel) | 29 | +TEST(Glm41vDecoderModelTest, DecoderModel) |
| 30 | { | 30 | { |
| 31 | GlobalMockObject::verify(); | 31 | GlobalMockObject::verify(); |
| 32 | 32 | ||
| @@ -44,11 +44,15 @@ TEST(Glm41vDecoderModelTest, Glm41vDecoderModel) | |||
| 44 | "\"enableSwiGLU\": false, " | 44 | "\"enableSwiGLU\": false, " |
| 45 | "\"rank\": 0, \"worldSize\": 2, \"backend\": \"hccl\", " | 45 | "\"rank\": 0, \"worldSize\": 2, \"backend\": \"hccl\", " |
| 46 | "\"positionEmbeddingType\": 0, \"linearHasBias\": [[true, false, false, false]], " | 46 | "\"positionEmbeddingType\": 0, \"linearHasBias\": [[true, false, false, false]], " |
| 47 | - "\"isPrefill\": false, \"supportLcoc\": false}"; | 47 | + "\"isPrefill\": false, \"enableLcoc\": false}"; |
| 48 | - atb_speed::glm41v::Glm41vDecoderModel decoderModel(param); | 48 | + atb_speed::glm41v::DecoderModel decoderModel(param); |
| 49 | EXPECT_TRUE(CheckGlm41vWeightCountPerLayer(decoderModel)); | 49 | EXPECT_TRUE(CheckGlm41vWeightCountPerLayer(decoderModel)); |
| 50 | + decoderModel.ConstructInTensorMap(); | ||
| 51 | + MOCKER(atb::CreateOperation<atb::GraphParam>).expects(atLeast(1)) | ||
| 52 | + .with(any(), any()).will(returnValue(0)); | ||
| 50 | atb::Operation *op = nullptr; | 53 | atb::Operation *op = nullptr; |
| 51 | atb::Status ret = decoderModel.CreateLayerOperation(&op, 0); | 54 | atb::Status ret = decoderModel.CreateLayerOperation(&op, 0); |
| 55 | + EXPECT_EQ(ret, atb::NO_ERROR); | ||
| 52 | } | 56 | } |
| 53 | 57 | ||
| 54 | } // namespace atb_speed | 58 | } // namespace atb_speed |
如果h_coords和w_coords是在npu上的tensor,无法直接转为numpy,需要先转cpu在转numpy