Spaces:
Running on Zero
Running on Zero
File size: 5,213 Bytes
b1aba72 51e9502 b1aba72 51e9502 b1aba72 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 | """
generation.py
-----------------------
Answer generation — ZeroGPU version.
"""
import spaces
from config import MAX_PROMPT_TOKENS, ZEROGPU_DURATION
from model_setup import text_gen_pipeline, tokenizer
from utils import format_prompt, clean_final_answer
from logging_config import get_logger
logger = get_logger(__name__)
# The ONLY GPU-touching function in the codebase
@spaces.GPU(duration=ZEROGPU_DURATION)
def _generate_on_gpu(prompt: str) -> str:
"""
Run the actual LLM inference. This is the single point in the
entire application where a real GPU is attached to the process.
Deliberately minimal: no logging of prompt content here (keep
the GPU-held window as short as possible), no retries, no
fallback logic — just the inference call. All the surrounding
orchestration (prompt building, history handling, streaming)
happens in generate_answer() below, OUTSIDE the GPU-decorated
boundary, where a real GPU is not needed and thus not consumed.
Returns the complete generated text (already stripped).
"""
output = text_gen_pipeline(prompt)
return output[0]["generated_text"].strip()
# Prompt building + streaming wrapper (no GPU here)
def generate_answer(state: dict):
"""
Generator that yields (token_text, prompt_token_count) tuples.
This function itself is NOT @spaces.GPU-decorated — it does no CUDA work directly.
It builds the prompt (CPU-side string operations), calls _generate_on_gpu() once to get the complete
response, then splits it into words and yields them one by one for the streaming UI effect in `app.py`.
Splitting the concerns this way means the GPU slot is held for the minimum possible time —
just the actual generation call — while all the prompt-assembly and history-formatting logic runs
on CPU without competing for the shared GPU resource.
"""
from ui_helpers import get_history_text, _safe_content
query = state["query"]
history = state.get("history", [])
retrieved_docs = _safe_content(state.get("retrieved_docs", ""))
web_results = _safe_content(state.get("web_results", ""))
# Safety valve: both retrieval phases returned nothing
if not retrieved_docs and not web_results:
yield (
"I could not retrieve relevant information for your question. "
"Please check that the TAVILY_API_KEY secret is set in the "
"Space settings, or try rephrasing your question.",
0,
)
return
# Build context block
context_parts = []
if retrieved_docs:
context_parts.append(f"## Internal Knowledge Base\n\n{retrieved_docs}")
if web_results:
context_parts.append(f"## Web Search Results\n\n{web_results}")
context = "\n\n".join(context_parts)
# System prompt
system_msg = (
"You are a precise, factual research assistant.\n\n"
"Rules:\n"
"1. Answer ONLY using the context and conversation history provided.\n"
"2. If the context is insufficient, say so clearly.\n"
"3. Cite factual claims with markdown links: [Title](URL).\n"
" Use ONLY URLs that appear explicitly in the context.\n"
"4. Be concise: 2-4 paragraphs.\n"
"5. Do NOT reproduce context verbatim. Synthesise and summarise.\n"
"6. For follow-up questions, use the conversation history to\n"
" understand what the user is referring to."
)
# User message: history (if any) + context + question
history_text = get_history_text(history, max_turns=3)
if history_text:
user_msg = (
f"Previous conversation:\n{history_text}\n\n"
f"Context:\n{context}\n\n"
f"---\n\n"
f"Current question: {query}\n\nAnswer:"
)
else:
user_msg = f"Context:\n{context}\n\n---\n\nQuestion: {query}\n\nAnswer:"
# format_prompt() uses tokenizer.apply_chat_template — CPU-side,
# no CUDA involved, safe to call outside the GPU-decorated function.
prompt = format_prompt(system=system_msg, user=user_msg)
# Token budget: trim from the LEFT if over limit (tokenizer.encode
# is also CPU-side — no GPU needed for this step).
token_ids = tokenizer.encode(prompt)
token_count = len(token_ids)
if token_count > MAX_PROMPT_TOKENS:
token_ids = token_ids[-MAX_PROMPT_TOKENS:]
token_count = len(token_ids) # MAX_PROMPT_TOKENS
prompt = tokenizer.decode(token_ids, skip_special_tokens=False)
logger.info(f"Prompt: {token_count} tokens | history turns: {len(history)}")
# --- The ONLY call into GPU-decorated code -----------------
# This is where the ZeroGPU backend actually attaches a real GPU
# to the process, runs inference, and releases the slot when done.
full_text = _generate_on_gpu(prompt)
# Post-process: fix Qwen's URL formatting quirks
full_text = clean_final_answer(full_text)
# --- Word-by-word yield (CPU-side, instant) -----------------
# Simulates streaming by splitting the already-complete response.
# This is identical to the CPU version's approach — the only
# difference is that _generate_on_gpu() ran on a real GPU instead
# of the CPU-bound text_gen_pipeline() call.
words = full_text.split()
for i, word in enumerate(words):
yield word + (" " if i < len(words) - 1 else ""), token_count
|