PraneshJs/FastVLM_Caption_Generator
0
1import gradio as gr2import torch3from PIL import Image4from transformers import AutoTokenizer, AutoModelForCausalLM5 6# ---- CPU-only config ----7MID = "apple/FastVLM-0.5B"8IMAGE_TOKEN_INDEX = -200 # special image token id used by FastVLM9 10tok = None11model = None12 13def load_model():14 global tok, model15 if tok is None or model is None:16 print("Loading model (CPU)…")17 tok = AutoTokenizer.from_pretrained(MID, trust_remote_code=True)18 # Force CPU + float32 (fp16 is unsafe on CPU)19 model = AutoModelForCausalLM.from_pretrained(20 MID,21 torch_dtype=torch.float32,22 device_map="cpu",23 trust_remote_code=True,24 )25 print("Model loaded successfully on CPU!")26 return tok, model27 28def caption_image(image, custom_prompt=None):29 """30 Generate a caption for the input image (CPU-only).31 """32 if image is None:33 return "Please upload an image first."34 35 try:36 tok, model = load_model()37 38 # Convert image to RGB if needed39 if image.mode != "RGB":40 image = image.convert("RGB")41 42 prompt = custom_prompt if custom_prompt else "Describe this image in detail."43 44 # Single-turn chat with an <image> placeholder45 messages = [{"role": "user", "content": f"<image>\n{prompt}"}]46 rendered = tok.apply_chat_template(messages, add_generation_prompt=True, tokenize=False)47 48 # Split around the literal "<image>"49 pre, post = rendered.split("<image>", 1)50 51 # Tokenize text around the image token52 pre_ids = tok(pre, return_tensors="pt", add_special_tokens=False).input_ids53 post_ids = tok(post, return_tensors="pt", add_special_tokens=False).input_ids54 55 # Derive device/dtype from the loaded model (CPU here, but future-proof)56 model_device = next(model.parameters()).device57 model_dtype = next(model.parameters()).dtype58 59 # Insert IMAGE token id at placeholder position60 img_tok = torch.tensor([[IMAGE_TOKEN_INDEX]], dtype=pre_ids.dtype, device=model_device)61 input_ids = torch.cat(62 [pre_ids.to(model_device), img_tok, post_ids.to(model_device)],63 dim=164 )65 attention_mask = torch.ones_like(input_ids, device=model_device)66 67 # Preprocess image using model's vision tower68 px = model.get_vision_tower().image_processor(69 images=image, return_tensors="pt"70 )["pixel_values"].to(device=model_device, dtype=model_dtype)71 72 # Generate caption (deterministic)73 with torch.no_grad():74 out = model.generate(75 inputs=input_ids,76 attention_mask=attention_mask,77 images=px,78 max_new_tokens=400,79 do_sample=False, # temperature is ignored when sampling is off80 )81 82 # Decode and slice to the assistant part if present83 generated_text = tok.decode(out[0], skip_special_tokens=True)84 if "Assistant:" in generated_text:85 response = generated_text.split("Assistant:", 1)[-1].strip()86 elif "assistant" in generated_text:87 response = generated_text.split("assistant", 1)[-1].strip()88 else:89 response = generated_text.strip()90 91 return response92 93 except Exception as e:94 return f"Error generating caption: {str(e)}"95 96# ---- Gradio UI (CPU) ----97with gr.Blocks(title="FastVLM Image Captioning (CPU)") as demo:98 gr.Markdown(99 """100 # 🖼️ FastVLM Image Captioning (CPU)101 Upload an image to generate a detailed caption using Apple's FastVLM-0.5B.102 This build runs on **CPU only**. Expect slower generation than GPU.103 """104 )105 106 with gr.Row():107 with gr.Column():108 image_input = gr.Image(type="pil", label="Upload Image", elem_id="image-upload")109 custom_prompt = gr.Textbox(110 label="Custom Prompt (Optional)",111 placeholder="Leave empty for default: 'Describe this image in detail.'",112 lines=2113 )114 with gr.Row():115 clear_btn = gr.ClearButton([image_input, custom_prompt])116 generate_btn = gr.Button("Generate Caption", variant="primary")117 118 with gr.Column():119 output = gr.Textbox(120 label="Generated Caption",121 lines=8,122 max_lines=15,123 show_copy_button=True124 )125 126 generate_btn.click(fn=caption_image, inputs=[image_input, custom_prompt], outputs=output)127 128 # Also generate on image upload if no custom prompt129 def _auto_caption(img, prompt):130 return caption_image(img, prompt) if (img is not None and not prompt) else None131 132 image_input.change(fn=_auto_caption, inputs=[image_input, custom_prompt], outputs=output)133 134 gr.Markdown(135 """136 ---137 **Model:** [apple/FastVLM-0.5B](https://huggingface.co/apple/FastVLM-0.5B) 138 **Note:** CPU-only run. For speed, switch to a CUDA GPU build or a GPU Space.139 """140 )141 142if __name__ == "__main__":143 demo.launch(144 share=False,145 show_error=True,146 server_name="0.0.0.0",147 server_port=7860148 )