import os import re import sys import types import spaces import gradio as gr import torch import torch.nn.functional as _F # --- flash_attn shim: Blackwell ZeroGPU does not ship a flash_attn wheel # compatible with the current torch/CUDA. transformers 4.39.2's # modeling_llama.py and oryx's vision tower both do a top-level # `from flash_attn import flash_attn_func, flash_attn_varlen_func`, so we # inject a stub module that routes to torch SDPA before any of those # imports run. def _shim_flash_attn(): if "flash_attn" in sys.modules: return import importlib.machinery mod = types.ModuleType("flash_attn") # transformers' is_flash_attn_2_available uses importlib.util.find_spec # which raises if __spec__ is None. Provide a spec so the lookup # completes; transformers' subsequent metadata-version lookup will # still fail (no .dist-info), so it will treat flash_attn as # "unavailable" and won't try to dispatch to it. Our explicit # `attn_implementation="sdpa"` setting keeps transformers on torch SDPA. mod.__spec__ = importlib.machinery.ModuleSpec("flash_attn", loader=None) def flash_attn_func(q, k, v, dropout_p=0.0, softmax_scale=None, causal=False, **kwargs): # q, k, v: (B, N, H, C) -> SDPA expects (B, H, N, C) q_ = q.transpose(1, 2) k_ = k.transpose(1, 2) v_ = v.transpose(1, 2) out = _F.scaled_dot_product_attention( q_, k_, v_, dropout_p=dropout_p, is_causal=bool(causal), scale=softmax_scale, ) return out.transpose(1, 2).contiguous() def flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=0.0, softmax_scale=None, causal=False, **kwargs): outs = [] cq = cu_seqlens_q.tolist() if hasattr(cu_seqlens_q, "tolist") else list(cu_seqlens_q) ck = cu_seqlens_k.tolist() if hasattr(cu_seqlens_k, "tolist") else list(cu_seqlens_k) for i in range(len(cq) - 1): qi = q[cq[i]:cq[i + 1]].unsqueeze(0).transpose(1, 2) ki = k[ck[i]:ck[i + 1]].unsqueeze(0).transpose(1, 2) vi = v[ck[i]:ck[i + 1]].unsqueeze(0).transpose(1, 2) oi = _F.scaled_dot_product_attention( qi, ki, vi, dropout_p=dropout_p, is_causal=bool(causal), scale=softmax_scale, ) outs.append(oi.transpose(1, 2).squeeze(0)) return torch.cat(outs, dim=0) def flash_attn_qkvpacked_func(qkv, dropout_p=0.0, softmax_scale=None, causal=False, **kwargs): q, k, v = qkv.unbind(dim=2) return flash_attn_func(q, k, v, dropout_p=dropout_p, softmax_scale=softmax_scale, causal=causal) mod.flash_attn_func = flash_attn_func mod.flash_attn_varlen_func = flash_attn_varlen_func mod.flash_attn_qkvpacked_func = flash_attn_qkvpacked_func # submodules that transformers / other libs sometimes import directly bert_padding = types.ModuleType("flash_attn.bert_padding") def _identity_unpad(hidden, attention_mask): return hidden, None, None, None def _identity_pad(*a, **k): raise NotImplementedError("flash_attn shim: pad_input not implemented") bert_padding.unpad_input = _identity_unpad bert_padding.pad_input = _identity_pad bert_padding.index_first_axis = lambda *a, **k: a[0] mod.bert_padding = bert_padding sys.modules["flash_attn"] = mod sys.modules["flash_attn.bert_padding"] = bert_padding _shim_flash_attn() # --- end flash_attn shim from decord import VideoReader, cpu from PIL import Image import numpy as np import transformers from typing import Dict, Optional, Sequence, List from oryx.conversation import conv_templates, SeparatorStyle from oryx.model.builder import load_pretrained_model from oryx.utils import disable_torch_init from oryx.mm_utils import tokenizer_image_token, get_model_name_from_path, KeywordsStoppingCriteria, process_anyres_video_genli,process_anyres_highres_image_genli from oryx.constants import IGNORE_INDEX, DEFAULT_IMAGE_TOKEN, IMAGE_TOKEN_INDEX model_path = "THUdyh/Oryx-1.5-7B" model_name = get_model_name_from_path(model_path) overwrite_config = {} overwrite_config["mm_resampler_type"] = "dynamic_compressor" overwrite_config["patchify_video_feature"] = False overwrite_config["attn_implementation"] = "sdpa" if torch.__version__ >= "2.1.2" else "eager" tokenizer, model, image_processor, context_len = load_pretrained_model(model_path, None, model_name, device_map="cpu", overwrite_config=overwrite_config) device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model.to(device).eval() cur_dir = os.path.dirname(os.path.abspath(__file__)) title_markdown = """
""" bibtext = """ ### Citation ``` @article{liu2024oryx, title={Oryx MLLM: On-Demand Spatial-Temporal Understanding at Arbitrary Resolution}, author={Liu, Zuyan and Dong, Yuhao and Liu, Ziwei and Hu, Winston and Lu, Jiwen and Rao, Yongming}, journal={arXiv preprint arXiv:2409.12961}, year={2024} } ``` """ def preprocess_qwen(sources, tokenizer: transformers.PreTrainedTokenizer, has_image: bool = False, max_len=2048, system_message: str = "You are a helpful assistant.") -> Dict: roles = {"human": "<|im_start|>user", "gpt": "<|im_start|>assistant"} im_start, im_end = tokenizer.additional_special_tokens_ids[:2] nl_tokens = tokenizer("\n").input_ids _system = tokenizer("system").input_ids + nl_tokens _user = tokenizer("user").input_ids + nl_tokens _assistant = tokenizer("assistant").input_ids + nl_tokens # Apply prompt templates input_ids, targets = [], [] source = sources if roles[source[0]["from"]] != roles["human"]: source = source[1:] input_id, target = [], [] system = [im_start] + _system + tokenizer(system_message).input_ids + [im_end] + nl_tokens input_id += system target += [im_start] + [IGNORE_INDEX] * (len(system) - 3) + [im_end] + nl_tokens assert len(input_id) == len(target) for j, sentence in enumerate(source): role = roles[sentence["from"]] if has_image and sentence["value"] is not None and "