diff --git a/studio/backend/core/inference/inference.py b/studio/backend/core/inference/inference.py index 23203613f2..6be7e63306 100644 --- a/studio/backend/core/inference/inference.py +++ b/studio/backend/core/inference/inference.py @@ -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}") diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index ae8fc31a46..3bb84d5c6b 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -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)) -