from __future__ import annotations import argparse import json import os from pathlib import Path import numpy as np from backend_runtime import create_backend_runner, load_runtime_inputs_npz from infer_axmodel_text import compare_embeddings, format_vector_preview, load_embed_matrix from utils.text_infer_func import TextInferManager, detect_prefill_len from utils.text_runtime_compat import load_text_runtime_config, load_tokenizer def load_multimodal_config(model_dir: Path) -> dict: with open(model_dir / "config.json", encoding="utf-8") as handle: return json.load(handle) def build_input_ids_for_multimodal_prompt( tokenizer, prompt: str, *, placeholder_token_id: int, placeholder_count: int, ) -> list[int]: encoded = tokenizer([prompt], return_tensors="np", padding=True, truncation=False) base_ids = encoded["input_ids"][0].tolist() current_count = sum(1 for token_id in base_ids if token_id == placeholder_token_id) if current_count == placeholder_count: return base_ids if current_count != 1: raise ValueError( f"Unexpected placeholder token count in prompt: got {current_count}, expected 1 or {placeholder_count}" ) expanded: list[int] = [] expanded_once = False for token_id in base_ids: if token_id == placeholder_token_id and not expanded_once: expanded.extend([placeholder_token_id] * placeholder_count) expanded_once = True continue expanded.append(token_id) return expanded def build_prefill_embeddings( embed_matrix: np.ndarray, input_ids: list[int], *, placeholder_token_id: int, replacement_tokens: np.ndarray, ) -> np.ndarray: prefill = np.take(embed_matrix, input_ids, axis=0).astype(np.float32, copy=False) flat_tokens = replacement_tokens.reshape(-1, replacement_tokens.shape[-1]).astype(np.float32, copy=False) positions = [index for index, token_id in enumerate(input_ids) if token_id == placeholder_token_id] if len(positions) != flat_tokens.shape[0]: raise ValueError( f"Replacement token count mismatch: prompt has {len(positions)} placeholders, " f"encoder produced {flat_tokens.shape[0]} tokens" ) for token_index, position in enumerate(positions): prefill[position, :] = flat_tokens[token_index] return prefill def run_multimodal_encoder_tokens( *, encoder_runner, mode: str, runtime_inputs: dict[str, np.ndarray], ) -> np.ndarray: if mode in {"vision_embedding", "audio_embedding"}: return encoder_runner.run(runtime_inputs) if mode == "video_embedding": pixel_values_frames = runtime_inputs["pixel_values_frames"] outputs = [] for frame in pixel_values_frames: output = encoder_runner.run({"pixel_values": np.ascontiguousarray(frame)}) outputs.append(np.array(output, copy=True)) if not outputs: raise ValueError("No video frames available for encoder inference") return np.concatenate(outputs, axis=1) raise ValueError(f"Unsupported mode: {mode}") def infer_embedding( *, hf_model: Path, llm_axmodel_path: Path, encoder_axmodel_path: Path, mode: str, prepared_inputs: Path, meta_json: Path, embed_file: Path, slice_len: int | None, ) -> tuple[np.ndarray, np.ndarray, list[int], dict]: if mode not in {"vision_embedding", "audio_embedding", "video_embedding"}: raise ValueError(f"Unsupported mode: {mode}") runtime_inputs = load_runtime_inputs_npz(prepared_inputs) meta = json.loads(meta_json.read_text(encoding="utf-8")) config = load_multimodal_config(hf_model) tokenizer = load_tokenizer(hf_model) encoder_runner = create_backend_runner("axmodel", encoder_axmodel_path) encoder_tokens = run_multimodal_encoder_tokens( encoder_runner=encoder_runner, mode=mode, runtime_inputs=runtime_inputs, ) placeholder_count = int(encoder_tokens.shape[1]) if mode == "vision_embedding": placeholder_token_id = int(config["image_token_id"]) elif mode == "audio_embedding": placeholder_token_id = int(config["audio_token_id"]) else: placeholder_token_id = int(config["video_token_id"]) prompt = meta["prompt_meta"]["prompt"] input_ids = build_input_ids_for_multimodal_prompt( tokenizer, prompt, placeholder_token_id=placeholder_token_id, placeholder_count=placeholder_count, ) embed_matrix = load_embed_matrix( embed_file, vocab_size=int(config["tokens_embed_num"]), hidden_size=int(config["tokens_embed_size"]), ) prefill = build_prefill_embeddings( embed_matrix, input_ids, placeholder_token_id=placeholder_token_id, replacement_tokens=encoder_tokens, ) text_config = load_text_runtime_config(hf_model) actual_slice_len = slice_len or detect_prefill_len(str(llm_axmodel_path), default=128) mask_mode = str(config.get("prefill_mask_mode", "causal")).strip().lower() old_mask_mode = os.environ.get("AXERA_TEXT_MASK_MODE") os.environ["AXERA_TEXT_MASK_MODE"] = mask_mode manager = TextInferManager(text_config, str(llm_axmodel_path), max_seq_len=2047) try: embedding = manager.embed_text(input_ids, prefill, slice_len=actual_slice_len) finally: manager.close() if old_mask_mode is None: os.environ.pop("AXERA_TEXT_MASK_MODE", None) else: os.environ["AXERA_TEXT_MASK_MODE"] = old_mask_mode extra = { "prompt": prompt, "placeholder_count": placeholder_count, "slice_len": actual_slice_len, "mask_mode": mask_mode, "runtime_input_shapes": {name: list(value.shape) for name, value in runtime_inputs.items()}, } return embedding, encoder_tokens, input_ids, extra def main() -> None: parser = argparse.ArgumentParser() parser.add_argument( "--hf-model", required=True, type=Path, help="Tokenizer/config/chat-template dir used at runtime; no safetensors required", ) parser.add_argument("--llm-axmodel-path", required=True, type=Path) parser.add_argument("--encoder-axmodel-path", required=True, type=Path) parser.add_argument("--mode", required=True, choices=["vision_embedding", "audio_embedding", "video_embedding"]) parser.add_argument("--prepared-inputs", required=True, type=Path) parser.add_argument("--meta-json", required=True, type=Path) parser.add_argument("--embed-file", type=Path, default=None) parser.add_argument("--embed-npy", type=Path, default=None, help=argparse.SUPPRESS) parser.add_argument("--slice-len", type=int, default=None) parser.add_argument("--reference-npy", type=Path) parser.add_argument("--reference-token-npy", type=Path) parser.add_argument("--save-npy", type=Path) parser.add_argument("--save-token-npy", type=Path) args = parser.parse_args() embed_file = args.embed_file or args.embed_npy or (args.llm_axmodel_path / "model.embed_tokens.weight.bfloat16.bin") embedding, encoder_tokens, input_ids, extra = infer_embedding( hf_model=args.hf_model, llm_axmodel_path=args.llm_axmodel_path, encoder_axmodel_path=args.encoder_axmodel_path, mode=args.mode, prepared_inputs=args.prepared_inputs, meta_json=args.meta_json, embed_file=embed_file, slice_len=args.slice_len, ) if args.save_npy: args.save_npy.parent.mkdir(parents=True, exist_ok=True) np.save(args.save_npy, embedding) if args.save_token_npy: args.save_token_npy.parent.mkdir(parents=True, exist_ok=True) np.save(args.save_token_npy, encoder_tokens) summary = { "hf_model": str(args.hf_model), "llm_axmodel_path": str(args.llm_axmodel_path), "encoder_axmodel_path": str(args.encoder_axmodel_path), "embed_file": str(embed_file), "mode": args.mode, "prepared_inputs": str(args.prepared_inputs), "meta_json": str(args.meta_json), "sequence_length": len(input_ids), "multimodal_token_count": int(encoder_tokens.shape[1]), "shape": list(embedding.shape), "dtype": str(embedding.dtype), "preview": format_vector_preview(embedding), "l2_norm": float(np.linalg.norm(embedding[0])), "runtime_input_shapes": extra["runtime_input_shapes"], "slice_len": extra["slice_len"], "mask_mode": extra["mask_mode"], } if args.reference_token_npy: reference_tokens = np.load(args.reference_token_npy) summary["reference_token_compare"] = compare_embeddings(reference_tokens, encoder_tokens) if args.reference_npy: reference = np.load(args.reference_npy) summary["reference_compare"] = compare_embeddings(reference, embedding) print(json.dumps(summary, indent=2, ensure_ascii=False)) if __name__ == "__main__": main()