feat: add cancelation support for chat generation and streaming tasks
This commit is contained in:
parent
9a4f71c939
commit
571959e383
2 changed files with 135 additions and 38 deletions
|
|
@ -489,6 +489,7 @@ class InferenceBackend:
|
|||
def generate_with_adapter_control(
|
||||
self,
|
||||
use_adapter: Optional[Union[bool, str]] = None,
|
||||
cancel_event=None,
|
||||
**gen_kwargs,
|
||||
) -> Generator[str, None, None]:
|
||||
"""
|
||||
|
|
@ -505,7 +506,7 @@ class InferenceBackend:
|
|||
with self._generation_lock:
|
||||
self._apply_adapter_state(use_adapter)
|
||||
# Delegate to the lock-free generation path
|
||||
yield from self._generate_chat_response_inner(**gen_kwargs)
|
||||
yield from self._generate_chat_response_inner(cancel_event=cancel_event, **gen_kwargs)
|
||||
|
||||
def generate_chat_response(self,
|
||||
messages: list,
|
||||
|
|
@ -515,7 +516,8 @@ class InferenceBackend:
|
|||
top_p: float = 0.9,
|
||||
top_k: int = 40,
|
||||
max_new_tokens: int = 256,
|
||||
repetition_penalty: float = 1.1) -> Generator[str, None, None]:
|
||||
repetition_penalty: float = 1.1,
|
||||
cancel_event=None) -> Generator[str, None, None]:
|
||||
"""
|
||||
Generate response for text or vision models.
|
||||
Acquires the generation lock. For adapter-controlled generation,
|
||||
|
|
@ -531,6 +533,7 @@ class InferenceBackend:
|
|||
top_k=top_k,
|
||||
max_new_tokens=max_new_tokens,
|
||||
repetition_penalty=repetition_penalty,
|
||||
cancel_event=cancel_event,
|
||||
)
|
||||
|
||||
def _generate_chat_response_inner(self,
|
||||
|
|
@ -541,7 +544,8 @@ class InferenceBackend:
|
|||
top_p: float = 0.9,
|
||||
top_k: int = 40,
|
||||
max_new_tokens: int = 256,
|
||||
repetition_penalty: float = 1.1) -> Generator[str, None, None]:
|
||||
repetition_penalty: float = 1.1,
|
||||
cancel_event=None) -> Generator[str, None, None]:
|
||||
"""
|
||||
Inner generation logic (no lock). Called by both generate_chat_response
|
||||
and generate_with_adapter_control.
|
||||
|
|
@ -558,7 +562,8 @@ class InferenceBackend:
|
|||
# Vision model generation
|
||||
yield from self._generate_vision_response(
|
||||
messages, system_prompt, image,
|
||||
temperature, top_p, top_k, max_new_tokens, repetition_penalty
|
||||
temperature, top_p, top_k, max_new_tokens, repetition_penalty,
|
||||
cancel_event=cancel_event,
|
||||
)
|
||||
else:
|
||||
# Text model: Use training pipeline approach
|
||||
|
|
@ -600,12 +605,13 @@ class InferenceBackend:
|
|||
|
||||
# Step 3: Generate
|
||||
yield from self.generate_stream(
|
||||
formatted_prompt, temperature, top_p, top_k, max_new_tokens, repetition_penalty
|
||||
formatted_prompt, temperature, top_p, top_k, max_new_tokens, repetition_penalty,
|
||||
cancel_event=cancel_event,
|
||||
)
|
||||
|
||||
def _generate_vision_response(self, messages, system_prompt, image,
|
||||
temperature, top_p, top_k, max_new_tokens,
|
||||
repetition_penalty) -> Generator[str, None, None]:
|
||||
repetition_penalty, cancel_event=None) -> Generator[str, None, None]:
|
||||
"""Handle vision model generation with true token-by-token streaming."""
|
||||
model_info = self.models[self.active_model_name]
|
||||
model = model_info["model"]
|
||||
|
|
@ -651,7 +657,10 @@ class InferenceBackend:
|
|||
import threading
|
||||
|
||||
streamer = TextIteratorStreamer(
|
||||
processor.tokenizer, skip_prompt=True, skip_special_tokens=True
|
||||
processor.tokenizer,
|
||||
skip_prompt=True,
|
||||
skip_special_tokens=True,
|
||||
timeout=0.2,
|
||||
)
|
||||
|
||||
generation_kwargs = dict(
|
||||
|
|
@ -664,23 +673,50 @@ class InferenceBackend:
|
|||
top_k=top_k,
|
||||
)
|
||||
|
||||
err: dict[str, str] = {}
|
||||
|
||||
def generate_fn():
|
||||
try:
|
||||
model.generate(**generation_kwargs)
|
||||
except Exception as e:
|
||||
err["msg"] = str(e)
|
||||
logger.error(f"Vision generation error in thread: {e}")
|
||||
finally:
|
||||
try:
|
||||
streamer.end()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
thread = threading.Thread(target=generate_fn)
|
||||
thread.start()
|
||||
|
||||
output = ""
|
||||
for new_token in streamer:
|
||||
if new_token:
|
||||
output += new_token
|
||||
cleaned = self._clean_generated_text(output)
|
||||
yield cleaned
|
||||
from queue import Empty
|
||||
try:
|
||||
while True:
|
||||
if cancel_event is not None and cancel_event.is_set():
|
||||
break
|
||||
try:
|
||||
new_token = next(streamer)
|
||||
except StopIteration:
|
||||
break
|
||||
except Empty:
|
||||
if not thread.is_alive():
|
||||
break
|
||||
continue
|
||||
if new_token:
|
||||
output += new_token
|
||||
cleaned = self._clean_generated_text(output)
|
||||
yield cleaned
|
||||
finally:
|
||||
if cancel_event is not None:
|
||||
cancel_event.set()
|
||||
thread.join(timeout=10)
|
||||
if thread.is_alive():
|
||||
logger.warning("Vision generation thread did not exit after cancel/join timeout")
|
||||
|
||||
thread.join()
|
||||
if err.get("msg"):
|
||||
yield f"Error: {err['msg']}"
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Vision generation error: {e}")
|
||||
|
|
@ -693,7 +729,8 @@ class InferenceBackend:
|
|||
top_p: float = 0.9,
|
||||
top_k: int = 40,
|
||||
max_new_tokens: int = 256,
|
||||
repetition_penalty: float = 1.1) -> Generator[str, None, None]:
|
||||
repetition_penalty: float = 1.1,
|
||||
cancel_event=None) -> Generator[str, None, None]:
|
||||
"""Generate streaming text response (text models only)."""
|
||||
if not self.active_model_name:
|
||||
yield "Error: No active model"
|
||||
|
|
@ -709,7 +746,12 @@ class InferenceBackend:
|
|||
from transformers import TextIteratorStreamer
|
||||
import threading
|
||||
|
||||
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)
|
||||
streamer = TextIteratorStreamer(
|
||||
tokenizer,
|
||||
skip_prompt=True,
|
||||
skip_special_tokens=True,
|
||||
timeout=0.2,
|
||||
)
|
||||
|
||||
generation_kwargs = dict(
|
||||
**inputs,
|
||||
|
|
@ -723,24 +765,66 @@ class InferenceBackend:
|
|||
eos_token_id=tokenizer.eos_token_id,
|
||||
pad_token_id=tokenizer.eos_token_id if tokenizer.pad_token_id is None else tokenizer.pad_token_id,
|
||||
)
|
||||
if cancel_event is not None:
|
||||
from transformers.generation.stopping_criteria import (
|
||||
StoppingCriteria,
|
||||
StoppingCriteriaList,
|
||||
)
|
||||
|
||||
class _CancelCriteria(StoppingCriteria):
|
||||
def __init__(self, ev):
|
||||
self.ev = ev
|
||||
|
||||
def __call__(self, input_ids, scores, **kwargs):
|
||||
return self.ev.is_set()
|
||||
|
||||
generation_kwargs["stopping_criteria"] = StoppingCriteriaList(
|
||||
[_CancelCriteria(cancel_event)]
|
||||
)
|
||||
|
||||
def generate_fn():
|
||||
try:
|
||||
model.generate(**generation_kwargs)
|
||||
except Exception as e:
|
||||
err["msg"] = str(e)
|
||||
logger.error(f"Generation error: {e}")
|
||||
finally:
|
||||
try:
|
||||
streamer.end()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
err: dict[str, str] = {}
|
||||
thread = threading.Thread(target=generate_fn)
|
||||
thread.start()
|
||||
|
||||
output = ""
|
||||
for new_token in streamer:
|
||||
if new_token:
|
||||
output += new_token
|
||||
cleaned = self._clean_generated_text(output)
|
||||
yield cleaned
|
||||
from queue import Empty
|
||||
try:
|
||||
while True:
|
||||
if cancel_event is not None and cancel_event.is_set():
|
||||
break
|
||||
try:
|
||||
new_token = next(streamer)
|
||||
except StopIteration:
|
||||
break
|
||||
except Empty:
|
||||
if not thread.is_alive():
|
||||
break
|
||||
continue
|
||||
if new_token:
|
||||
output += new_token
|
||||
cleaned = self._clean_generated_text(output)
|
||||
yield cleaned
|
||||
finally:
|
||||
if cancel_event is not None:
|
||||
cancel_event.set()
|
||||
thread.join(timeout=10)
|
||||
if thread.is_alive():
|
||||
logger.warning("Generation thread did not exit after cancel/join timeout")
|
||||
|
||||
thread.join()
|
||||
if err.get("msg"):
|
||||
yield f"Error: {err['msg']}"
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error during generation: {e}")
|
||||
|
|
|
|||
|
|
@ -5,11 +5,13 @@ import sys
|
|||
import time
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from fastapi.responses import StreamingResponse, JSONResponse
|
||||
from typing import Optional
|
||||
import json
|
||||
import logging
|
||||
import asyncio
|
||||
import threading
|
||||
|
||||
|
||||
|
||||
|
|
@ -304,7 +306,7 @@ def _extract_content_parts(
|
|||
|
||||
|
||||
@router.post("/chat/completions")
|
||||
async def openai_chat_completions(request: ChatCompletionRequest):
|
||||
async def openai_chat_completions(payload: ChatCompletionRequest, request: Request):
|
||||
"""
|
||||
OpenAI-compatible chat completions endpoint.
|
||||
|
||||
|
|
@ -324,7 +326,7 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
|
||||
# ── Parse messages (handles multimodal content parts) ─────
|
||||
system_prompt, chat_messages, extracted_image_b64 = _extract_content_parts(
|
||||
request.messages
|
||||
payload.messages
|
||||
)
|
||||
|
||||
# If no non-system messages were provided, error out
|
||||
|
|
@ -336,7 +338,7 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
|
||||
# ── Decode image (from content parts OR legacy field) ─────
|
||||
# Content-part images take priority; fall back to legacy field
|
||||
image_b64 = extracted_image_b64 or request.image_base64
|
||||
image_b64 = extracted_image_b64 or payload.image_base64
|
||||
image = None
|
||||
|
||||
if image_b64:
|
||||
|
|
@ -366,31 +368,35 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
messages=chat_messages,
|
||||
system_prompt=system_prompt,
|
||||
image=image,
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
top_k=request.top_k,
|
||||
max_new_tokens=request.max_tokens or 512,
|
||||
repetition_penalty=request.repetition_penalty,
|
||||
temperature=payload.temperature,
|
||||
top_p=payload.top_p,
|
||||
top_k=payload.top_k,
|
||||
max_new_tokens=payload.max_tokens or 512,
|
||||
repetition_penalty=payload.repetition_penalty,
|
||||
)
|
||||
|
||||
# ── Choose generation path (adapter-controlled or standard) ──
|
||||
if request.use_adapter is not None:
|
||||
cancel_event = threading.Event()
|
||||
|
||||
if payload.use_adapter is not None:
|
||||
# Compare mode: toggle adapter state atomically with generation
|
||||
def generate():
|
||||
return backend.generate_with_adapter_control(
|
||||
use_adapter=request.use_adapter, **gen_kwargs
|
||||
use_adapter=payload.use_adapter,
|
||||
cancel_event=cancel_event,
|
||||
**gen_kwargs,
|
||||
)
|
||||
else:
|
||||
# Standard path: no adapter toggling
|
||||
def generate():
|
||||
return backend.generate_chat_response(**gen_kwargs)
|
||||
return backend.generate_chat_response(cancel_event=cancel_event, **gen_kwargs)
|
||||
|
||||
model_name = backend.active_model_name or request.model
|
||||
model_name = backend.active_model_name or payload.model
|
||||
completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||
created = int(time.time())
|
||||
|
||||
# ── Streaming response ────────────────────────────────────────
|
||||
if request.stream:
|
||||
if payload.stream:
|
||||
async def stream_chunks():
|
||||
try:
|
||||
# First chunk: send the role
|
||||
|
|
@ -409,6 +415,10 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
# text, so we diff to get incremental deltas.
|
||||
prev_text = ""
|
||||
for cumulative in generate():
|
||||
if await request.is_disconnected():
|
||||
cancel_event.set()
|
||||
backend.reset_generation_state()
|
||||
return
|
||||
new_text = cumulative[len(prev_text):]
|
||||
prev_text = cumulative
|
||||
if not new_text:
|
||||
|
|
@ -437,6 +447,10 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
yield f"data: {final_chunk.model_dump_json(exclude_none=True)}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
except asyncio.CancelledError:
|
||||
cancel_event.set()
|
||||
backend.reset_generation_state()
|
||||
raise
|
||||
except Exception as e:
|
||||
backend.reset_generation_state()
|
||||
logger.error(f"Error during OpenAI streaming: {e}", exc_info=True)
|
||||
|
|
@ -477,4 +491,3 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
backend.reset_generation_state()
|
||||
logger.error(f"Error during OpenAI completion: {e}", exc_info=True)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue