已合并
【新需求】GLM4.1V适配lcoc, flash comm和权重预取 #251
凡银创建于 1月15日
【新需求】GLM4.1V适配lcoc, flash comm和权重预取 #251
已合并
凡银创建于 1月15日
从已删除 :glm4v合入到Ascend/MindIE-LLMdev
共 8 个文件变更+111-90
@@ -16,11 +16,9 @@ namespace glm41v {
16 16 
17static const uint64_t NUM2 = 2;17static const uint64_t NUM2 = 2;
18 18 
19-Glm41vDecoderLayer::Glm41vDecoderLayer(19+DecoderLayer::DecoderLayer(
20- const Glm41vLayerParam &param) : atb_speed::base::DecoderLayer<atb::infer::RmsNormParam>(param)20+ const atb_speed::base::LayerParam &param) : 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 @@
20namespace atb_speed {20namespace atb_speed {
21namespace glm41v {21namespace 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> {
26public:25public:
27- explicit Glm41vDecoderLayer(const Glm41vLayerParam &param);26+ explicit DecoderLayer(const atb_speed::base::LayerParam &param);
28- ~Glm41vDecoderLayer() override {};27+ ~DecoderLayer() override {};
29protected:28protected:
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 count17// Weight count
18const uint64_t GLM4_WEIGHT_COUNT_PER_LAYER = 52;18const uint64_t GLM4_WEIGHT_COUNT_PER_LAYER = 52;
19 19 
20-Glm41vDecoderModel::Glm41vDecoderModel(const std::string &param) : atb_speed::base::DecoderModel(param)20+DecoderModel::DecoderModel(const std::string &param) : 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 @@
22namespace atb_speed {22namespace atb_speed {
23namespace glm41v {23namespace 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 {
28public:26public:
29- explicit Glm41vDecoderModel(const std::string &param);27+ explicit DecoderModel(const std::string &param);
30protected:28protected:
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 glm41v34} // namespace glm41v
38} // namespace atb_speed35} // namespace atb_speed
39#endif36#endif
@@ -27,9 +27,10 @@ from torch import nn
27import torch_npu27import torch_npu
28from atb_llm.models.base.flash_causal_lm import FlashForCausalLM28from atb_llm.models.base.flash_causal_lm import FlashForCausalLM
29from atb_llm.models.base.modeling import FlashAttention, FlashLayer, MLP29from atb_llm.models.base.modeling import FlashAttention, FlashLayer, MLP
30-from atb_llm.models.base.graph_manager.graph_manager import ATBGraphManager
31from atb_llm.models.base.inputs_modifier.qlen_modifier import QLenModifier30from atb_llm.models.base.inputs_modifier.qlen_modifier import QLenModifier
32-from atb_llm.models.base.graph_manager import DapGraphWrapper, SpeculateGraphWrapper31+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
33from atb_llm.utils.initial import NPUSocInfo34from atb_llm.utils.initial import NPUSocInfo
34from atb_llm.utils.layers import TensorParallelRowLinear, RMSNorm, TensorEmbedding, TensorHead, \35from atb_llm.utils.layers import TensorParallelRowLinear, RMSNorm, TensorEmbedding, TensorHead, \
35 load_column_multi, PositionRotaryEmbedding, AttentionMask36 load_column_multi, PositionRotaryEmbedding, AttentionMask
@@ -41,6 +42,12 @@ from atb_llm.utils.log.error_code import ErrorCode
41from atb_llm.utils.log import logger42from 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+ 
44class Glm41vTextAttention(FlashAttention):51class 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 None128 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_heads130 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 management177 # 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 pass183 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 True250+ "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": False256+ "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 **kwargs341 **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_param349 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_out454+ 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+ 
51class Glm41vVisionPatchEmbed(nn.Module):56class 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 freqs118+ self.build_freq_table()
119+ return self.freqs
108 120 
109 121 
110class Glm41vVisionEmbeddings(nn.Module):122class Glm41vVisionEmbeddings(nn.Module):
@@ -116,6 +128,7 @@ class Glm41vVisionEmbeddings(nn.Module):
116 self.patch_size = config.patch_size128 self.patch_size = config.patch_size
117 self.num_patches = (self.image_size // self.patch_size) ** 2129 self.num_patches = (self.image_size // self.patch_size) ** 2
118 self.num_positions = self.num_patches130 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.dtype141 dtype = pos_embed_weight.dtype
129- device = torch.device("cpu")142+ device = pos_embed_weight.device
130 143 
131 # Move coordinates to correct device144 # 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()
lijian
lijianlijian1月16日

如果h_coords和w_coords是在npu上的tensor,无法直接转为numpy,需要先转cpu在转numpy

likedislike
凡银
1月16日 评论:
133- 
134 # Handle empty sequence case146 # 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 needed150 # 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 embedding155 # 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 patch164 # 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.float32166+ [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.float32169+ [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_sample171 # 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 - 1174 norm_w = ((w_coords + 0.5) / target_w) * 2 - 1
166 norm_h = ((h_coords + 0.5) / target_h) * 2 - 1175 norm_h = ((h_coords + 0.5) / target_h) * 2 - 1
167- 
168 # Create sampling grid176 # 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 interpolation179 # 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 dtype190 # 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 
18namespace atb_speed {18namespace 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_speed58} // namespace atb_speed