from __future__ import annotations import argparse import json import os from pathlib import Path import numpy as np from utils.text_infer_func import TextInferManager from utils.text_infer_func import detect_prefill_len from utils.text_runtime_compat import build_text_inputs from utils.text_runtime_compat import load_text_runtime_config from utils.text_runtime_compat import load_tokenizer def format_vector_preview(array: np.ndarray, count: int = 8) -> list[float]: flat = array.reshape(-1) return [float(v) for v in flat[:count]] 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 compare_embeddings(reference: np.ndarray, output: np.ndarray) -> dict: diff = np.abs(reference - output) return { "reference_shape": list(reference.shape), "max_abs_diff": float(diff.max()), "mean_abs_diff": float(diff.mean()), "cosine_similarity": cosine_similarity(reference, output), } def load_embed_matrix(embed_file: Path, *, vocab_size: int, hidden_size: int) -> np.ndarray: if embed_file.suffix == ".npy": return np.load(embed_file).astype(np.float32) if embed_file.suffix == ".bin": raw = np.fromfile(embed_file, dtype=np.uint16) expected = vocab_size * hidden_size if raw.size != expected: raise ValueError(f"Unexpected embed bin size: got {raw.size}, expected {expected}") as_fp32 = (raw.astype(np.uint32) << 16).view(np.float32) return as_fp32.reshape(vocab_size, hidden_size) raise ValueError(f"Unsupported embedding file: {embed_file}") 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("--axmodel-path", required=True, type=Path) parser.add_argument("--text", required=True, type=str) parser.add_argument("--prompt-name", choices=["query", "document"], default="query") 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("--max-length", type=int, default=None) parser.add_argument("--reference-npy", type=Path) parser.add_argument("--save-npy", type=Path) args = parser.parse_args() config = load_text_runtime_config(args.hf_model) tokenizer = load_tokenizer(args.hf_model) max_length = args.max_length or int(config.max_position_embeddings) encoded = build_text_inputs( tokenizer, text=args.text, prompt_name=args.prompt_name, max_length=max_length, ) token_ids = encoded["input_ids"][0].tolist() with open(args.axmodel_path / "config.json", encoding="utf-8") as handle: runtime_config = json.load(handle) embed_file = args.embed_file or args.embed_npy or (args.axmodel_path / "model.embed_tokens.weight.bfloat16.bin") embed_matrix = load_embed_matrix( embed_file, vocab_size=int(runtime_config["tokens_embed_num"]), hidden_size=int(runtime_config["tokens_embed_size"]), ) prefill_data = np.take(embed_matrix, token_ids, axis=0) slice_len = args.slice_len or detect_prefill_len(str(args.axmodel_path), default=128) mask_mode = str(runtime_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(config, str(args.axmodel_path), max_seq_len=2047) try: output = manager.embed_text(token_ids, prefill_data, slice_len=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 if args.save_npy: args.save_npy.parent.mkdir(parents=True, exist_ok=True) np.save(args.save_npy, output) summary = { "hf_model": str(args.hf_model), "axmodel_path": str(args.axmodel_path), "embed_file": str(embed_file), "prompt_name": args.prompt_name, "text": args.text, "token_count": len(token_ids), "slice_len": slice_len, "mask_mode": mask_mode, "shape": list(output.shape), "dtype": str(output.dtype), "preview": format_vector_preview(output), "l2_norm": float(np.linalg.norm(output[0])), } if args.reference_npy: reference = np.load(args.reference_npy) summary["reference_compare"] = compare_embeddings(reference, output) print(json.dumps(summary, indent=2, ensure_ascii=False)) if __name__ == "__main__": main()