Spaces:
Running on Zero
Running on Zero
| import os | |
| os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") | |
| import spaces | |
| import torch | |
| import sys | |
| import json | |
| import tempfile | |
| from pathlib import Path | |
| from PIL import Image | |
| # Add the UniICL code directory to path | |
| sys.path.insert(0, str(Path(__file__).parent / "UniICL")) | |
| from modeling.uniicl import ( | |
| BagelConfig, Bagel, Qwen2Config, Qwen2ForCausalLM, | |
| SiglipVisionConfig, SiglipVisionModel | |
| ) | |
| from modeling.qwen2 import Qwen2Tokenizer | |
| from modeling.autoencoder import load_ae | |
| from models.capm import CapmConfig | |
| from data.transforms import ImageTransform | |
| from data.data_utils import add_special_tokens, pil_img2rgb | |
| from inferencer import InterleaveInferencer | |
| from copy import deepcopy | |
| from safetensors import safe_open | |
| from safetensors.torch import load_file | |
| MODEL_ID = "xuyicheng-zju/UniICL" | |
| def load_model(): | |
| """Load the UniICL model from Hugging Face Hub.""" | |
| from huggingface_hub import snapshot_download | |
| model_path = snapshot_download(MODEL_ID) | |
| # LLM config | |
| llm_config = Qwen2Config.from_json_file(os.path.join(model_path, "llm_config.json")) | |
| llm_config.qk_norm = True | |
| llm_config.tie_word_embeddings = False | |
| llm_config.layer_module = "Qwen2MoTDecoderLayer" | |
| # ViT config | |
| vit_config = SiglipVisionConfig.from_json_file(os.path.join(model_path, "vit_config.json")) | |
| vit_config.rope = False | |
| vit_config.num_hidden_layers = vit_config.num_hidden_layers - 1 | |
| # VAE | |
| vae_model, vae_config = load_ae(local_path=os.path.join(model_path, "ae.safetensors")) | |
| vae_model = vae_model.to(device="cuda", dtype=torch.bfloat16) | |
| vae_model.eval() | |
| # CAPM config | |
| capm_config = None | |
| capm_config_path = os.path.join(model_path, "capm_config.json") | |
| if os.path.exists(capm_config_path): | |
| print(f"Loading CAPM config from {capm_config_path}...") | |
| with open(capm_config_path, "r") as f: | |
| capm_dict = json.load(f) | |
| capm_config = CapmConfig(**capm_dict) | |
| print(f"[OK] CAPM config loaded: d_capm={capm_config.d_capm}, num_inject_layers={capm_config.num_inject_layers}") | |
| # Bagel config | |
| config = BagelConfig( | |
| visual_gen=True, | |
| visual_und=True, | |
| llm_config=llm_config, | |
| vit_config=vit_config, | |
| vae_config=vae_config, | |
| capm_config=capm_config, | |
| vit_max_num_patch_per_side=70, | |
| connector_act="gelu_pytorch_tanh", | |
| latent_patch_size=2, | |
| max_latent_size=64, | |
| ) | |
| tokenizer = Qwen2Tokenizer.from_pretrained(model_path) | |
| tokenizer, new_token_ids, _ = add_special_tokens(tokenizer) | |
| # Transforms | |
| vae_transform = ImageTransform(1024, 512, 16) | |
| vit_transform = ImageTransform(980, 224, 14) | |
| default_dtype = torch.get_default_dtype() | |
| torch.set_default_dtype(torch.bfloat16) | |
| try: | |
| language_model = Qwen2ForCausalLM(llm_config) | |
| vit_model = SiglipVisionModel(vit_config) | |
| model = Bagel(language_model, vit_model, config) | |
| model.vit_model.vision_model.embeddings.convert_conv2d_to_linear(vit_config, meta=False) | |
| finally: | |
| torch.set_default_dtype(default_dtype) | |
| # Load checkpoint | |
| checkpoint_path = os.path.join(model_path, "ema.safetensors") | |
| if not os.path.exists(checkpoint_path): | |
| checkpoint_path = os.path.join(model_path, "model.safetensors") | |
| print(f"Loading checkpoint from {checkpoint_path}...") | |
| def _load_checkpoint_streaming(model, checkpoint_path): | |
| model_state = model.state_dict() | |
| loaded_keys = set() | |
| unexpected_keys = [] | |
| remapped_count = 0 | |
| with safe_open(checkpoint_path, framework="pt", device="cpu") as f: | |
| for key in f.keys(): | |
| load_key = key | |
| if key.startswith("prism."): | |
| candidate = "capm." + key[len("prism."):] | |
| if candidate in model_state: | |
| load_key = candidate | |
| remapped_count += 1 | |
| elif key in model_state: | |
| pass | |
| else: | |
| load_key = candidate | |
| if load_key not in model_state: | |
| unexpected_keys.append(load_key) | |
| continue | |
| tensor = f.get_tensor(key) | |
| target = model_state[load_key] | |
| if tuple(tensor.shape) != tuple(target.shape): | |
| continue | |
| if tensor.dtype != target.dtype: | |
| tensor = tensor.to(dtype=target.dtype) | |
| with torch.no_grad(): | |
| target.copy_(tensor) | |
| loaded_keys.add(load_key) | |
| if remapped_count > 0: | |
| print(f" Remapped {remapped_count} legacy PRISM keys to CAPM keys") | |
| missing_keys = [key for key in model_state.keys() if key not in loaded_keys] | |
| if missing_keys: | |
| print(f" Missing keys: {len(missing_keys)}") | |
| if unexpected_keys: | |
| print(f" Unexpected keys: {len(unexpected_keys)}") | |
| _load_checkpoint_streaming(model, checkpoint_path) | |
| print("Moving model to CUDA...") | |
| model = model.to(device="cuda", dtype=torch.bfloat16) | |
| model.eval() | |
| return model, vae_model, tokenizer, vae_transform, vit_transform, new_token_ids | |
| print("Loading UniICL model...") | |
| model, vae_model, tokenizer, vae_transform, vit_transform, new_token_ids = load_model() | |
| print("Model loaded successfully!") | |
| inferencer = InterleaveInferencer( | |
| model=model, | |
| vae_model=vae_model, | |
| tokenizer=tokenizer, | |
| vae_transform=vae_transform, | |
| vit_transform=vit_transform, | |
| new_token_ids=new_token_ids, | |
| ) | |
| def run_understanding( | |
| image: Image.Image, | |
| prompt: str, | |
| max_new_tokens: int = 512, | |
| temperature: float = 0.3, | |
| do_sample: bool = False, | |
| ) -> str: | |
| """Answer a question about the provided image using UniICL's understanding capability. | |
| Args: | |
| image: Input image to analyze. | |
| prompt: Question or instruction about the image. | |
| max_new_tokens: Maximum number of tokens to generate. | |
| temperature: Sampling temperature for text generation. | |
| do_sample: Whether to use sampling instead of greedy decoding. | |
| """ | |
| if image is None: | |
| return "Please provide an image." | |
| if not prompt: | |
| return "Please provide a prompt/question." | |
| image = pil_img2rgb(image) | |
| result = inferencer( | |
| image=image, | |
| text=prompt, | |
| understanding_output=True, | |
| max_think_token_n=max_new_tokens, | |
| do_sample=do_sample, | |
| text_temperature=temperature, | |
| ) | |
| return result.get("text", "No output generated.") | |
| def run_generation( | |
| prompt: str, | |
| num_timesteps: int = 28, | |
| cfg_text_scale: float = 3.0, | |
| cfg_img_scale: float = 1.5, | |
| image: Image.Image = None, | |
| ) -> Image.Image: | |
| """Generate an image from a text prompt using UniICL's generation capability. | |
| Args: | |
| prompt: Text description of the image to generate. | |
| num_timesteps: Number of diffusion timesteps (more = higher quality, slower). | |
| cfg_text_scale: Text classifier-free guidance scale. | |
| cfg_img_scale: Image classifier-free guidance scale. | |
| image: Optional reference image for image-to-image generation. | |
| """ | |
| if not prompt: | |
| return None | |
| input_list = [] | |
| if image is not None: | |
| input_list.append(pil_img2rgb(image)) | |
| input_list.append(prompt) | |
| result = inferencer.interleave_inference( | |
| input_list, | |
| understanding_output=False, | |
| num_timesteps=num_timesteps, | |
| cfg_text_scale=cfg_text_scale, | |
| cfg_img_scale=cfg_img_scale, | |
| ) | |
| for item in result: | |
| if isinstance(item, Image.Image): | |
| return item | |
| return None | |
| import gradio as gr | |
| CSS = """ | |
| #col-container { max-width: 1100px; margin: 0 auto; } | |
| .dark .gradio-container { color: var(--body-text-color); } | |
| """ | |
| with gr.Blocks(theme=gr.themes.Citrus(), css=CSS) as demo: | |
| gr.Markdown( | |
| """ | |
| # 🧠 UniICL: Unified Multimodal In-context Learning | |
| UniICL is a unified multimodal in-context learning framework for both visual understanding and generation. | |
| It demonstrates in-context learning across 6 capability categories: perception, imitation, conception, deduction, analogy, and discernment. | |
| [Paper](https://arxiv.org/abs/2603.24690) | [GitHub](https://github.com/xuyicheng-zju/UniICL) | [Model](https://huggingface.co/xuyicheng-zju/UniICL) | |
| """ | |
| ) | |
| with gr.Tabs(): | |
| # ── Understanding Tab ── | |
| with gr.Tab("Visual Understanding"): | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| und_image = gr.Image(label="Input Image", type="pil") | |
| und_prompt = gr.Textbox( | |
| label="Question / Instruction", | |
| placeholder="Describe what you see in this image.", | |
| lines=3, | |
| ) | |
| und_btn = gr.Button("Analyze", variant="primary") | |
| with gr.Accordion("Advanced Settings", open=False): | |
| und_max_tokens = gr.Slider( | |
| minimum=64, maximum=2048, value=512, step=64, | |
| label="Max New Tokens" | |
| ) | |
| und_temp = gr.Slider( | |
| minimum=0.0, maximum=2.0, value=0.3, step=0.1, | |
| label="Temperature" | |
| ) | |
| und_sample = gr.Checkbox(label="Use Sampling", value=False) | |
| with gr.Column(scale=1): | |
| und_output = gr.Textbox(label="Response", lines=12, interactive=False) | |
| und_btn.click( | |
| run_understanding, | |
| inputs=[und_image, und_prompt, und_max_tokens, und_temp, und_sample], | |
| outputs=und_output, | |
| api_name="understand", | |
| ) | |
| gr.Examples( | |
| examples=[ | |
| ["astronaut.jpg", "What is happening in this image? Describe the scene in detail."], | |
| ["autumn_forest_path.jpg", "Describe the colors and mood of this forest path."], | |
| ["bird_kingfisher.jpg", "What species of bird is this? Describe its key features."], | |
| ], | |
| inputs=[und_image, und_prompt], | |
| outputs=und_output, | |
| fn=run_understanding, | |
| cache_examples=True, | |
| cache_mode="lazy", | |
| ) | |
| # ── Generation Tab ── | |
| with gr.Tab("Image Generation"): | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| gen_prompt = gr.Textbox( | |
| label="Generation Prompt", | |
| placeholder="Generate an image of a lighthouse at dusk.", | |
| lines=3, | |
| ) | |
| gen_ref_image = gr.Image( | |
| label="Reference Image (optional, for image editing)", | |
| type="pil", | |
| ) | |
| gen_btn = gr.Button("Generate", variant="primary") | |
| with gr.Accordion("Advanced Settings", open=False): | |
| gen_steps = gr.Slider( | |
| minimum=4, maximum=100, value=28, step=4, | |
| label="Diffusion Steps" | |
| ) | |
| gen_cfg_text = gr.Slider( | |
| minimum=1.0, maximum=10.0, value=3.0, step=0.5, | |
| label="Text CFG Scale" | |
| ) | |
| gen_cfg_img = gr.Slider( | |
| minimum=1.0, maximum=5.0, value=1.5, step=0.5, | |
| label="Image CFG Scale" | |
| ) | |
| with gr.Column(scale=1): | |
| gen_output = gr.Image(label="Generated Image", type="pil") | |
| gen_btn.click( | |
| run_generation, | |
| inputs=[gen_prompt, gen_steps, gen_cfg_text, gen_cfg_img, gen_ref_image], | |
| outputs=gen_output, | |
| api_name="generate", | |
| ) | |
| gr.Examples( | |
| examples=[ | |
| ["A serene lighthouse at dusk, with warm golden light reflecting off calm ocean waves", 28, 3.0, 1.5, None], | |
| ["A cozy cabin in a snow-covered forest, smoke rising from the chimney", 28, 3.0, 1.5, None], | |
| ["An abstract painting of a cosmic nebula with vibrant purples and blues", 28, 3.0, 1.5, None], | |
| ], | |
| inputs=[gen_prompt, gen_steps, gen_cfg_text, gen_cfg_img, gen_ref_image], | |
| outputs=gen_output, | |
| fn=run_generation, | |
| cache_examples=True, | |
| cache_mode="lazy", | |
| ) | |
| demo.launch(mcp_server=True) |