Spaces:
Running on Zero
Running on Zero
multimodalart HF Staff
Fix garbled/off-distribution outputs: align empty-subtitle prompt to reference get_prompt_for_multimodal template
6cae6f5 verified | import os | |
| os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") | |
| import spaces # MUST come before any torch/CUDA-touching import | |
| import torch | |
| import torchvision | |
| import torchvision.transforms.functional as _F | |
| import sys | |
| import types | |
| # Compatibility shim: torchvision >= 0.21 removed functional_tensor, which pytorchvideo still imports | |
| if not hasattr(torchvision.transforms, 'functional_tensor'): | |
| _mod = types.ModuleType('torchvision.transforms.functional_tensor') | |
| for _name in dir(_F): | |
| setattr(_mod, _name, getattr(_F, _name)) | |
| sys.modules['torchvision.transforms.functional_tensor'] = _mod | |
| torchvision.transforms.functional_tensor = _mod | |
| import gradio as gr | |
| import numpy as np | |
| import tempfile | |
| import time | |
| from pathlib import Path | |
| from huggingface_hub import snapshot_download, hf_hub_download | |
| from omegaconf import OmegaConf | |
| import copy | |
| import einops | |
| # Patch config module to use HF Hub paths before any model imports | |
| import config as _config_module | |
| # Download pretrained models from HF Hub and patch config paths | |
| _qwen3_06b_path = snapshot_download("Qwen/Qwen3-0.6B") | |
| _clip_path = snapshot_download("openai/clip-vit-base-patch16") | |
| _hubert_path = snapshot_download("TencentGameMate/chinese-hubert-base") | |
| _config_module.PATH_TO_LLM['Qwen3_0.6B'] = _qwen3_06b_path | |
| _config_module.PATH_TO_VISUAL['CLIP_VIT_BASE'] = _clip_path | |
| _config_module.PATH_TO_AUDIO['HUBERT_BASE'] = _hubert_path | |
| # Now import the model code (after config is patched) | |
| from my_affectgpt.common.registry import registry | |
| from my_affectgpt.models.affectgpt import AffectGPT | |
| from my_affectgpt.models.encoder import * # registers encoders | |
| from my_affectgpt.conversation.conversation_video import Chat | |
| from my_affectgpt.processors.video_processor import ( | |
| AlproVideoEvalProcessor, | |
| load_video, | |
| ) | |
| from my_affectgpt.processors.blip_processors import Blip2ImageEvalProcessor | |
| from my_affectgpt.models.ImageBind.data import load_audio, transform_audio | |
| # ---- Model config matching stage1-swdh-qwen3-0.6b/config.yaml ---- | |
| MODEL_CFG = OmegaConf.create({ | |
| "arch": "affectgpt", | |
| "model_type": "pretrain_vicuna", | |
| "frozen_video_proj": False, | |
| "frozen_video_Qformer": False, | |
| "frozen_audio_Qformer": False, | |
| "frozen_audio_proj": False, | |
| "frozen_multi_Qformer": False, | |
| "frozen_multi_llama_proj": False, | |
| "frozen_llm": False, | |
| "multi_fusion_type": "attention", | |
| "video_fusion_type": "attention", | |
| "audio_fusion_type": "attention", | |
| "image_fusion_type": "mean", | |
| "ckpt": "", | |
| "ckpt_2": "", | |
| "llama_model": "Qwen3_0.6B", | |
| "acoustic_encoder": "HUBERT_BASE", | |
| "visual_encoder": "CLIP_VIT_BASE", | |
| "num_audio_query_token": 1, | |
| "num_video_query_token": 1, | |
| "num_multi_query_token": 1, | |
| "num_image_query_token": 1, | |
| "max_length": 1024, | |
| "lora_r": 16, | |
| "vis_processor": { | |
| "train": { | |
| "name": "alpro_video_eval", | |
| "n_frms": 8, | |
| "image_size": 224, | |
| } | |
| }, | |
| }) | |
| # Download the Light-MER Stage 1 checkpoint | |
| ckpt_path = hf_hub_download( | |
| "kevin233333/Light-MER", | |
| "stage1-swdh-qwen3-0.6b/checkpoint_000060_loss_1.291.pth", | |
| ) | |
| MODEL_CFG.ckpt_3 = ckpt_path | |
| print("Loading Light-MER model...") | |
| model_cls = registry.get_model_class(MODEL_CFG.arch) | |
| model = model_cls.from_config(MODEL_CFG) | |
| model = model.to("cuda").eval() | |
| chat = Chat(model, MODEL_CFG, device="cuda") | |
| # Video processor for inference | |
| vis_processor = AlproVideoEvalProcessor( | |
| image_size=224, n_frms=8, | |
| mean=(0.48145466, 0.4578275, 0.40821073), | |
| std=(0.26862954, 0.26130258, 0.27577711), | |
| ) | |
| # Default prompt (matching the dataset's description question) | |
| DEFAULT_QUESTION = "Please infer the person's emotional state and provide your reasoning process." | |
| FACE_OR_FRAME = "multiface_audio_face_text" | |
| def recognize_emotion(video_path, subtitle, question, audio_path=None, | |
| max_new_tokens=512, temperature=1.0, top_p=0.9, | |
| do_sample=True, progress=gr.Progress(track_tqdm=True)): | |
| """Recognize emotions from video/audio using Light-MER. | |
| Args: | |
| video_path: Path to the input video file. | |
| audio_path: Path to the input audio file (optional if video has audio). | |
| subtitle: Subtitle text for the video (optional). | |
| question: The question to ask about the emotional state. | |
| max_new_tokens: Maximum number of tokens to generate. | |
| temperature: Sampling temperature. | |
| top_p: Nucleus sampling probability. | |
| do_sample: Whether to use sampling for generation. | |
| """ | |
| if video_path is None and (audio_path is None or audio_path == ""): | |
| return "Please provide a video or audio file." | |
| start_time = time.time() | |
| # Build sample_data like the inference script does | |
| sample_data = { | |
| 'frame': None, 'raw_frame': None, | |
| 'face': None, 'raw_face': None, | |
| 'audio': None, 'raw_audio': None, | |
| 'image': None, 'raw_image': None, | |
| } | |
| # Process video -> frames (and faces from frames) | |
| if video_path is not None: | |
| try: | |
| raw_frame, msg = load_video( | |
| video_path=video_path, | |
| n_frms=8, | |
| height=224, | |
| width=224, | |
| sampling="uniform", | |
| return_msg=True, | |
| ) | |
| frame = vis_processor.transform(raw_frame) | |
| sample_data['frame'] = frame | |
| sample_data['raw_frame'] = raw_frame | |
| # Use frames as "faces" (since we don't have OpenFace preprocessed faces) | |
| sample_data['face'] = frame | |
| sample_data['raw_face'] = raw_frame | |
| except Exception as e: | |
| print(f"Video loading error: {e}") | |
| # Process audio | |
| if audio_path is not None and audio_path != "": | |
| try: | |
| raw_audio = load_audio([audio_path], "cpu", clips_per_video=8)[0] | |
| audio = transform_audio(raw_audio, "cpu") | |
| sample_data['audio'] = audio | |
| sample_data['raw_audio'] = raw_audio | |
| except Exception as e: | |
| print(f"Audio loading error: {e}") | |
| # If we have video but no separate audio, try to extract audio from video | |
| if video_path is not None and sample_data['audio'] is None: | |
| try: | |
| raw_audio = load_audio([video_path], "cpu", clips_per_video=8)[0] | |
| audio = transform_audio(raw_audio, "cpu") | |
| sample_data['audio'] = audio | |
| sample_data['raw_audio'] = raw_audio | |
| except Exception as e: | |
| print(f"Audio extraction from video failed: {e}") | |
| if subtitle is None or subtitle.strip() == "": | |
| subtitle = "" | |
| if question is None or question.strip() == "": | |
| question = DEFAULT_QUESTION | |
| # Encode multimodal features | |
| audio_hiddens, audio_llms = chat.postprocess_audio(sample_data) | |
| frame_hiddens, frame_llms = chat.postprocess_frame(sample_data) | |
| face_hiddens, face_llms = chat.postprocess_face(sample_data) | |
| _, image_llms = chat.postprocess_image(sample_data) | |
| multi_llms = None | |
| if face_hiddens is not None and audio_hiddens is not None: | |
| _, multi_llms = chat.postprocess_multi(face_hiddens, audio_hiddens) | |
| img_list = { | |
| 'audio': audio_llms, | |
| 'frame': frame_llms, | |
| 'face': face_llms, | |
| 'image': image_llms, | |
| 'multi': multi_llms, | |
| } | |
| # Build prompt dynamically based on available features | |
| prompt_parts = ["###Human: "] | |
| if multi_llms is not None: | |
| prompt_parts.append("The audio and video merged info is: <Multi><MultiHere></Multi>. ") | |
| if audio_llms is not None: | |
| prompt_parts.append("The audio content is as follows: <Audio><AudioHere></Audio>. ") | |
| if face_llms is not None: | |
| prompt_parts.append("Meanwhile, we uniformly sample raw frames from the video and extract faces from these frames: <Video><FaceHere></Video>. ") | |
| elif frame_llms is not None: | |
| prompt_parts.append("Meanwhile, we uniformly sample raw frames from the video: <Video><FrameHere></Video>. ") | |
| # Always use the exact trained <Subtitle> framing from the reference | |
| # implementation (base_dataset.get_prompt_for_multimodal). During training the | |
| # subtitle can be empty, in which case the reference simply passes an empty | |
| # string through the same template (<Subtitle></Subtitle>). Any deviation from | |
| # this framing (e.g. injecting "there is no subtitle / do not invent" style | |
| # instructions) pushes the model off-distribution and produces garbled, | |
| # self-contradictory hedging. | |
| prompt_parts.append(f"The subtitle of this video is: <Subtitle>{subtitle}</Subtitle>. ") | |
| prompt_parts.append(f"Now, please answer my question based on all the provided information. {question} ###Assistant: ") | |
| prompt = "".join(prompt_parts) | |
| # Run inference | |
| response = chat.answer_sample( | |
| prompt=prompt, | |
| img_list=img_list, | |
| num_beams=1, | |
| temperature=temperature, | |
| do_sample=do_sample, | |
| top_p=top_p, | |
| max_new_tokens=max_new_tokens, | |
| max_length=2000, | |
| ) | |
| elapsed = time.time() - start_time | |
| result = f"{response}\n\n---\nInference time: {elapsed:.2f}s" | |
| return result | |
| # ---- Gradio UI ---- | |
| 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("# Light-MER: Lightweight Multimodal Emotion Recognition") | |
| gr.Markdown( | |
| "A sub-1B parameter multimodal emotion recognition model that processes " | |
| "video, audio, and text to recognize emotions. Upload a video or audio file " | |
| "and ask a question about the emotional state." | |
| ) | |
| gr.Markdown( | |
| "Paper: [Do We Really Need Multimodal Emotion Language Models Larger Than 1B Parameters?](https://arxiv.org/abs/2607.12787) | " | |
| "[GitHub](https://github.com/GAIR-Lab/Light-MER) | " | |
| "[Model](https://huggingface.co/kevin233333/Light-MER)" | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=2): | |
| video_input = gr.Video(label="Video Input") | |
| audio_input = gr.Audio(label="Audio Input (optional)", type="filepath") | |
| subtitle_input = gr.Textbox( | |
| label="Subtitle (optional)", | |
| placeholder="Enter subtitle text if available...", | |
| lines=2, | |
| ) | |
| question_input = gr.Textbox( | |
| label="Question", | |
| value=DEFAULT_QUESTION, | |
| lines=2, | |
| ) | |
| run_btn = gr.Button("Recognize Emotion", variant="primary") | |
| with gr.Column(scale=3): | |
| output_text = gr.Textbox( | |
| label="Emotion Recognition Result", | |
| lines=15, | |
| show_copy_button=True, | |
| ) | |
| with gr.Accordion("Advanced Settings", open=False): | |
| max_tokens = gr.Slider( | |
| label="Max New Tokens", minimum=64, maximum=1024, value=512, step=64 | |
| ) | |
| temp = gr.Slider( | |
| label="Temperature", minimum=0.1, maximum=2.0, value=1.0, step=0.1 | |
| ) | |
| top_p_val = gr.Slider( | |
| label="Top P", minimum=0.1, maximum=1.0, value=0.9, step=0.05 | |
| ) | |
| do_sample_chk = gr.Checkbox(label="Do Sample", value=True) | |
| gr.Examples( | |
| examples=[ | |
| ["examples/man_laughing.mp4", "", "Please infer the person's emotional state and provide your reasoning process."], | |
| ["examples/man_sad.mp4", "", "Please infer the person's emotional state and provide your reasoning process."], | |
| ["examples/woman_laughing_studio.mp4", "", "Please infer the person's emotional state and provide your reasoning process."], | |
| ], | |
| inputs=[video_input, subtitle_input, question_input], | |
| outputs=output_text, | |
| fn=recognize_emotion, | |
| cache_examples=False, | |
| run_on_click=True, | |
| ) | |
| run_btn.click( | |
| fn=recognize_emotion, | |
| inputs=[video_input, subtitle_input, question_input, audio_input, | |
| max_tokens, temp, top_p_val, do_sample_chk], | |
| outputs=output_text, | |
| api_name="recognize_emotion", | |
| ) | |
| demo.launch() |