from __future__ import annotations import ctypes.util import importlib import os import sys from dataclasses import dataclass from pathlib import Path from typing import Any import numpy as np try: import onnxruntime as ort except ImportError: ort = None @dataclass class PreparedBackendInputs: mode: str runtime_inputs: dict[str, np.ndarray] torch_inputs: dict prompt_meta: dict class BackendRunner: backend_name = "unknown" def run(self, runtime_inputs: dict[str, np.ndarray]) -> np.ndarray: # pragma: no cover - interface only raise NotImplementedError class OnnxBackendRunner(BackendRunner): backend_name = "onnx" def __init__(self, artifact_path: str | Path, providers: list[str] | None = None): if ort is None: raise ImportError("onnxruntime is required for ONNX backend inference") self.artifact_path = Path(artifact_path) self.session = ort.InferenceSession( str(self.artifact_path), providers=providers or ["CPUExecutionProvider"], ) def run(self, runtime_inputs: dict[str, np.ndarray]) -> np.ndarray: return self.session.run(None, runtime_inputs)[0] class AxModelBackendRunner(BackendRunner): backend_name = "axmodel" def __init__(self, artifact_path: str | Path): self.artifact_path = Path(artifact_path) self.axengine = import_axengine() self.session = self.axengine.InferenceSession( str(self.artifact_path), providers=[self.axengine.axengine_provider_name], ) def run(self, runtime_inputs: dict[str, np.ndarray]) -> np.ndarray: outputs = self.session.run(None, runtime_inputs) return select_backend_output(outputs) def create_backend_runner(backend: str, artifact_path: str | Path) -> BackendRunner: if backend == "onnx": return OnnxBackendRunner(artifact_path) if backend == "axmodel": return AxModelBackendRunner(artifact_path) raise ValueError(f"Unsupported backend: {backend}") def _candidate_axengine_roots() -> list[Path]: roots: list[Path] = [] for env_name in ("AXENGINE_PYTHON_ROOT", "PYAXENGINE_ROOT"): value = os.environ.get(env_name) if value: roots.append(Path(value)) roots.extend([Path("/root/yongqiang/pyaxengine"), Path("/root/yongqiang/lhj/pyaxengine")]) return roots def _candidate_ax_libraries() -> dict[str, str]: mapping = { "ax_engine": os.environ.get("AXENGINE_LIBAX_ENGINE", "/soc/lib/libax_engine.so"), "ax_sys": os.environ.get("AXENGINE_LIBAX_SYS", "/soc/lib/libax_sys.so"), "axcl_rt": os.environ.get("AXENGINE_LIBAXCL_RT", ""), } return {name: path for name, path in mapping.items() if path} def _ensure_ax_library_path() -> None: if os.environ.get("AXENGINE_LD_LIBRARY_PATH_READY") == "1": return candidate_dirs = ["/soc/lib", "/mnt/oss/npu-ci/ax650a/lib"] existing = [entry for entry in os.environ.get("LD_LIBRARY_PATH", "").split(":") if entry] prepend = [entry for entry in candidate_dirs if Path(entry).is_dir() and entry not in existing] if not prepend: return env = os.environ.copy() env["LD_LIBRARY_PATH"] = ":".join(prepend + existing) env["AXENGINE_LD_LIBRARY_PATH_READY"] = "1" os.execvpe(sys.executable, [sys.executable, *sys.argv], env) def import_axengine(): _ensure_ax_library_path() for root in reversed(_candidate_axengine_roots()): if root.is_dir(): root_str = str(root) if root_str not in sys.path: sys.path.insert(0, root_str) library_mapping = _candidate_ax_libraries() original_find_library = getattr(ctypes.util, "_axera_original_find_library", ctypes.util.find_library) def patched_find_library(name: str) -> str | None: mapped = library_mapping.get(name) if mapped and Path(mapped).is_file(): return mapped return original_find_library(name) ctypes.util._axera_original_find_library = original_find_library ctypes.util.find_library = patched_find_library return importlib.import_module("axengine") def select_backend_output(outputs: list[np.ndarray], expected_last_dim: int | None = None) -> np.ndarray: if not outputs: raise ValueError("Backend returned no outputs") if len(outputs) == 1: return outputs[0] if expected_last_dim is not None: for output in outputs: if output.ndim >= 2 and output.shape[-1] == expected_last_dim: return output ranked = sorted( outputs, key=lambda output: (int(output.ndim >= 2), output.shape[-1] if output.ndim >= 1 else -1), reverse=True, ) return ranked[0] def load_runtime_inputs_npz(path: str | Path) -> dict[str, np.ndarray]: path = Path(path) with np.load(path, allow_pickle=False) as data: return {key: np.ascontiguousarray(data[key]) for key in data.files} def create_prepared_inputs_from_runtime_inputs( mode: str, runtime_inputs: dict[str, np.ndarray], *, prompt_meta: dict[str, Any] | None = None, ) -> PreparedBackendInputs: return PreparedBackendInputs( mode=mode, runtime_inputs=runtime_inputs, torch_inputs={}, prompt_meta=prompt_meta or {}, ) def prepare_backend_inputs( model_dir: str | Path, mode: str, *, image_path: str | Path | None = None, audio_path: str | Path | None = None, video_path: str | Path | None = None, width: int = 448, height: int = 448, prompt_name: str = "query", num_frames: int = 8, max_prefill_tokens: int | None = None, ) -> PreparedBackendInputs: from jina_omni_utils import ( load_processor, prepare_fixed_audio_inputs, prepare_fixed_image_inputs, prepare_fixed_video_inputs, ) model_dir = Path(model_dir) if mode == "vision_tokens": if image_path is None: raise ValueError("--image-path is required for vision_tokens") processor = load_processor(model_dir, pixel_budget=width * height) torch_inputs, prompt_meta = prepare_fixed_image_inputs( model_dir, image_path=image_path, width=width, height=height, prompt_name=prompt_name, processor=processor, ) runtime_inputs = {"pixel_values": torch_inputs["pixel_values"].cpu().numpy().astype(np.float32)} return PreparedBackendInputs( mode=mode, runtime_inputs=runtime_inputs, torch_inputs=torch_inputs, prompt_meta=prompt_meta, ) if mode == "audio_tokens": if audio_path is None: raise ValueError("--audio-path is required for audio_tokens") processor = load_processor(model_dir) torch_inputs, prompt_meta = prepare_fixed_audio_inputs( model_dir, audio_path=audio_path, prompt_name=prompt_name, processor=processor, ) runtime_inputs = {"input_features": torch_inputs["input_features"].cpu().numpy().astype(np.float32)} return PreparedBackendInputs( mode=mode, runtime_inputs=runtime_inputs, torch_inputs=torch_inputs, prompt_meta=prompt_meta, ) if mode == "video_tokens": if video_path is None: raise ValueError("--video-path is required for video_tokens") processor = load_processor(model_dir, pixel_budget=width * height) torch_inputs, prompt_meta = prepare_fixed_video_inputs( model_dir, video_path=video_path, num_frames=num_frames, width=width, height=height, prompt_name=prompt_name, max_prefill_tokens=max_prefill_tokens, processor=processor, ) runtime_inputs = { "pixel_values_frames": torch_inputs["pixel_values_frames"].cpu().numpy().astype(np.float32), } return PreparedBackendInputs( mode=mode, runtime_inputs=runtime_inputs, torch_inputs=torch_inputs, prompt_meta=prompt_meta, ) raise ValueError(f"Unsupported mode: {mode}") def run_backend_runner_on_prepared(runner: BackendRunner, prepared: PreparedBackendInputs) -> np.ndarray: if prepared.mode != "video_tokens": return runner.run(prepared.runtime_inputs) pixel_values_frames = prepared.runtime_inputs["pixel_values_frames"] outputs: list[np.ndarray] = [] for frame in pixel_values_frames: output = runner.run({"pixel_values": np.ascontiguousarray(frame)}) outputs.append(np.array(output, copy=True)) if not outputs: raise ValueError("No video frames provided for backend inference") return np.concatenate(outputs, axis=1) def run_torch_reference( model_dir: str | Path, prepared: PreparedBackendInputs, *, task: str = "retrieval", device: str | None = None, dtype_name: str = "float32", merge_task_lora: bool = False, ) -> np.ndarray: import torch from jina_omni_utils import ( get_vision_tokens_from_processed, load_base_model, move_tensor_inputs, to_numpy, ) model_dir = Path(model_dir) if prepared.mode == "vision_tokens": model = load_base_model( model_dir, modality="vision", task=task, dtype_name=dtype_name, merge_task_lora=merge_task_lora, device=device, attn_implementation="eager", ) torch_inputs = move_tensor_inputs(prepared.torch_inputs, next(model.parameters()).device) output = get_vision_tokens_from_processed( model, torch_inputs["pixel_values"], torch_inputs["image_grid_thw"], ) return to_numpy(output) if prepared.mode == "audio_tokens": model = load_base_model( model_dir, modality="audio", task=task, dtype_name=dtype_name, merge_task_lora=merge_task_lora, device=device, attn_implementation="eager", ) torch_inputs = move_tensor_inputs(prepared.torch_inputs, next(model.parameters()).device) output = model.get_audio_features(torch_inputs["input_features"]).unsqueeze(0) return to_numpy(output) if prepared.mode == "video_tokens": model = load_base_model( model_dir, modality="vision", task=task, dtype_name=dtype_name, merge_task_lora=merge_task_lora, device=device, attn_implementation="eager", ) torch_inputs = move_tensor_inputs(prepared.torch_inputs, next(model.parameters()).device) frame_grid_thw = torch_inputs["image_grid_thw"] outputs = [ get_vision_tokens_from_processed(model, frame_pixel_values, frame_grid_thw) for frame_pixel_values in torch_inputs["pixel_values_frames"] ] return to_numpy(torch.cat(outputs, dim=1)) raise ValueError(f"Unsupported mode: {prepared.mode}") def summarize_backend_output( *, backend: str, artifact_path: str | Path, prepared: PreparedBackendInputs, output: np.ndarray, ) -> dict: prompt_meta = dict(prepared.prompt_meta) prompt = prompt_meta.get("prompt") if isinstance(prompt, str): prompt_meta["prompt_preview"] = prompt[:120] prompt_meta["prompt_length"] = len(prompt) del prompt_meta["prompt"] return { "backend": backend, "artifact_path": str(artifact_path), "mode": prepared.mode, "inputs": {name: list(value.shape) for name, value in prepared.runtime_inputs.items()}, "shape": list(output.shape), "dtype": str(output.dtype), "preview": format_vector_preview(output), "prompt_meta": prompt_meta, } def compare_backend_to_torch(torch_output: np.ndarray, backend_output: np.ndarray) -> dict: diff = np.abs(torch_output - backend_output) return { "torch_shape": list(torch_output.shape), "max_abs_diff": float(diff.max()), "mean_abs_diff": float(diff.mean()), "cosine_similarity": cosine_similarity(torch_output, backend_output), } def cosine_similarity(a: np.ndarray, b: np.ndarray) -> float: lhs = a.reshape(-1).astype(np.float64) rhs = b.reshape(-1).astype(np.float64) denom = (np.linalg.norm(lhs) * np.linalg.norm(rhs)) + 1e-12 return float(np.dot(lhs, rhs) / denom) def format_vector_preview(array: np.ndarray, count: int = 8) -> list[float]: flat = array.reshape(-1) return [float(v) for v in flat[:count]]