| 721 | |
| 722 | |
| 723 | class ChatGLMForConditionalGeneration(ChatGLMPreTrainedModel): |
| 724 | def __init__(self, config: ChatGLMConfig, empty_init=True, device=None): |
| 725 | super().__init__(config) |
| 726 | |
| 727 | self.max_sequence_length = config.max_length |
| 728 | self.transformer = ChatGLMModel(config, empty_init=empty_init, device=device) |
| 729 | self.config = config |
| 730 | self.pack_loss = False |
| 731 | |
| 732 | def _update_model_kwargs_for_generation( |
| 733 | self, |
| 734 | outputs: ModelOutput, |
| 735 | model_kwargs: Dict[str, Any], |
| 736 | is_encoder_decoder: bool = False, |
| 737 | standardize_cache_format: bool = False, |
| 738 | ) -> Dict[str, Any]: |
| 739 | # update past_key_values |
| 740 | model_kwargs["past_key_values"] = self._extract_past_from_model_output( |
| 741 | outputs, standardize_cache_format=standardize_cache_format |
| 742 | ) |
| 743 | |
| 744 | # update attention mask |
| 745 | if "attention_mask" in model_kwargs: |
| 746 | attention_mask = model_kwargs["attention_mask"] |
| 747 | model_kwargs["attention_mask"] = torch.cat( |
| 748 | [attention_mask, attention_mask.new_ones((attention_mask.shape[0], 1))], dim=-1 |
| 749 | ) |
| 750 | |
| 751 | # update position ids |
| 752 | if "position_ids" in model_kwargs: |
| 753 | position_ids = model_kwargs["position_ids"] |
| 754 | new_position_id = position_ids[..., -1:].clone() |
| 755 | new_position_id += 1 |
| 756 | model_kwargs["position_ids"] = torch.cat( |
| 757 | [position_ids, new_position_id], dim=-1 |
| 758 | ) |
| 759 | |
| 760 | model_kwargs["is_first_forward"] = False |
| 761 | return model_kwargs |
| 762 | |
| 763 | def prepare_inputs_for_generation( |
| 764 | self, |
| 765 | input_ids: torch.LongTensor, |
| 766 | past_key_values: Optional[torch.Tensor] = None, |
| 767 | attention_mask: Optional[torch.Tensor] = None, |
| 768 | position_ids: Optional[torch.Tensor] = None, |
| 769 | is_first_forward: bool = True, |
| 770 | **kwargs |
| 771 | ) -> dict: |
| 772 | # only last token for input_ids if past is not None |
| 773 | if position_ids is None: |
| 774 | position_ids = self.get_position_ids(input_ids, device=input_ids.device) |
| 775 | if not is_first_forward: |
| 776 | position_ids = position_ids[..., -1:] |
| 777 | input_ids = input_ids[:, -1:] |
| 778 | return { |
| 779 | "input_ids": input_ids, |
| 780 | "past_key_values": past_key_values, |
nothing calls this directly
no outgoing calls
no test coverage detected