| 3914 | |
| 3915 | @dataclass |
| 3916 | class PromptPlan: |
| 3917 | text: str |
| 3918 | generation_prompt: str |
| 3919 | text_tokens: List[int] |
| 3920 | identity_tokens: List[int] |
| 3921 | segments: List[PromptSegment] |
| 3922 | text_token_index_by_pos: Dict[int, int] = field(default_factory=dict) |
| 3923 | |
| 3924 | @classmethod |
| 3925 | def from_tokens( |
| 3926 | cls, |
| 3927 | text: str, |
| 3928 | tokens: List[int], |
| 3929 | *, |
| 3930 | generation_prompt: str = "", |
| 3931 | ) -> "PromptPlan": |
| 3932 | segment = PromptSegment( |
| 3933 | kind="text", |
| 3934 | start_pos=0, |
| 3935 | n_pos=len(tokens), |
| 3936 | identity_tokens=list(tokens), |
| 3937 | decode_start_pos=0, |
| 3938 | decode_n_pos=len(tokens), |
| 3939 | text_tokens=list(tokens), |
| 3940 | ) |
| 3941 | return cls( |
| 3942 | text=text, |
| 3943 | generation_prompt=generation_prompt, |
| 3944 | text_tokens=list(tokens), |
| 3945 | identity_tokens=list(tokens), |
| 3946 | segments=[segment] if tokens else [], |
| 3947 | text_token_index_by_pos={pos: pos for pos in range(len(tokens))}, |
| 3948 | ) |
| 3949 | |
| 3950 | @property |
| 3951 | def length(self) -> int: |
| 3952 | return len(self.identity_tokens) |
| 3953 | |
| 3954 | @property |
| 3955 | def eval_token_count(self) -> int: |
| 3956 | return self.length |
| 3957 | |
| 3958 | def position_increments_up_to(self, pos: int) -> List[int]: |
| 3959 | increments: List[int] = [] |
| 3960 | for segment in self.segments: |
| 3961 | if pos <= segment.start_pos: |
| 3962 | break |
| 3963 | take = min(pos, segment.end_pos) - segment.start_pos |
| 3964 | if take <= 0: |
| 3965 | continue |
| 3966 | increments.extend(segment.decoder_position_increments[:take]) |
| 3967 | if pos <= segment.end_pos: |
| 3968 | break |
| 3969 | return increments |
| 3970 | |
| 3971 | def is_boundary(self, pos: int) -> bool: |
| 3972 | if pos <= 0 or pos >= self.length: |
| 3973 | return True |
no outgoing calls
no test coverage detected
searching dependent graphs…