MCPcopy Create free account
hub / github.com/ZinYY/TreeLoRA / chat_completion

Method chat_completion

utils/data/data_utils.py:126–209  ·  view source on GitHub ↗
(
        model,
        tokenizer,
        dialogs: List[Dialog],
        temperature: float = 0.6,
        top_p: float = 0.9,
        max_gen_len: Optional[int] = None,
        logprobs: bool = False,
    )

Source from the content-addressed store, hash-verified

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,

Callers

nothing calls this directly

Calls 1

generateMethod · 0.45

Tested by

no test coverage detected