import base64 import json import time from typing import List, Optional from pydantic import BaseModel import httpx from fastapi import FastAPI, Request, HTTPException, Form, UploadFile, File from fastapi.responses import HTMLResponse, StreamingResponse from fastapi.middleware.cors import CORSMiddleware import io import threading from threading import Thread from PIL import Image import torch try: from transformers import AutoProcessor, AutoModelForImageTextToText, TextIteratorStreamer except Exception as e: print("FATAL ERROR: Failed to import transformers classes.") import traceback traceback.print_exc() raise e app = FastAPI( title="Gemma 4 E2B API Space", description="FastAPI wrapper for Gemma 4 E2B via native HF Transformers.", version="1.0.0" ) app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) MODEL_NAME = "google/gemma-4-E2B" # Global references for model and processor processor = None model = None model_loaded = False load_error = None loading_lock = threading.Lock() def load_model_if_needed(): global processor, model, model_loaded, load_error with loading_lock: if model_loaded: return try: print(f"Loading processor and model: {MODEL_NAME}...") processor = AutoProcessor.from_pretrained(MODEL_NAME) # Setup fallback chat template if the official template file is missing/uncached tokenizer = None if hasattr(processor, "tokenizer"): tokenizer = processor.tokenizer elif hasattr(processor, "image_processor") and hasattr(processor, "tokenizer"): tokenizer = processor.tokenizer has_template = False if hasattr(processor, "chat_template") and processor.chat_template: has_template = True elif tokenizer and hasattr(tokenizer, "chat_template") and tokenizer.chat_template: has_template = True if not has_template: print("Setting fallback chat template for Gemma...") fallback_template = ( "{{ bos_token }}" "{% for message in messages %}" "{{ message['role'] }}\n" "{% if message['content'] is string %}" "{{ message['content'] }}" "{% else %}" "{% for part in message['content'] %}" "{% if part['type'] == 'text' %}" "{{ part['text'] }}" "{% elif part['type'] == 'image' %}" "" "{% endif %}" "{% endfor %}" "{% endif %}" "\n" "{% endfor %}" "{% if add_generation_prompt %}" "model\n" "{% endif %}" ) try: processor.chat_template = fallback_template except Exception: pass if tokenizer: tokenizer.chat_template = fallback_template try: print(f"Attempting to load model in bfloat16: {MODEL_NAME}...") model = AutoModelForImageTextToText.from_pretrained( MODEL_NAME, device_map="cpu", torch_dtype=torch.bfloat16, low_cpu_mem_usage=True ) print("Model loaded in bfloat16 successfully!") except Exception as e_bf16: print(f"bfloat16 load failed ({e_bf16}). Falling back to float32...") model = AutoModelForImageTextToText.from_pretrained( MODEL_NAME, device_map="cpu", torch_dtype=torch.float32, low_cpu_mem_usage=True ) print("Model loaded in float32 successfully!") model_loaded = True print("Model loaded successfully!") except Exception as e: load_error = str(e) print(f"Error loading model: {load_error}") raise e def background_load_model(): try: load_model_if_needed() except Exception: pass @app.on_event("startup") async def startup_event(): # Start loading the model in a background thread to prevent HF Spaces boot timeout thread = threading.Thread(target=background_load_model, daemon=True) thread.start() class GenerateRequest(BaseModel): prompt: str system: Optional[str] = None temperature: Optional[float] = 0.7 class ChatMessage(BaseModel): role: str content: str class ChatRequest(BaseModel): messages: List[ChatMessage] temperature: Optional[float] = 0.7 # --------------------------------------------------------------------------- # HTML UI # --------------------------------------------------------------------------- HTML_PAGE = """ Gemma 4 E2B — Live Console
Gemma 4 E2B gemma4:e2b CPU · Ollama · FastAPI
📋 Request / Response Log
""" # --------------------------------------------------------------------------- # Routes # --------------------------------------------------------------------------- # --------------------------------------------------------------------------- # Helper function to extract PIL Images and filter messages list # --------------------------------------------------------------------------- def extract_images_and_clean_messages(messages): images = [] cleaned_messages = [] for msg in messages: role = msg.get("role") content = msg.get("content") if isinstance(content, list): cleaned_parts = [] for part in content: if part.get("type") == "image": img = part.get("image") if img: images.append(img) # We keep the dict structure but remove raw PIL objects for serialization, # or keep it if processor needs it. Standard transformers template # accepts type="image" and ignores or expects it. cleaned_parts.append(part) else: cleaned_parts.append(part) cleaned_messages.append({"role": role, "content": cleaned_parts}) else: cleaned_messages.append({"role": role, "content": content}) return cleaned_messages, images # --------------------------------------------------------------------------- # Helper function for streaming generation # --------------------------------------------------------------------------- async def run_generation(messages): global model_loaded, load_error if not model_loaded: if load_error: yield f"data: [ERROR] Model load failed: {load_error}\n\n" else: yield "data: [ERROR] Model is still loading. Please try again in a few seconds.\n\n" return try: # Extract images and clean messages cleaned_messages, images = extract_images_and_clean_messages(messages) # Apply chat template (yields text prompt) text = processor.apply_chat_template( cleaned_messages, tokenize=False, add_generation_prompt=True ) # Build processor inputs if images: inputs = processor(text=text, images=images, return_tensors="pt") else: inputs = processor(text=text, return_tensors="pt") # Move to CPU / model device inputs = {k: v.to(model.device) for k, v in inputs.items()} streamer = TextIteratorStreamer( processor.tokenizer, skip_prompt=True, skip_special_tokens=True ) generate_kwargs = dict( **inputs, streamer=streamer, max_new_tokens=1024, do_sample=True, temperature=0.7 ) # Run model generate in a daemon thread so it does not block the event loop thread = Thread(target=model.generate, kwargs=generate_kwargs, daemon=True) thread.start() for new_text in streamer: if new_text: text_escaped = new_text.replace("\n", "\\n") yield f"data: {text_escaped}\n\n" except Exception as e: yield f"data: [ERROR] Generation error: {str(e)}\n\n" def run_generation_sync(messages, temperature=0.7): if not model_loaded: if load_error: raise HTTPException(status_code=503, detail=f"Model load failed: {load_error}") raise HTTPException(status_code=503, detail="Model is still loading. Please try again.") cleaned_messages, images = extract_images_and_clean_messages(messages) text = processor.apply_chat_template( cleaned_messages, tokenize=False, add_generation_prompt=True ) if images: inputs = processor(text=text, images=images, return_tensors="pt") else: inputs = processor(text=text, return_tensors="pt") inputs = {k: v.to(model.device) for k, v in inputs.items()} with torch.no_grad(): outputs = model.generate( **inputs, max_new_tokens=1024, do_sample=True, temperature=temperature ) input_len = inputs["input_ids"].shape[1] generated_tokens = outputs[0][input_len:] return processor.decode(generated_tokens, skip_special_tokens=True) # --------------------------------------------------------------------------- # Routes # --------------------------------------------------------------------------- @app.get("/") def root(request: Request): accept = request.headers.get("accept", "") if "text/html" in accept: return HTMLResponse(content=HTML_PAGE) return { "status": "online", "model": MODEL_NAME, "message": "Gemma 4 E2B API is running natively.", "endpoints": { "chat_form": "/chat", "vision": "/vision", "generate": "/api/generate", "chat_json": "/api/chat", "health": "/health", "docs": "/docs" } } @app.get("/health") async def health_check(): global model_loaded, load_error if model_loaded: return {"status": "ok", "model": MODEL_NAME, "loaded": True} elif load_error: raise HTTPException(status_code=503, detail=f"Model failed to load: {load_error}") else: return {"status": "loading", "model": MODEL_NAME, "loaded": False} # --------------------------------------------------------------------------- # POST /chat – FormData → SSE # --------------------------------------------------------------------------- @app.post("/chat") async def chat_form(prompt: str = Form(...)): messages = [ { "role": "user", "content": [{"type": "text", "text": prompt}] } ] return StreamingResponse( run_generation(messages), media_type="text/event-stream", headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"} ) # --------------------------------------------------------------------------- # POST /vision – FormData (prompt + imagen) → SSE # --------------------------------------------------------------------------- @app.post("/vision") async def vision_form(prompt: str = Form(...), imagen: UploadFile = File(...)): try: image_bytes = await imagen.read() pil_image = Image.open(io.BytesIO(image_bytes)).convert("RGB") except Exception as e: raise HTTPException(status_code=400, detail=f"Invalid image file: {str(e)}") messages = [ { "role": "user", "content": [ {"type": "image", "image": pil_image}, {"type": "text", "text": prompt} ] } ] return StreamingResponse( run_generation(messages), media_type="text/event-stream", headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"} ) # --------------------------------------------------------------------------- # JSON endpoints # --------------------------------------------------------------------------- @app.post("/api/generate") async def generate_completion(payload: GenerateRequest): messages = [ { "role": "user", "content": [{"type": "text", "text": payload.prompt}] } ] if payload.system: messages.insert(0, { "role": "system", "content": [{"type": "text", "text": payload.system}] }) try: response_text = run_generation_sync(messages, temperature=payload.temperature) return { "model": MODEL_NAME, "response": response_text, "done": True } except HTTPException as he: raise he except Exception as e: raise HTTPException(status_code=500, detail=f"Generation failed: {str(e)}") @app.post("/api/chat") async def chat_completion(payload: ChatRequest): messages = [] for msg in payload.messages: messages.append({ "role": msg.role, "content": [{"type": "text", "text": msg.content}] }) try: response_text = run_generation_sync(messages, temperature=payload.temperature) return { "model": MODEL_NAME, "message": { "role": "assistant", "content": response_text }, "done": True } except HTTPException as he: raise he except Exception as e: raise HTTPException(status_code=503, detail=f"Generation failed: {str(e)}")