feat: add cancelation support for chat generation and streaming tasks

This commit is contained in:
Shine1i 2026-02-15 18:23:27 +01:00
commit 571959e383
2 changed files with 135 additions and 38 deletions

View file

@ -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}")

View file

@ -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))