()
| 212 | nv.set_neuron_interventions(interventions) |
| 213 | |
| 214 | def event_generator(): |
| 215 | gen = cast( |
| 216 | Generator[int | GenerateOutput, None, None], |
| 217 | nv.send_message( |
| 218 | SUBJECT, |
| 219 | message, |
| 220 | max_new_tokens=max_new_tokens or 64, |
| 221 | temperature=temperature or 1.0, |
| 222 | stream=True, |
| 223 | ), |
| 224 | ) |
| 225 | chat_tokens: list[ChatToken] = [] |
| 226 | |
| 227 | for update in gen: |
| 228 | if isinstance(update, int): |
| 229 | chat_tokens.append(ChatToken(token=SUBJECT.decode(update))) |
| 230 | data = json.dumps([ct.model_dump() for ct in chat_tokens]) |
| 231 | yield f"data: {data}\n\n" |
| 232 | else: |
| 233 | # Collate tokenwise log probs |
| 234 | index_to_log_probs: dict[int, list[tuple[str, float]]] = {} |
| 235 | tokenwise_log_probs = update.tokenwise_log_probs |
| 236 | for i, (token_ids, log_probs) in enumerate(tokenwise_log_probs): |
| 237 | cur_log_probs = [ |
| 238 | (SUBJECT.decode(token_id), log_prob) |
| 239 | for token_id, log_prob in zip(token_ids[0], log_probs[0]) |
| 240 | ] |
| 241 | # output_ids_BT.shape[1] is the total sequence length, len(tokenwise_log_probs) is the number of tokens generated |
| 242 | index_to_log_probs[ |
| 243 | i + update.output_ids_BT.shape[1] - len(tokenwise_log_probs) |
| 244 | ] = cur_log_probs |
| 245 | |
| 246 | # Update chat tokens with top log probs |
| 247 | chat_tokens = [ |
| 248 | ChatToken(**(ct.model_dump() | {"top_log_probs": index_to_log_probs.get(i)})) |
| 249 | for i, ct in enumerate(chat_tokens) |
| 250 | ] |
| 251 | data = json.dumps([ct.model_dump() for ct in chat_tokens]) |
| 252 | |
| 253 | yield f"data: {data}\n\n" |
| 254 | yield "data: [DONE]\n\n" |
| 255 | |
| 256 | nv.clear_neuron_interventions() |
| 257 | SESSIONS[session_id] = (ChatConversation(**nv.model_input.model_dump()), nv.filter) |
| 258 | |
| 259 | return StreamingResponse(event_generator(), media_type="text/event-stream") |
| 260 |
no test coverage detected