| 124 | |
| 125 | @staticmethod |
| 126 | def chat_completion( |
| 127 | model, |
| 128 | tokenizer, |
| 129 | dialogs: List[Dialog], |
| 130 | temperature: float = 0.6, |
| 131 | top_p: float = 0.9, |
| 132 | max_gen_len: Optional[int] = None, |
| 133 | logprobs: bool = False, |
| 134 | ) -> List[ChatPrediction]: |
| 135 | if max_gen_len is None: |
| 136 | max_gen_len = model.params.max_seq_len - 1 |
| 137 | prompt_tokens = [] |
| 138 | for dialog in dialogs: |
| 139 | if dialog[0]["role"] != "system": |
| 140 | dialog = [ |
| 141 | { |
| 142 | "role": "system", |
| 143 | "content": DEFAULT_SYSTEM_PROMPT, |
| 144 | } |
| 145 | ] + dialog |
| 146 | dialog = [ |
| 147 | { |
| 148 | "role": dialog[1]["role"], |
| 149 | "content": B_SYS |
| 150 | + dialog[0]["content"] |
| 151 | + E_SYS |
| 152 | + dialog[1]["content"], |
| 153 | } |
| 154 | ] + dialog[2:] |
| 155 | |
| 156 | assert all([msg["role"] == "user" for msg in dialog[::2]]) and all( |
| 157 | [msg["role"] == "assistant" for msg in dialog[1::2]] |
| 158 | ), ( |
| 159 | "model only supports 'system', 'user' and 'assistant' roles, " |
| 160 | "starting with 'system', then 'user' and alternating (u/a/u/a/u...)" |
| 161 | ) |
| 162 | |
| 163 | dialog_tokens: List[int] = sum( |
| 164 | [ |
| 165 | tokenizer.encode( |
| 166 | f"{B_INST} {(prompt['content']).strip()} {E_INST} {(answer['content']).strip()} ", |
| 167 | bos=True, |
| 168 | eos=True, |
| 169 | ) |
| 170 | for prompt, answer in zip( |
| 171 | dialog[::2], |
| 172 | dialog[1::2], |
| 173 | ) |
| 174 | ], |
| 175 | [], |
| 176 | ) |
| 177 | assert ( |
| 178 | dialog[-1]["role"] == "user" |
| 179 | ), f"Last message must be from user, got {dialog[-1]['role']}" |
| 180 | dialog_tokens += tokenizer.encode( |
| 181 | f"{B_INST} {(dialog[-1]['content']).strip()} {E_INST}", |
| 182 | bos=True, |
| 183 | eos=False, |