yongqiang
Bootstrap nano AXERA package scaffold
df5639f
Raw
History Blame
5.07 kB
from __future__ import annotations
import argparse
import json
from pathlib import Path
import urllib.request
import numpy as np
PACKAGE_ROOT = Path(__file__).resolve().parent.parent
CASE_ROOT = PACKAGE_ROOT / "python" / "testdata" / "service_cases"
def post_json(url: str, payload: dict) -> dict:
data = json.dumps(payload, ensure_ascii=False).encode("utf-8")
req = urllib.request.Request(
url,
data=data,
headers={
"Content-Type": "application/json",
"Authorization": "Bearer not-needed",
},
method="POST",
)
with urllib.request.urlopen(req, timeout=600) as resp:
return json.loads(resp.read().decode("utf-8"))
def get_default_model(api_url: str) -> str:
with urllib.request.urlopen(api_url.rstrip("/") + "/models", timeout=60) as resp:
payload = json.loads(resp.read().decode("utf-8"))
data = payload.get("data") or []
if not data:
raise RuntimeError("No model found from /v1/models")
return str(data[0]["id"])
def cosine_similarity(lhs: np.ndarray, rhs: np.ndarray) -> float:
lhs64 = lhs.reshape(-1).astype(np.float64)
rhs64 = rhs.reshape(-1).astype(np.float64)
denom = (np.linalg.norm(lhs64) * np.linalg.norm(rhs64)) + 1e-12
return float(np.dot(lhs64, rhs64) / 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 build_request(model: str, case_meta: dict) -> dict:
modality = case_meta["modality"]
prompt_name = case_meta["prompt_name"]
if modality == "text":
return {
"model": model,
"input": case_meta["text"],
"prompt_name": prompt_name,
"encoding_format": "float",
}
asset_path = (PACKAGE_ROOT / case_meta["asset_path"]).resolve()
if modality == "image":
media_part = {"type": "image_url", "image_url": {"url": str(asset_path)}}
elif modality == "audio":
media_part = {"type": "audio_url", "audio_url": {"url": str(asset_path)}}
elif modality == "video":
media_part = {"type": "video_url", "video_url": {"url": str(asset_path)}}
else:
raise ValueError(f"Unsupported modality: {modality}")
return {
"model": model,
"prompt_name": prompt_name,
"encoding_format": "float",
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": case_meta.get("text_prefix", "")},
media_part,
],
}
],
}
def load_cases(explicit_cases: list[str]) -> list[Path]:
if explicit_cases:
return [CASE_ROOT / case_name for case_name in explicit_cases]
return sorted(path for path in CASE_ROOT.iterdir() if path.is_dir())
def main() -> None:
parser = argparse.ArgumentParser(description="Compare axllm /v1/embeddings outputs against packaged HF references")
parser.add_argument("--api-url", default="http://127.0.0.1:8000/v1", type=str)
parser.add_argument("--model", default=None, type=str)
parser.add_argument("--case", action="append", default=[], help="Repeatable case name under python/testdata/service_cases")
parser.add_argument("--save-summary", type=Path, default=None)
args = parser.parse_args()
model = args.model or get_default_model(args.api_url)
cases = load_cases(args.case)
summary_cases = []
for case_dir in cases:
meta = json.loads((case_dir / "meta.json").read_text(encoding="utf-8"))
reference = np.load(case_dir / "torch_embedding.npy").astype(np.float32)
payload = build_request(model, meta)
response = post_json(args.api_url.rstrip("/") + "/embeddings", payload)
output = np.asarray(response["data"][0]["embedding"], dtype=np.float32).reshape(1, -1)
comparison = compare_embeddings(reference, output)
result = {
"case_name": meta["case_name"],
"modality": meta["modality"],
"output_shape": list(output.shape),
"l2_norm": float(np.linalg.norm(output[0])),
"comparison": comparison,
}
if "soft_token_count" in meta:
result["soft_token_count"] = int(meta["soft_token_count"])
if "used_num_frames" in meta:
result["used_num_frames"] = int(meta["used_num_frames"])
summary_cases.append(result)
summary = {
"api_url": args.api_url,
"model": model,
"num_cases": len(summary_cases),
"cases": summary_cases,
}
if args.save_summary is not None:
args.save_summary.write_text(json.dumps(summary, indent=2, ensure_ascii=False), encoding="utf-8")
print(json.dumps(summary, indent=2, ensure_ascii=False))
if __name__ == "__main__":
main()