""" Gemma 4 Any-to-Any — Flask API Architecture: TARGET : google/gemma-4-e2b-it — full Gemma 4 E2B multimodal model (text + audio + image, runs standalone) DRAFTER : gsstec322/gemma — Gemma4Assistant MTP drafter (passed as assistant_model for speculative decoding — up to 3x speedup, same quality) The drafter CANNOT run standalone: it requires inputs_embeds + shared_kv_states from the target model and must only be used as assistant_model in generate(). Endpoints ───────── GET /health → liveness probe POST /audio-to-text → audio file → transcription / translation POST /text-to-audio → text prompt → Gemma reply + WAV audio (base64) """ import base64 import io import os import re import tempfile import textwrap import traceback import numpy as np import soundfile as sf import torch from flask import Flask, jsonify, request from transformers import AutoModelForCausalLM, AutoProcessor # ── Constants ────────────────────────────────────────────────────────────────── TARGET_MODEL_ID = "google/gemma-4-e2b-it" # full model — all inference goes here ASSISTANT_MODEL_ID = "gsstec322/gemma" # MTP drafter — accelerates generation DEVICE = "cuda" if torch.cuda.is_available() else "cpu" DTYPE = torch.bfloat16 # ── Model loading (once at startup) ─────────────────────────────────────────── print(f"[init] Loading processor from {TARGET_MODEL_ID} …") processor = AutoProcessor.from_pretrained(TARGET_MODEL_ID) print(f"[init] Loading target model {TARGET_MODEL_ID} on {DEVICE} ({DTYPE}) …") model = AutoModelForCausalLM.from_pretrained( TARGET_MODEL_ID, dtype=DTYPE, device_map="auto", ) model.eval() print("[init] Target model ready ✓") print(f"[init] Loading assistant/drafter {ASSISTANT_MODEL_ID} …") assistant_model = AutoModelForCausalLM.from_pretrained( ASSISTANT_MODEL_ID, dtype=DTYPE, device_map="auto", ) assistant_model.eval() print("[init] Assistant model ready ✓") # ── Flask app ────────────────────────────────────────────────────────────────── app = Flask(__name__) # ── Internal helpers ────────────────────────────────────────────────────────── def _apply_template(messages: list) -> str: """Apply the processor's chat template, with a manual Gemma 4 fallback.""" try: return processor.apply_chat_template( messages, tokenize=False, add_generation_prompt=True ) except Exception: # tokenizer.chat_template not set on this checkpoint — render manually text = "" for msg in messages: role = msg["role"] content = msg["content"] if isinstance(msg["content"], str) else "" text += f"<|turn>{role}\n{content}" text += "<|turn>model\n" return text def _parse_model_response(raw: str) -> str: """Unwrap processor.parse_response to a plain string.""" try: parsed = processor.parse_response(raw) if isinstance(parsed, dict): return parsed.get("text", raw) return str(parsed) except Exception: return re.sub(r"<[^>]+>", "", raw).strip() def _run_text_only(messages: list, max_new_tokens: int = 512) -> str: """Generate from text-only messages using speculative decoding.""" text = _apply_template(messages) inputs = processor(text=text, return_tensors="pt").to(DEVICE) input_len = inputs["input_ids"].shape[-1] with torch.inference_mode(): outputs = model.generate( **inputs, assistant_model=assistant_model, max_new_tokens=max_new_tokens, do_sample=True, temperature=1.0, top_p=0.95, top_k=64, ) raw = processor.decode(outputs[0][input_len:], skip_special_tokens=False) return _parse_model_response(raw) def _run_with_audio( messages: list, audio_arrays: list, sample_rates: list, max_new_tokens: int = 512 ) -> str: """Generate from messages that include audio, using speculative decoding.""" text = _apply_template(messages) audios = [ {"array": arr, "sampling_rate": sr} for arr, sr in zip(audio_arrays, sample_rates) ] inputs = processor( text=text, audios=audios if audios else None, return_tensors="pt", ).to(DEVICE) input_len = inputs["input_ids"].shape[-1] with torch.inference_mode(): outputs = model.generate( **inputs, assistant_model=assistant_model, max_new_tokens=max_new_tokens, do_sample=True, temperature=1.0, top_p=0.95, top_k=64, ) raw = processor.decode(outputs[0][input_len:], skip_special_tokens=False) return _parse_model_response(raw) def _normalise_waveform(waveform: np.ndarray) -> np.ndarray: """Convert any integer PCM to float32 in [-1, 1].""" waveform = waveform.astype(np.float32) peak = np.max(np.abs(waveform)) if peak > 1.0: waveform /= peak return waveform def _tts(text: str) -> tuple: """ Best-effort TTS: gTTS (online) → pyttsx3 (offline) → silence. Returns (sample_rate, float32 waveform). """ # 1. gTTS try: from gtts import gTTS tts = gTTS(text=text, lang="en") buf = io.BytesIO() tts.write_to_fp(buf) buf.seek(0) waveform, sr = sf.read(buf, dtype="float32") return sr, waveform except Exception: pass # 2. pyttsx3 try: import pyttsx3 engine = pyttsx3.init() tmp = tempfile.mktemp(suffix=".wav") engine.save_to_file(text, tmp) engine.runAndWait() waveform, sr = sf.read(tmp, dtype="float32") os.remove(tmp) return sr, waveform except Exception: pass # 3. Silent fallback sr = 22050 return sr, np.zeros(sr, dtype=np.float32) def _waveform_to_base64_wav(sr: int, waveform: np.ndarray) -> str: """Encode a waveform as a base64 WAV string.""" buf = io.BytesIO() sf.write(buf, waveform, sr, format="WAV", subtype="PCM_16") buf.seek(0) return base64.b64encode(buf.read()).decode("utf-8") # ── Routes ──────────────────────────────────────────────────────────────────── @app.get("/") def index(): """API index.""" return jsonify({ "target_model": TARGET_MODEL_ID, "assistant_model": ASSISTANT_MODEL_ID, "device": DEVICE, "endpoints": { "GET /health": "Liveness probe", "POST /audio-to-text": "Upload audio file → transcription or translation (multipart/form-data)", "POST /text-to-audio": "Send text prompt → Gemma reply + base64 WAV (application/json)", }, }) @app.get("/health") def health(): """Liveness probe.""" return jsonify({ "status": "ok", "target_model": TARGET_MODEL_ID, "assistant_model": ASSISTANT_MODEL_ID, "device": DEVICE, }) @app.post("/audio-to-text") def audio_to_text(): """ Convert uploaded audio to text (transcription or translation). Multipart form fields ───────────────────── audio : audio file (wav, mp3, flac, ogg — max 30 s) task : "transcribe" | "translate" (default: "transcribe") source_language: e.g. "English" (default: "English") target_language: e.g. "French" (default: "English", translate only) max_new_tokens : int (default: 512) Response JSON → { "text": "..." } """ try: if "audio" not in request.files: return jsonify({"error": "No audio file. Send under the 'audio' key."}), 400 audio_file = request.files["audio"] task = request.form.get("task", "transcribe").lower() source_language = request.form.get("source_language", "English") target_language = request.form.get("target_language", "English") max_new_tokens = int(request.form.get("max_new_tokens", 512)) try: buf = io.BytesIO(audio_file.read()) waveform, sample_rate = sf.read(buf, dtype="float32") except Exception as exc: return jsonify({"error": f"Could not read audio file: {exc}"}), 422 waveform = _normalise_waveform(waveform) if task == "translate": instruction = textwrap.dedent(f"""\ Transcribe the following speech segment in {source_language}, \ then translate it into {target_language}. When formatting the answer, first output the transcription in \ {source_language}, then one newline, then output '{target_language}: ', \ then the translation.""") else: instruction = textwrap.dedent(f"""\ Transcribe the following speech segment in {source_language}. * Only output the transcription, with no newlines. * Write numbers as digits (e.g. 3, not three).""") messages = [ { "role": "user", "content": [ {"type": "text", "text": instruction}, {"type": "audio", "array": waveform, "sampling_rate": sample_rate}, ], } ] result = _run_with_audio( messages, audio_arrays=[waveform], sample_rates=[sample_rate], max_new_tokens=max_new_tokens, ) return jsonify({"text": result}) except Exception as exc: return jsonify({"error": str(exc), "trace": traceback.format_exc()}), 500 @app.post("/text-to-audio") def text_to_audio(): """ Generate a Gemma 4 reply from a text prompt and return it as spoken audio. JSON body ───────── { "text" : "What is the speed of light?", ← required "system_prompt" : "You are a helpful assistant.", ← optional "max_new_tokens": 256 ← optional (default 256) } Response JSON ───────────── { "reply_text" : "The speed of light is …", "audio_wav" : "", "sample_rate": 22050 } """ try: body = request.get_json(force=True, silent=True) or {} user_text = body.get("text", "").strip() if not user_text: return jsonify({"error": "'text' field is required and must not be empty."}), 400 system_prompt = body.get("system_prompt", "You are a helpful assistant.") max_new_tokens = int(body.get("max_new_tokens", 256)) # Prepend system prompt into the user turn (avoids relying on a # "system" role the tokenizer template may not support). merged = f"{system_prompt}\n\n{user_text}" if system_prompt else user_text messages = [{"role": "user", "content": merged}] reply_text = _run_text_only(messages, max_new_tokens=max_new_tokens) try: sr, waveform = _tts(reply_text) audio_b64 = _waveform_to_base64_wav(sr, waveform) except Exception as tts_exc: return jsonify({ "reply_text" : reply_text, "audio_wav" : None, "sample_rate": None, "tts_error" : str(tts_exc), }) return jsonify({ "reply_text" : reply_text, "audio_wav" : audio_b64, "sample_rate": sr, }) except Exception as exc: return jsonify({"error": str(exc), "trace": traceback.format_exc()}), 500 # ── Entry point ─────────────────────────────────────────────────────────────── if __name__ == "__main__": app.run(host="0.0.0.0", port=7860, debug=False)