import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import spaces # noqa: E402 import re import numpy as np import PIL.Image import torch import gradio as gr from huggingface_hub import hf_hub_download, snapshot_download from transformers import AutoModelForCausalLM, AutoModel # --------------------------------------------------------------------------- # Configuration # --------------------------------------------------------------------------- REPO_ID = "gaoyuan-ai/SA-IQA-model" SUBFOLDER = "sa-iqa-prompt4" LABELS = ["excellent", "good", "fair", "poor", "bad"] LABEL_TO_SCORE = {"excellent": 5.0, "good": 4.0, "fair": 3.0, "poor": 2.0, "bad": 1.0} DIMENSIONS = ["distortion", "harmony", "layout", "lighting"] # The reference pipeline (SA-IQA tools/infer.py) drives the checkpoint through # ms-swift's `ovis2_5` template. The values below mirror that template exactly so # this Space reproduces the authors' numbers: # * system prompt -- tools/train_sft.sh passes `--system 'You are a helpful assistant.'`, # and swift's PtEngine replays it at inference time from the checkpoint's args.json. # * min/max pixels -- Ovis2_5Template.init_env_args in ms-swift. # * top_logprobs = 5 -- RequestConfig in tools/infer.py::run_inference. SYSTEM_PROMPT = "You are a helpful assistant." MIN_PIXELS = 448 * 448 MAX_PIXELS = 1344 * 1792 TOP_LOGPROBS = 5 # Prompt version 4 (the released final checkpoint), verbatim from # tools/prompt_configs.py::PROMPTS[4]. PROMPTS = { "distortion": ( "Please evaluate the spatial aesthetic distortion quality level of this image. " "The distortion dimension assesses whether soft furnishings (e.g., cabinets, carpets) or fixed " "structures (e.g., floors, walls) appear deformed or misaligned. Additionally, evaluate the realism " "and material accuracy of textures, and judge whether any distortion negatively impacts the overall " "aesthetic quality of the image." ), "harmony": ( "Please evaluate the spatial aesthetic harmony quality level of this image. " "The harmony dimension focuses on stylistic consistency, color coordination, and overall visual cohesion. " "Examine how well the combination of elements creates a balanced and visually pleasant composition, " "avoiding clashes or imbalances in style and color." ), "layout": ( "Please evaluate the spatial aesthetic layout quality level of this image. " "The layout dimension describes the spatial distribution, positional relationships, and quantity of major " "elements within the space. Consider how the layout supports the overall visual order, maintains balance, " "and enhances the functional aesthetics of the image." ), "lighting": ( "Please evaluate the spatial aesthetic lighting quality level of this image. " "The lighting dimension examines the quality of light effects, shadow interactions, and the realism of " "light sources. Assess how well lighting contributes to the overall depth, mood, and authenticity of the " "image, emphasizing both natural and artificial lighting scenarios." ), } # --------------------------------------------------------------------------- # Load model at module scope (ZeroGPU rule #2) # --------------------------------------------------------------------------- print(f"Downloading model weights from {REPO_ID}/{SUBFOLDER} ...") # Download only the sa-iqa-prompt4 subfolder (the released checkpoint), # not the redundant Ovis2.5-9B base model copy. local_model_dir = snapshot_download( repo_id=REPO_ID, repo_type="model", allow_patterns=[f"{SUBFOLDER}/*"], ) model_path = os.path.join(local_model_dir, SUBFOLDER) print(f"Loading model from {model_path} ...") # Patch the custom modeling code on disk for compatibility with transformers 5.x. # The Ovis2.5 modeling code was written for transformers 4.x and uses # attributes/methods that changed in 5.x. _modeling_file = os.path.join(model_path, "modeling_ovis2_5.py") # Always start from a pristine copy of the upstream file. The patch below is # *prepended*, so patching an already-patched blob (possible when the HF cache # survives a restart) would stack definitions and let a stale one win. _pristine = hf_hub_download( repo_id=REPO_ID, repo_type="model", filename=f"{SUBFOLDER}/modeling_ovis2_5.py", force_download=True, ) with open(_pristine, "r") as f: _src = f.read() assert "_transformers_5x_compat_v2" not in _src, "pristine copy is already patched" if True: # 1. Add is_parallelizable = False to nn.Module (used in __init__) # 2. Fix tie_weights to accept missing_keys and recompute_mapping kwargs # 3. Provide pure-torch stand-ins for the two flash_attn helpers the SigLIP2 # vision tower imports. The ZeroGPU runtime now ships torch 2.13, for # which no prebuilt flash-attn wheel exists (ABI mismatch), and the # reference SA-IQA inference path (ms-swift PtEngine) runs the vision # tower with plain SDPA anyway. Exact (non-approximate) attention. _patch = ''' # _transformers_5x_compat_v2 import torch as _torch import torch.nn as nn from torch.nn import functional as _F if not hasattr(nn.Module, 'is_parallelizable'): nn.Module.is_parallelizable = False def apply_rotary_emb(x, cos, sin, interleaved=False, **kwargs): """Pure-torch stand-in for flash_attn.layers.rotary.apply_rotary_emb. x: (batch, seqlen, nheads, headdim) cos, sin: (seqlen, rotary_dim // 2) """ ro_dim = cos.shape[-1] * 2 assert ro_dim <= x.shape[-1] cos = _torch.cat([cos, cos], dim=-1).unsqueeze(-2) sin = _torch.cat([sin, sin], dim=-1).unsqueeze(-2) x_ro = x[..., :ro_dim] x1, x2 = x_ro.chunk(2, dim=-1) out = x_ro * cos + _torch.cat((-x2, x1), dim=-1) * sin if ro_dim < x.shape[-1]: out = _torch.cat([out, x[..., ro_dim:]], dim=-1) return out def flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, *args, **kwargs): """Pure-torch (SDPA) stand-in for flash_attn.flash_attn_varlen_func. q/k/v: (total_tokens, nheads, headdim); non-causal self-attention over the variable-length segments described by cu_seqlens. Each segment is attended to independently rather than via one big block-diagonal `attn_mask`. Passing an explicit mask forces PyTorch to fall back to the `math` SDPA backend, which runs the softmax in bf16 and is both slower and markedly less accurate; without a mask SDPA can use the flash / mem-efficient kernels, which accumulate in fp32 exactly like the real flash_attn_varlen_func this replaces. That keeps vision-tower outputs numerically stable and reproducible across machines. """ bounds = cu_seqlens_q.tolist() outs = [] for i in range(1, len(bounds)): s, e = bounds[i - 1], bounds[i] if e <= s: continue outs.append( _F.scaled_dot_product_attention( q[s:e].transpose(0, 1).unsqueeze(0), k[s:e].transpose(0, 1).unsqueeze(0), v[s:e].transpose(0, 1).unsqueeze(0), dropout_p=0.0, ).squeeze(0).transpose(0, 1) ) return _torch.cat(outs, dim=0) ''' _src = _patch + "\n" + _src # Drop the flash_attn imports; the stand-ins above replace them. _src = _src.replace( "from flash_attn import flash_attn_varlen_func\n" "from flash_attn.layers.rotary import apply_rotary_emb\n", "", ) # Fix tie_weights to accept the new transformers 5.x signature _src = _src.replace( " def tie_weights(self):\n self.llm.tie_weights()", " def tie_weights(self, **kwargs):\n self.llm.tie_weights(**kwargs)" ) with open(_modeling_file, "w") as f: f.write(_src) # Patch PreTrainedModel for compatibility with transformers 5.x. from transformers import PreTrainedModel if not hasattr(PreTrainedModel, 'all_tied_weights_keys'): PreTrainedModel.all_tied_weights_keys = {} # Monkey-patch AutoModel.from_config to handle Siglip2NavitConfig. _original_from_config = AutoModel.from_config.__func__ @classmethod def _patched_from_config(cls, config, *args, **kwargs): if hasattr(config, 'model_type') and config.model_type == 'siglip2_navit': import importlib.util, sys, types modeling_file = os.path.join(model_path, "modeling_ovis2_5.py") pkg_name = "_ovis_pkg" if pkg_name not in sys.modules: pkg = types.ModuleType(pkg_name) pkg.__path__ = [model_path] sys.modules[pkg_name] = pkg cfg_file = os.path.join(model_path, "configuration_ovis2_5.py") if f"{pkg_name}.configuration_ovis2_5" not in sys.modules: cfg_spec = importlib.util.spec_from_file_location(f"{pkg_name}.configuration_ovis2_5", cfg_file) cfg_mod = importlib.util.module_from_spec(cfg_spec) sys.modules[f"{pkg_name}.configuration_ovis2_5"] = cfg_mod cfg_spec.loader.exec_module(cfg_mod) if f"{pkg_name}.modeling_ovis2_5" not in sys.modules: spec2 = importlib.util.spec_from_file_location(f"{pkg_name}.modeling_ovis2_5", modeling_file) mod2 = importlib.util.module_from_spec(spec2) sys.modules[f"{pkg_name}.modeling_ovis2_5"] = mod2 spec2.loader.exec_module(mod2) return sys.modules[f"{pkg_name}.modeling_ovis2_5"].Siglip2NavitModel(config) return _original_from_config(cls, config, *args, **kwargs) AutoModel.from_config = _patched_from_config model = AutoModelForCausalLM.from_pretrained( model_path, torch_dtype=torch.bfloat16, trust_remote_code=True, ).to("cuda").eval() # Fix: re-initialize indicator_token_indices. # `indicator_token_indices` is registered as a NON-PERSISTENT buffer, so it is not # part of state_dict(): transformers 5.x's post-loading steps corrupt it, and on # ZeroGPU it is also not carried across when the model is moved into the GPU worker # process. That silently produces wrong visual-indicator embeddings, which varies # between container restarts. Restore it explicitly, on the model's current device, # immediately before every inference call. _NUM_INDICATOR_IDS = 4 # len(INDICATOR_IDS) in modeling_ovis2_5.py def _restore_indicator_token_indices(): want = torch.arange( model.config.visual_vocab_size - _NUM_INDICATOR_IDS, model.config.visual_vocab_size, dtype=torch.long, device=model.vte.weight.device, ) cur = getattr(model, "indicator_token_indices", None) if cur is None or cur.device != want.device or not torch.equal(cur, want): model.indicator_token_indices = want _restore_indicator_token_indices() tokenizer = model.text_tokenizer print("Model loaded successfully.") # --------------------------------------------------------------------------- # Scoring (ported from tools/infer.py) # --------------------------------------------------------------------------- def calculate_iqa_score(logits: dict) -> float: """Convert class log-probabilities into a continuous score in [1, 5].""" safe = {} for key in LABELS: safe[key] = logits.get(key, -50.0) logprobs = np.array([safe[k] for k in LABELS], dtype=np.float32) logprobs = logprobs - np.max(logprobs) probs = np.exp(logprobs) / np.sum(np.exp(logprobs)) score_values = np.array([5, 4, 3, 2, 1], dtype=np.float32) return float(np.inner(probs, score_values)) def _decode_token(token_id) -> str: """Decode a single token the way tools/infer.py normalises top-logprob tokens.""" return tokenizer.decode(int(token_id)).strip().split(" ")[-1] def _extract_top_logprobs_as_dict(logits_per_step, generated_ids): """Port of tools/infer.py::extract_top_logprobs_as_dict (returns (dict, index)). The reference implementation reads the *top-5 logprobs of the rating token only*. With the SA-IQA answer template "The spatial aesthetic quality level of this image is ." the generated token stream ends with [..., " ", ".", "<|im_end|>"], so the rating token sits at index -3 -- exactly what infer.py indexes. """ n = min(len(logits_per_step), len(generated_ids)) if n == 0: return {}, -1 idx = n - 3 # template-dependent heuristic, as in the original implementation if idx < 0 or _decode_token(generated_ids[idx]).lower() not in LABELS: # Fall back gracefully if the model deviates from the answer template. found = [i for i in range(n) if _decode_token(generated_ids[i]).lower() in LABELS] idx = found[-1] if found else max(n - 3, 0) log_probs = torch.log_softmax(logits_per_step[idx][0].float(), dim=-1) top = torch.topk(log_probs, TOP_LOGPROBS) logits_dict = {} for logprob, token_id in zip(top.values.tolist(), top.indices.tolist()): token = _decode_token(token_id) logits_dict[token] = logprob return logits_dict, idx def _extract_label_from_text(text: str) -> str: """Find the quality label in the model output text.""" for l in LABELS: if re.search(r'\b' + l + r'\b', text, re.IGNORECASE): return l return "unknown" def _run_single_dimension(image: PIL.Image.Image, dimension: str) -> dict: """Run inference for one dimension. Must be called inside @spaces.GPU.""" prompt_text = PROMPTS[dimension] # The model's preprocess_inputs expects content as a list of typed dicts, # not a plain string — otherwise the image is never extracted. The chat # template strips the literal "" out of the text part and emits the # image placeholder followed by "\n", which is byte-for-byte what ms-swift's # Ovis2_5Template.replace_tag ([[-200], '\n']) produced during training. content = [ {"type": "image", "image": image}, {"type": "text", "text": prompt_text}, ] messages = [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": content}, ] # enable_thinking=None -> the chat template does NOT inject an empty # "\n\n" block. ms-swift's ovis2_5 template never injects one, # so the SA-IQA checkpoint was trained/evaluated to answer straight after # "<|im_start|>assistant\n". input_ids, pixel_values, grid_thws = model.preprocess_inputs( messages=messages, add_generation_prompt=True, enable_thinking=None, min_pixels=MIN_PIXELS, max_pixels=MAX_PIXELS, ) _restore_indicator_token_indices() input_ids = input_ids.to("cuda") if pixel_values is not None: pixel_values = pixel_values.to("cuda") if grid_thws is not None: grid_thws = grid_thws.to("cuda") with torch.inference_mode(): output = model.generate( inputs=input_ids, pixel_values=pixel_values, grid_thws=grid_thws, max_new_tokens=64, do_sample=False, eos_token_id=tokenizer.eos_token_id, pad_token_id=tokenizer.pad_token_id, use_cache=True, enable_thinking=False, return_dict_in_generate=True, # Raw (unprocessed) logits, which is what ms-swift's PtEngine turns # into logprobs (`generation_config.output_logits = True`). output_logits=True, ) generated_ids = output.sequences[0] response_text = tokenizer.decode(generated_ids, skip_special_tokens=True).strip() # Extract the top-5 logprobs of the rating token, exactly as tools/infer.py does. logits_dict, rating_idx = {}, -1 try: logits_dict, rating_idx = _extract_top_logprobs_as_dict(output.logits, generated_ids) except Exception as e: print(f"Logits extraction error: {e}") if rating_idx >= 0 and _decode_token(generated_ids[rating_idx]).lower() in LABELS: label = _decode_token(generated_ids[rating_idx]).lower() else: label = _extract_label_from_text(response_text) if logits_dict: score = calculate_iqa_score(logits_dict) # Also compute individual probabilities for display safe = {k: logits_dict.get(k, -50.0) for k in LABELS} lp_arr = np.array([safe[k] for k in LABELS], dtype=np.float32) lp_arr = lp_arr - np.max(lp_arr) probs = np.exp(lp_arr) / np.sum(np.exp(lp_arr)) prob_display = {l: float(p) for l, p in zip(LABELS, probs)} else: score = LABEL_TO_SCORE.get(label, 3.0) prob_display = {} return { "dimension": dimension, "label": label, "score": score, "text": response_text, "probabilities": prob_display, } # --------------------------------------------------------------------------- # Gradio inference functions # --------------------------------------------------------------------------- @spaces.GPU(duration=120) def assess(image: PIL.Image.Image, mode: str = "all", dimension: str = "lighting"): """Assess the spatial aesthetic quality of an interior image. Args: image: An interior image to evaluate. mode: "all" to assess all four dimensions, or "single" for one dimension. dimension: The dimension to assess in single mode (distortion, harmony, layout, lighting). Returns: Markdown report with assessment results. """ if image is None: return "⚠️ Please provide an image to assess." if mode == "single": result = _run_single_dimension(image, dimension) md = f"### {dimension.capitalize()} Assessment\n\n" md += f"| Field | Value |\n|-------|-------|\n" md += f"| **Label** | {result['label']} |\n" md += f"| **Score** | {result['score']:.2f} / 5.00 |\n\n" md += f"**Model output:** {result['text']}\n\n" if result["probabilities"]: md += "### Label Probabilities\n\n" for l in LABELS: p = result["probabilities"].get(l, 0.0) bar = "█" * int(p * 20) md += f"- **{l}**: {p:.1%} {bar}\n" return md # All dimensions results = [] for dim in DIMENSIONS: r = _run_single_dimension(image, dim) results.append(r) md = "### Full Assessment Report\n\n" md += "| Dimension | Label | Score |\n" md += "|-----------|--------|-------|\n" for r in results: md += f"| {r['dimension'].capitalize()} | {r['label']} | {r['score']:.2f} |\n" avg = np.mean([r["score"] for r in results]) md += f"\n**Overall Score: {avg:.2f} / 5.00**\n\n" md += "---\n\n" for r in results: md += f"**{r['dimension'].capitalize()}**: {r['text']}\n\n" return md # --------------------------------------------------------------------------- # Gradio UI # --------------------------------------------------------------------------- CSS = """ #col-container { max-width: 1100px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } """ with gr.Blocks() as demo: gr.Markdown(""" # SA-IQA: Spatial Aesthetics Image Quality Assessment Assess the spatial aesthetic quality of interior images across four dimensions: **distortion**, **harmony**, **layout**, and **lighting**. Based on [gaoyuan-ai/SA-IQA-model](https://huggingface.co/gaoyuan-ai/SA-IQA-model), a fine-tuned Ovis2.5-9B vision-language model. [Paper (CVPRW 2026)](https://arxiv.org/abs/2512.05098) | [GitHub](https://github.com/AlibabaResearch/SA-IQA) """) with gr.Row(): with gr.Column(scale=1): image_input = gr.Image(label="Interior Image", type="pil", height=400) mode = gr.Radio( choices=["all", "single"], value="all", label="Evaluation Mode", info="'all' assesses all four dimensions; 'single' assesses one chosen dimension.", ) dimension = gr.Dropdown( choices=DIMENSIONS, value="lighting", label="Dimension (single mode only)", interactive=True, ) assess_btn = gr.Button("Assess", variant="primary") with gr.Column(scale=1): result_output = gr.Markdown( "Upload an interior image and click **Assess** to evaluate its spatial aesthetics.", label="Result", ) gr.Examples( examples=[ ["examples/lighting_example.jpg", "all", "lighting"], ["examples/distortion_example.jpg", "single", "distortion"], ["examples/harmony_example.jpg", "single", "harmony"], ["examples/layout_example.jpg", "single", "layout"], ], inputs=[image_input, mode, dimension], fn=assess, outputs=[result_output], cache_examples=False, run_on_click=True, ) gr.Markdown(""" ### Dimension Definitions | Dimension | Description | |-----------|-------------| | **Distortion** | Whether soft furnishings or fixed structures appear deformed or misaligned; realism and material accuracy of textures. | | **Harmony** | Stylistic consistency, color coordination, and overall visual cohesion of the composition. | | **Layout** | Spatial distribution, positional relationships, and quantity of major elements; overall visual order and balance. | | **Lighting** | Quality of light effects, shadow interactions, and realism of light sources; contribution to depth and mood. | Labels: **excellent** (5) > **good** (4) > **fair** (3) > **poor** (2) > **bad** (1). Continuous scores are the probability-weighted average of those label values, computed from the top-5 log-probabilities of the rating token — the same conversion as `tools/infer.py` in the [SA-IQA repository](https://github.com/AlibabaResearch/SA-IQA), using the `sa-iqa-prompt4` checkpoint and the prompt-v4 templates. """) assess_btn.click( fn=assess, inputs=[image_input, mode, dimension], outputs=[result_output], api_name="assess", ) demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)