import argparse import os import socket import time from typing import Generator, List, Optional, Tuple import gradio as gr import numpy as np from ml_dtypes import bfloat16 from PIL import Image from transformers import AutoConfig, AutoProcessor, AutoTokenizer from axengine import InferenceSession from utils.infer_func import InferManager from utils.vision_output import describe_output_shapes, select_vit_output try: import onnxruntime as ort except Exception: ort = None TASK_PROMPTS = { "ocr": "OCR:", "table": "Table Recognition:", "formula": "Formula Recognition:", "chart": "Chart Recognition:", "spotting": "Spotting:", "seal": "Seal Recognition:", } def _list_host_ips() -> List[str]: ips = set() try: hostname = socket.gethostname() infos = socket.getaddrinfo(hostname, None, family=socket.AF_INET) for info in infos: ip = info[4][0] if ip and not ip.startswith("127."): ips.add(ip) except Exception: pass if not ips: ips.add("127.0.0.1") return sorted(ips) def _prepare_image(image: Image.Image, task: str) -> Tuple[Image.Image, int]: image = image.convert("RGB") resize_h, resize_w = 576, 768 image = image.resize((resize_w, resize_h)) # AX vision model is compiled with fixed 576x768 token layout. # Keep spotting path aligned to avoid variable token counts. max_pixels = 2048 * 28 * 28 if task == "spotting" else 1280 * 28 * 28 return image, max_pixels def _run_vit_onnx( session, pixel_values: np.ndarray, target_hidden_size: int, expected_tokens: Optional[int] = None ) -> Tuple[np.ndarray, List[Tuple[int, ...]]]: outputs = session.run(None, {"pixel_values": pixel_values}) return ( select_vit_output(outputs, target_hidden_size, expected_tokens=expected_tokens), describe_output_shapes(outputs), ) def _run_vit_axmodel( session, pixel_values: np.ndarray, target_hidden_size: int, expected_tokens: Optional[int] = None ) -> Tuple[np.ndarray, List[Tuple[int, ...]]]: outputs = session.run(None, {"pixel_values": pixel_values}) return ( select_vit_output(outputs, target_hidden_size, expected_tokens=expected_tokens), describe_output_shapes(outputs), ) def _expected_image_features(image_grid_thw) -> int: return int(sum(int(t) * int(h) * int(w) for t, h, w in image_grid_thw)) def _expected_image_tokens(image_grid_thw, merge_size: int) -> int: merge_area = int(merge_size) * int(merge_size) return int(sum(int(t) * int(h) * int(w) // merge_area for t, h, w in image_grid_thw)) def _replace_image_tokens( token_ids: List[int], token_embeds: np.ndarray, image_embeds: np.ndarray, image_token_id: int ) -> np.ndarray: image_positions = [idx for idx, token_id in enumerate(token_ids) if token_id == image_token_id] if not image_positions: return token_embeds flat_image_embeds = image_embeds.reshape(-1, image_embeds.shape[-1]) if len(image_positions) != flat_image_embeds.shape[0]: raise ValueError( f"Image tokens and image features do not match: tokens={len(image_positions)}, " f"features={flat_image_embeds.shape[0]}" ) if token_embeds.shape[-1] != flat_image_embeds.shape[-1]: raise ValueError( f"Embedding dim mismatch: token_dim={token_embeds.shape[-1]}, image_dim={flat_image_embeds.shape[-1]}" ) token_embeds[image_positions, :] = flat_image_embeds return token_embeds class PaddleOCRVLGradioDemo: def __init__(self, hf_model: str, axmodel_dir: str, vit_model: str, max_seq_len: int = 2047): self.hf_model = hf_model self.axmodel_dir = axmodel_dir self.vit_model = vit_model self.embeds = np.load(os.path.join(axmodel_dir, "python/model.embed_tokens.weight.npy")) self.tokenizer = AutoTokenizer.from_pretrained(self.hf_model, trust_remote_code=True) self.processor = AutoProcessor.from_pretrained(self.hf_model, trust_remote_code=True) self.config = AutoConfig.from_pretrained(self.hf_model, trust_remote_code=True) self.merge_size = self.config.vision_config.spatial_merge_size if self.vit_model.endswith(".axmodel"): self.vit_session = InferenceSession(self.vit_model) self.vit_mode = "axmodel" else: if ort is None: raise ImportError("onnxruntime is required when --vit_model is an onnx file") providers = ["CPUExecutionProvider"] if "CUDAExecutionProvider" in ort.get_available_providers(): providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] self.vit_session = ort.InferenceSession(self.vit_model, providers=providers) self.vit_mode = "onnx" self.infer_manager = InferManager(self.config, self.axmodel_dir, max_seq_len=max_seq_len) def _build_prompt_inputs(self, image: Image.Image, task: str, user_text: str): image, max_pixels = _prepare_image(image, task) prompt_text = user_text.strip() if user_text.strip() else TASK_PROMPTS[task] messages = [ { "role": "user", "content": [ {"type": "image", "image": image}, {"type": "text", "text": prompt_text}, ], } ] inputs = self.processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", images_kwargs={ "size": { "shortest_edge": self.processor.image_processor.min_pixels, "longest_edge": max_pixels, } }, ) return prompt_text, inputs def _prepare_model_inputs(self, inputs): token_ids = inputs.input_ids[0].cpu().numpy().tolist() image_grid_thw = inputs.image_grid_thw.cpu().numpy().tolist() expected_tokens = _expected_image_tokens(image_grid_thw, self.merge_size) expected_features = _expected_image_features(image_grid_thw) pixel_values = inputs.pixel_values if pixel_values.ndim == 4: pixel_values = pixel_values.unsqueeze(0) pixel_values = pixel_values.cpu().numpy().astype(np.float32) if self.vit_mode == "axmodel": image_embeds, vit_output_shapes = _run_vit_axmodel( self.vit_session, pixel_values, target_hidden_size=self.config.hidden_size, expected_tokens=expected_tokens, ) else: image_embeds, vit_output_shapes = _run_vit_onnx( self.vit_session, pixel_values, target_hidden_size=self.config.hidden_size, expected_tokens=expected_tokens, ) if image_embeds.ndim == 3: image_embeds = image_embeds[0] image_seq_len = image_embeds.shape[0] if image_seq_len != expected_tokens: if image_seq_len == expected_features: raise ValueError( "Vision output is pre-projector features. " f"got={image_seq_len}, expected_projected_tokens={expected_tokens}. " "Please re-export VIT ONNX with projector included (model_convert/export_onnx.py), " f"then re-compile to .axmodel. vit_output_shapes={vit_output_shapes}" ) raise ValueError( "Unexpected image feature length. " f"got={image_seq_len}, expected_projected_tokens={expected_tokens}, " f"expected_pre_projector_features={expected_features}, vit_output_shapes={vit_output_shapes}" ) projected_embeds = image_embeds prefill_data = np.take(self.embeds, token_ids, axis=0) prefill_data = _replace_image_tokens( token_ids, prefill_data, projected_embeds, image_token_id=self.config.image_token_id, ) prefill_data = prefill_data.astype(bfloat16) return token_ids, prefill_data def _stream_generate(self, token_ids: List[int], prefill_data: np.ndarray, max_new_tokens: int = 1024): for k_cache in self.infer_manager.k_caches: k_cache.fill(0) for v_cache in self.infer_manager.v_caches: v_cache.fill(0) eos_token_id = None if isinstance(self.config.eos_token_id, list) and len(self.config.eos_token_id) > 1: eos_token_id = self.config.eos_token_id slice_len = 128 t_start = time.time() token_ids = self.infer_manager.prefill(self.tokenizer, token_ids, prefill_data, slice_len=slice_len) mask = np.zeros((1, 1, self.infer_manager.max_seq_len + 1), dtype=np.float32).astype(bfloat16) mask[:, :, :self.infer_manager.max_seq_len] -= 65536 seq_len = len(token_ids) - 1 if slice_len > 0: mask[:, :, :seq_len] = 0 ttft_ms: Optional[float] = (time.time() - t_start) * 1000 decode_tokens = 0 decode_elapsed_ms: float = 0.0 generated_text = self.tokenizer.decode(token_ids[seq_len:], skip_special_tokens=True) yield generated_text, ttft_ms, None, 1, False remaining_decode_budget = max(0, int(max_new_tokens) - 1) for step_idx in range(self.infer_manager.max_seq_len): if remaining_decode_budget <= 0: break if slice_len > 0 and step_idx < seq_len: continue cur_token = token_ids[step_idx] indices = np.array([step_idx], np.uint32).reshape((1, 1)) data = self.embeds[cur_token, :].reshape((1, 1, self.config.hidden_size)).astype(bfloat16) for layer_idx in range(self.config.num_hidden_layers): input_feed = { "K_cache": self.infer_manager.k_caches[layer_idx], "V_cache": self.infer_manager.v_caches[layer_idx], "indices": indices, "input": data, "mask": mask, } outputs = self.infer_manager.decoder_sessions[layer_idx].run(None, input_feed, shape_group=0) self.infer_manager.k_caches[layer_idx][:, step_idx, :] = outputs[0][:, :, :] self.infer_manager.v_caches[layer_idx][:, step_idx, :] = outputs[1][:, :, :] data = outputs[2] mask[..., step_idx] = 0 if step_idx < seq_len - 1: continue post_out = self.infer_manager.post_process_session.run(None, {"input": data})[0] next_token, _, _ = self.infer_manager.post_process(post_out, temperature=0.7) if eos_token_id is not None and next_token in eos_token_id: break if next_token == self.tokenizer.eos_token_id: break token_ids.append(next_token) remaining_decode_budget -= 1 generated_text = self.tokenizer.decode(token_ids[seq_len:], skip_special_tokens=True) decode_tokens += 1 decode_elapsed_ms = (time.time() - t_start) * 1000 - ttft_ms avg_decode = decode_elapsed_ms / decode_tokens if decode_tokens > 0 else None total_tokens = 1 + decode_tokens yield generated_text, ttft_ms, avg_decode, total_tokens, False avg_decode = decode_elapsed_ms / decode_tokens if decode_tokens > 0 else None total_tokens = 1 + decode_tokens yield generated_text, ttft_ms, avg_decode, total_tokens, True def chat(self, user_input: str, image: Optional[Image.Image], task: str) -> Generator: if image is None: err = "请先上传图片,再执行识别。" metrics = ( "