File size: 7,126 Bytes
c255e8f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
# app.py
import os
import gradio as gr
from huggingface_hub import InferenceApi
from typing import List, Tuple

HF_TOKEN = os.environ.get("HUGGINGFACEHUB_API_TOKEN")
if not HF_TOKEN:
    raise RuntimeError("Set HUGGINGFACEHUB_API_TOKEN in your Space secrets/settings.")

# Default models (you can replace these with your own hosted model names)
TEXT_MODEL = os.environ.get("YAGA_TEXT_MODEL", "mistralai/Mistral-7B-v0.1")  # hosted generation model
IMAGE_CAPTION_MODEL = os.environ.get("YAGA_IMAGE_MODEL", "Salesforce/blip-image-captioning-large")
WHISPER_MODEL = os.environ.get("YAGA_WHISPER_MODEL", "openai/whisper-large-v2")

# Inference API clients
text_api = InferenceApi(repo_id=TEXT_MODEL, token=HF_TOKEN)
image_api = InferenceApi(repo_id=IMAGE_CAPTION_MODEL, token=HF_TOKEN)
whisper_api = InferenceApi(repo_id=WHISPER_MODEL, token=HF_TOKEN)

SYSTEM_PROMPT = (
    "You are YAGA — an advanced, helpful, and safe AI assistant. "
    "Answer concisely and clearly. If user provided an image or audio, "
    "use the provided descriptions/transcripts as context."
)

# Simple in-memory conversation store per session id (Gradio session)
# For Spaces this is ephemeral; for production, persist to DB.
CONV_STORE = {}

def build_prompt(history: List[Tuple[str, str]], extra_context: str | None):
    parts = [f"SYSTEM: {SYSTEM_PROMPT}"]
    if extra_context:
        parts.append(f"CONTEXT: {extra_context}")
    for role, text in history:
        parts.append(f"{role}: {text}")
    parts.append("Assistant:")
    return "\n".join(parts)

def get_session_history(session_id: str):
    return CONV_STORE.setdefault(session_id, [])

def clear_chat(session_id: str):
    CONV_STORE[session_id] = []
    return []

def transcribe_audio(audio_bytes):
    # Whisper inference: input is bytes or file-like; InferenceApi accepts bytes
    try:
        out = whisper_api(inputs=audio_bytes)
        # InferenceApi may return plain text or dict depending on model; handle common shapes
        if isinstance(out, dict) and "text" in out:
            return out["text"]
        if isinstance(out, str):
            return out
        # fallback: try join if list
        return str(out)
    except Exception as e:
        return f"[transcription_error: {e}]"

def caption_image(image_bytes):
    try:
        out = image_api(inputs=image_bytes)
        # Many image caption models return plain text
        if isinstance(out, dict) and "generated_text" in out:
            return out["generated_text"]
        if isinstance(out, str):
            return out
        # fallback
        return str(out)
    except Exception as e:
        return f"[caption_error: {e}]"

def generate_reply(prompt: str, max_tokens: int = 256, temperature: float = 0.7):
    params = {"max_new_tokens": max_tokens, "temperature": temperature, "top_p":0.95}
    # Some Inference API models return dicts or plain text; just return textual result
    out = text_api(inputs=prompt, parameters=params)
    if isinstance(out, dict):
        # common place: {"generated_text": "..."}
        if "generated_text" in out:
            return out["generated_text"]
        # some endpoints return list of candidates
        if "error" in out:
            return f"[generation_error: {out['error']}]"
        try:
            return str(out)
        except:
            return "[generation_error: unexpected format]"
    if isinstance(out, str):
        return out
    return str(out)

def respond(session_id, user_message, image_file, audio_file, max_tokens=256, temperature=0.7):
    history = get_session_history(session_id)

    extra_ctx_parts = []
    if image_file is not None:
        img_bytes = image_file.read()
        caption = caption_image(img_bytes)
        extra_ctx_parts.append(f"Image caption: {caption}")
    if audio_file is not None:
        audio_bytes = audio_file.read()
        transcript = transcribe_audio(audio_bytes)
        extra_ctx_parts.append(f"Audio transcript: {transcript}")

    extra_context = "\n".join(extra_ctx_parts) if extra_ctx_parts else None

    # Append user message to history
    history.append(("User", user_message))

    prompt = build_prompt(history, extra_context)
    reply = generate_reply(prompt, max_tokens=int(max_tokens), temperature=float(temperature))

    history.append(("Assistant", reply))
    # Keep last N turns
    CONV_STORE[session_id] = history[-12*2:]
    # Convert to chat-like pairs for the frontend
    chat_pairs = [(r[0], r[1]) for r in CONV_STORE[session_id]]
    return chat_pairs, ""

with gr.Blocks(title="YAGA — HF Space Prototype") as demo:
    gr.Markdown("## YAGA — Multi-turn chat (Hugging Face Space)\nText, image and audio inputs supported.\n\n**Important:** Add your Hugging Face token to Space secrets as `HUGGINGFACEHUB_API_TOKEN`.")
    with gr.Row():
        with gr.Column(scale=3):
            chatbot = gr.Chatbot()
            msg = gr.Textbox(placeholder="Type your message here...", label="Message")
            with gr.Row():
                img = gr.Image(type="bytes", label="Optional image")
                aud = gr.Audio(source="upload", type="file", label="Optional audio")
            with gr.Row():
                max_t = gr.Slider(32, 1024, value=256, step=32, label="Max new tokens")
                temp = gr.Slider(0.0, 1.0, value=0.7, step=0.05, label="Temperature")
            with gr.Row():
                send = gr.Button("Send")
                clear = gr.Button("Clear chat")
        with gr.Column(scale=1):
            gr.Markdown("### Session Info")
            session_id_display = gr.Textbox(value="", interactive=False, label="Session ID (auto)")
            gr.Markdown("### Hints")
            gr.Markdown("- Upload an image to give the model visual context (it will be captioned).\n- Upload audio to transcribe via Whisper.\n- Replace `TEXT_MODEL` env var with your model repo for a custom model.")
    # use gradio session_state for session id
    def get_sid():
        import uuid
        return str(uuid.uuid4())
    sid = get_sid()
    session_id_display.value = sid

    def on_send(message, image_file, audio_file, max_tokens, temperature):
        pairs, _ = respond(sid, message, image_file, audio_file, max_tokens, temperature)
        # convert pairs to Chatbot messages: list of [user, assistant] pairs
        chat_display = []
        it = iter(pairs)
        # pairs is list of (role, text)
        # Build sequential chat list
        for role, text in pairs:
            # only show in alternating style: we will append tuple (role, text)
            chat_display.append((role, text))
        # For gradio Chatbot, we need list of [user, assistant] pairs grouped.
        # But we'll just return the last assistant message to append:
        last_user = pairs[-2][1] if len(pairs) >= 2 else ""
        last_assistant = pairs[-1][1]
        return gr.update(value=chat_display), ""

    send.click(on_send, [msg, img, aud, max_t, temp], [chatbot, msg])
    def on_clear():
        clear_chat(sid)
        return [], ""
    clear.click(on_clear, [], [chatbot, msg])

if __name__ == "__main__":
    demo.launch()