(
self,
source_sequence_id: int,
dest_sequence_id: int,
keep_len: int,
*,
copy_all_state: bool = False,
)
| 14878 | self.truncate_sequence_metadata(seq_id, current_len, keep_len) |
| 14879 | |
| 14880 | def copy_sequence_state( |
| 14881 | self, |
| 14882 | source_sequence_id: int, |
| 14883 | dest_sequence_id: int, |
| 14884 | keep_len: int, |
| 14885 | *, |
| 14886 | copy_all_state: bool = False, |
| 14887 | ) -> None: |
| 14888 | if keep_len <= 0: |
| 14889 | return |
| 14890 | source_length = self.radix_trie.length(source_sequence_id) |
| 14891 | keep_pos = self.sequence_history.position_length_for_prefix( |
| 14892 | source_sequence_id, |
| 14893 | keep_len, |
| 14894 | ) |
| 14895 | copy_p0 = 0 |
| 14896 | copy_p1 = keep_pos |
| 14897 | if copy_all_state or not self.model.kv_unified: |
| 14898 | copy_p0 = -1 |
| 14899 | copy_p1 = -1 |
| 14900 | llama_cpp.llama_memory_seq_cp( |
| 14901 | self.model.mem, |
| 14902 | source_sequence_id, |
| 14903 | dest_sequence_id, |
| 14904 | copy_p0, |
| 14905 | copy_p1, |
| 14906 | ) |
| 14907 | self.model.copy_draft_sequence( |
| 14908 | source_sequence_id, |
| 14909 | dest_sequence_id, |
| 14910 | copy_p0, |
| 14911 | copy_p1, |
| 14912 | ) |
| 14913 | if copy_all_state and source_length > keep_len: |
| 14914 | if not llama_cpp.llama_memory_seq_rm( |
| 14915 | self.model.mem, |
| 14916 | dest_sequence_id, |
| 14917 | keep_pos, |
| 14918 | -1, |
| 14919 | ): |
| 14920 | raise RuntimeError( |
| 14921 | f"failed to truncate copied model sequence {dest_sequence_id} " |
| 14922 | f"at position {keep_pos}" |
| 14923 | ) |
| 14924 | self.model.truncate_draft_sequence(dest_sequence_id, keep_pos) |
| 14925 | self.radix_trie.copy(source_sequence_id, dest_sequence_id, keep_len) |
| 14926 | self.sequence_history.copy( |
| 14927 | source_sequence_id, |
| 14928 | dest_sequence_id, |
| 14929 | source_length, |
| 14930 | keep_len, |
| 14931 | ) |
| 14932 | |
| 14933 | def truncate_sequence_metadata( |
| 14934 | self, |
no test coverage detected