TARN / app.py
Phase-Technologies's picture
Update app.py
c376ada verified
Raw
History Blame Contribute Delete
4.97 kB
import os
import gc
import torch
import gradio as ui
import spaces # ADDED: Hugging Face Spaces SDK for ZeroGPU support
from transformers import pipeline, TextIteratorStreamer
from threading import Thread
# 1. Initialize Pipeline
MODEL_ID = "Xerv-AI/tarn"
print("Loading tarn architecture into memory...")
# FIX: Replaced 'torch_dtype' with 'dtype' to clear the transformers deprecation warning.
pipe = pipeline(
"image-text-to-text",
model=MODEL_ID,
model_kwargs={
"dtype": torch.float16,
"device_map": "auto"
}
)
print("tarn is fully initialized.")
# 2. Define the Inference Logic
@spaces.GPU # ADDED: Required by HF Spaces to allocate ZeroGPU dynamically
def process_chat(message, history):
"""
Handles incoming messages (both images and text), formats them into
the proper Qwen 3.5 VL structure, and yields streamed tokens.
"""
# Force clean any leftover GPU allocation trash before processing
gc.collect()
torch.cuda.empty_cache()
# Reconstruct history into Hugging Face chat template format
formatted_messages = []
# Process previous conversation turns
for user_turn, assistant_turn in history:
if user_turn:
# Check if user input in history was a file/image dict or text
if isinstance(user_turn, dict) or (isinstance(user_turn, tuple) and len(user_turn) == 1):
img_path = user_turn[0] if isinstance(user_turn, tuple) else user_turn.get("path")
formatted_messages.append({"role": "user", "content": [{"type": "image", "url": img_path}]})
else:
formatted_messages.append({"role": "user", "content": [{"type": "text", "text": str(user_turn)}]})
if assistant_turn:
formatted_messages.append({"role": "assistant", "content": [{"type": "text", "text": str(assistant_turn)}]})
# Append current incoming turn
current_content = []
# Check if the user uploaded an image along with text
if message["files"]:
for file_info in message["files"]:
# Gradio 4+ passes files as dicts or named objects
file_path = file_info if isinstance(file_info, str) else file_info.get("path")
current_content.append({"type": "image", "url": file_path})
if message["text"]:
current_content.append({"type": "text", "text": message["text"]})
# Guard clause if user sent nothing
if not current_content:
yield ""
return
formatted_messages.append({"role": "user", "content": current_content})
# Initialize a non-blocking compilation streamer
streamer = TextIteratorStreamer(pipe.tokenizer, skip_prompt=True, skip_special_tokens=True)
# Inject VRAM patch configurations directly into generator parameters
generate_kwargs = {
"text": formatted_messages,
"max_new_tokens": 1024,
"min_pixels": 256 * 28 * 28,
"max_pixels": 512 * 28 * 28,
"generate_kwargs": {"streamer": streamer}
}
# Execute text generation on a separate background thread to keep UI interactive
thread = Thread(target=pipe, kwargs=generate_kwargs)
thread.start()
# Yield generated tokens back to UI as they arrive
partial_text = ""
for new_token in streamer:
partial_text += new_token
yield partial_text
# 3. Build Web Layout using Custom Calm Styling Themes
custom_theme = ui.themes.Soft(
primary_hue="orange", # Brings in that subtle "tangy" vibe softly
neutral_hue="slate", # Keeps the interface looking professional and clean
font=[ui.themes.GoogleFont("Source Sans Pro"), "Arial", "sans-serif"]
)
# FIX: Removed `theme=custom_theme` from Blocks (moved to launch).
with ui.Blocks(title="tarn Demo Space") as demo:
ui.Markdown(
"""
# 🌌 tarn Multimodal Reasoner
### Fine-tuned for deep visual context analysis and structured step-by-step logic.
*Built and maintained by **Xerv-AI**.*
"""
)
# Core multi-turn Gradio Chatbot interface
chatbot_interface = ui.ChatInterface(
fn=process_chat,
chatbot=ui.Chatbot(
height=550,
placeholder="📸 Upload an image or ask a structural reasoning question...",
# FIX: Removed `show_copy_button=True` as it is removed in newer Gradio versions.
avatar_images=(None, "https://huggingface.co/front/assets/huggingface/logos/huggingface-logo-dark.svg")
),
multimodal=True, # Explicitly turns on the combined image upload + text input array
textbox=ui.MultimodalTextbox(
placeholder="Type your question here... (Click the paperclip icon to pin an image)",
file_types=["image"],
scale=7
)
)
if __name__ == "__main__":
# FIX: Passed the theme to the launch method to satisfy Gradio 6.0 requirements.
demo.queue().launch(theme=custom_theme)