(
self, context: str, continuation: str, decode_mode="default"
)
| 812 | return logits |
| 813 | |
| 814 | def _encode_pair( |
| 815 | self, context: str, continuation: str, decode_mode="default" |
| 816 | ) -> Tuple[List[int], List[int]]: |
| 817 | |
| 818 | n_spaces = len(context) - len(context.rstrip()) |
| 819 | if n_spaces > 0: |
| 820 | continuation = context[-n_spaces:] + continuation |
| 821 | context = context[:-n_spaces] |
| 822 | |
| 823 | whole_enc = self.tok_encode(context + continuation, add_special_tokens=False) |
| 824 | context_enc = self.tok_encode(context, add_special_tokens=False) |
| 825 | |
| 826 | def custom_enc_dec_encode( |
| 827 | encodings, |
| 828 | whole_enc, |
| 829 | pad_token_id: int = 220, #50256, |
| 830 | enc_length: int=1024, |
| 831 | ): |
| 832 | encodings = torch.tensor(encodings).unsqueeze(0) |
| 833 | whole_enc = torch.tensor(whole_enc).unsqueeze(0) |
| 834 | batch, seqlen = whole_enc.shape |
| 835 | if seqlen < enc_length: |
| 836 | prefill_size_pad = enc_length - seqlen |
| 837 | padding = torch.full((batch, prefill_size_pad), pad_token_id, dtype=torch.long, device=encodings.device) |
| 838 | encodings = torch.cat([padding, encodings], dim=-1) |
| 839 | whole_enc = torch.cat([padding, whole_enc], dim=-1) |
| 840 | attn_mask = torch.ones_like(whole_enc) |
| 841 | attn_mask[:, :prefill_size_pad] = 0 |
| 842 | |
| 843 | # convert back to list |
| 844 | encodings = encodings.squeeze(0).tolist() |
| 845 | whole_enc = whole_enc.squeeze(0).tolist() |
| 846 | return encodings, whole_enc, attn_mask, enc_length |
| 847 | else: |
| 848 | # attn_mask = torch.ones_like(whole_enc) |
| 849 | attn_mask=None |
| 850 | encodings = encodings.squeeze(0).tolist() |
| 851 | whole_enc = whole_enc.squeeze(0).tolist() |
| 852 | return encodings, whole_enc, attn_mask, seqlen |
| 853 | if decode_mode == "default_left_pad": |
| 854 | context_enc, whole_enc, attn_mask, seqlen = custom_enc_dec_encode(context_enc, whole_enc) |
| 855 | else: |
| 856 | attn_mask = None |
| 857 | |
| 858 | # whole_enc = self.tok_encode(context + continuation) |
| 859 | # context_enc = self.tok_encode(context, add_special_tokens=False) |
| 860 | context_enc_len = len(context_enc) |
| 861 | continuation_enc = whole_enc[context_enc_len:] |
| 862 | return context_enc, continuation_enc, attn_mask |
| 863 | |
| 864 | def loglikelihood(self, requests: List[Instance], decode_mode="default") -> List[Tuple[float, bool]]: |
| 865 | new_reqs = [] |
no test coverage detected