MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / ChatGLMModel

Class ChatGLMModel

SwissArmyTransformer/sat/model/official/chatglm_model.py:165–265  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

163 return output
164
165class ChatGLMModel(BaseModel):
166 def __init__(self, args, transformer=None, **kwargs):
167 super(ChatGLMModel, self).__init__(args, transformer=transformer, activation_func=gelu, **kwargs)
168 del self.transformer.position_embeddings
169 self.add_mixin("chatglm-final", ChatGLMFinalMixin(args.vocab_size, args.hidden_size))
170 self.add_mixin("chatglm-attn", ChatGLMAttnMixin(args.hidden_size, args.num_attention_heads))
171 self.add_mixin("chatglm-layer", ChatGLMLayerMixin(args.num_layers))
172 self.bos_token_id = args.bos_token_id
173 self.mask_token_id = args.mask_token_id
174 self.gmask_token_id = args.gmask_token_id
175 self.pad_token_id = args.pad_token_id
176
177 def position_embedding_forward(self, position_ids, output_cross_layer, **kw_args):
178 return None
179
180 def forward(self, input_ids, position_ids=None, attention_mask=None, past_key_values=None, **kwargs):
181 if attention_mask is None and position_ids is None:
182 attention_mask, position_ids = self.get_inputs(input_ids, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, **kwargs)
183 if attention_mask is not None and attention_mask.dtype is torch.bool:
184 attention_mask = (~attention_mask).long()
185 if past_key_values is not None:
186 input_ids = input_ids[:, -1:]
187 position_ids = position_ids[..., -1:]
188 if input_ids.size(0) != 1:
189 attention_mask = attention_mask[:, :, -1:]
190 return super().forward(input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, **kwargs)
191
192 def get_inputs(self, input_ids, attention_mask=None, position_ids=None, past_key_values=None, **kwargs):
193 if attention_mask is None:
194 if past_key_values is not None and input_ids.size(0) == 1:
195 attention_mask = torch.tensor([[1]], dtype=torch.long, device=input_ids.device)
196 else:
197 attention_mask = self.get_masks(
198 input_ids=input_ids,
199 device=input_ids.device, **kwargs
200 )
201 if position_ids is None:
202 MASK, gMASK = self.mask_token_id, self.gmask_token_id
203 mask_token = gMASK if gMASK in input_ids else MASK
204 use_gmask = True if gMASK in input_ids else False
205
206 mask_positions = [seq.tolist().index(mask_token) for seq in input_ids]
207 position_ids = self.get_position_ids(
208 input_ids=input_ids,
209 mask_positions=mask_positions,
210 device=input_ids.device,
211 gmask=use_gmask, **kwargs
212 )
213 return attention_mask, position_ids
214
215 def get_pad_length(self, seq):
216 l = 0
217 while l < len(seq) and seq[l] == self.pad_token_id:
218 l += 1
219 return l
220
221 def get_masks(self, input_ids, device, **kwargs):
222 batch_size, seq_length = input_ids.shape

Callers 1

transform_param.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected