| 163 | return output |
| 164 | |
| 165 | class 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 |