yongqiang
Prepare Jina omni nano retrieval package
a2cbb72
Raw
History Blame
11.5 kB
from __future__ import annotations
import atexit
import os
import re
import numpy as np
from ml_dtypes import bfloat16
from backend_runtime import import_axengine
def release_ax_inference_session(session):
inner = getattr(session, "_sess", None)
unload = getattr(inner, "_unload", None)
if not callable(unload):
return
try:
unload()
except Exception as exc:
print(f"[WARN] Failed to unload axengine session cleanly: {exc}")
finally:
try:
inner._unload = lambda: None
except Exception:
pass
def detect_prefill_len(model_dir: str, default: int = 128) -> int:
layer_pattern = re.compile(r"^.*_p(?P<prefill>\d+)_l\d+_together\.axmodel$")
prefill_counts: dict[int, int] = {}
try:
for fname in os.listdir(model_dir):
match = layer_pattern.match(fname)
if match:
prefill = int(match.group("prefill"))
prefill_counts[prefill] = prefill_counts.get(prefill, 0) + 1
except FileNotFoundError:
return default
if not prefill_counts:
return default
return max(prefill_counts.items(), key=lambda kv: kv[1])[0]
def _find_axmodel_files(base_dir: str, expected_prefill: int = 128):
files = os.listdir(base_dir)
layer_pattern = re.compile(r"^(?P<prefix>.*)_p(?P<prefill>\d+)_l(?P<idx>\d+)_together\.axmodel$")
post_pattern = re.compile(r"^(?P<prefix>.*)_post\.axmodel$")
prefix_map: dict[str, list[tuple[int, str]]] = {}
for fname in files:
match = layer_pattern.match(fname)
if match:
prefix = match.group("prefix")
idx = int(match.group("idx"))
prefix_map.setdefault(prefix, []).append((idx, fname))
if not prefix_map:
raise FileNotFoundError(f"No layer axmodel files found under: {base_dir}")
prefix = max(prefix_map.items(), key=lambda kv: len(kv[1]))[0]
layer_files = sorted(prefix_map[prefix], key=lambda item: item[0])
post_file = None
for fname in files:
match = post_pattern.match(fname)
if match and match.group("prefix") == prefix:
post_file = fname
break
if post_file is None:
candidate = os.path.join(base_dir, f"{prefix}_post.axmodel")
if os.path.exists(candidate):
post_file = os.path.basename(candidate)
if post_file is None:
raise FileNotFoundError(f"Cannot find post axmodel file under: {base_dir}")
return layer_files, post_file, prefix
class TextInferManager:
def __init__(self, config, model_dir: str, max_seq_len: int = 2047):
self.axengine = import_axengine()
self.config = config
self.max_seq_len = int(max_seq_len)
self.hidden_size = int(config.hidden_size)
self.num_hidden_layers = int(config.num_hidden_layers)
self.num_key_value_heads = int(config.num_key_value_heads)
self.head_dim = int(config.head_dim or (self.hidden_size // int(config.num_attention_heads)))
self.kv_dim = self.head_dim * self.num_key_value_heads
self.use_mrope = False
layer_files, post_file, _ = _find_axmodel_files(model_dir)
self.decoder_sessions = [
self.axengine.InferenceSession(os.path.join(model_dir, fname))
for _, fname in layer_files
]
self.post_process_session = self.axengine.InferenceSession(os.path.join(model_dir, post_file))
self.decode_cache_lens = [self._decode_cache_len(session) for session in self.decoder_sessions]
self.cache_len = max(self.decode_cache_lens, default=self.max_seq_len + 1)
self._closed = False
self.mask_mode = os.environ.get("AXERA_TEXT_MASK_MODE", "causal").strip().lower()
atexit.register(self.close)
def close(self):
if self._closed:
return
sessions = list(self.decoder_sessions) + [self.post_process_session]
for session in sessions:
release_ax_inference_session(session)
self._closed = True
@staticmethod
def _session_output_names(session):
try:
return tuple(output.name for output in session.get_outputs())
except Exception:
return ()
@staticmethod
def _session_input_names(session):
try:
return tuple(input_meta.name for input_meta in session.get_inputs())
except Exception:
return ()
@staticmethod
def _session_input_shapes(session):
try:
return {input_meta.name: tuple(input_meta.shape) for input_meta in session.get_inputs()}
except Exception:
return {}
def _decode_cache_len(self, session):
input_shapes = self._session_input_shapes(session)
k_shape = input_shapes.get("K_cache")
if k_shape is not None and len(k_shape) >= 2 and k_shape[1] is not None:
return int(k_shape[1])
return self.max_seq_len + 1
def _decoder_output_names(self, session, shape_group: int):
available_names = self._session_output_names(session)
base_names = ("K_cache_out", "V_cache_out", "output")
if shape_group == 0:
return base_names
grouped_names = (
f"K_cache_out_{shape_group}",
f"V_cache_out_{shape_group}",
f"output_{shape_group}",
)
if all(name in available_names for name in grouped_names):
return grouped_names
return base_names
def _decoder_input_name_map(self, session, shape_group: int):
available_names = set(self._session_input_names(session))
logical_names = ["K_cache", "V_cache", "indices", "input", "mask"]
mapped_names = {}
for logical_name in logical_names:
grouped_name = f"{logical_name}_{shape_group}" if shape_group != 0 else logical_name
if grouped_name in available_names:
mapped_names[logical_name] = grouped_name
elif logical_name in available_names:
mapped_names[logical_name] = logical_name
return mapped_names
def _prepare_decoder_input(self, session, input_feed, shape_group: int):
name_map = self._decoder_input_name_map(session, shape_group)
return {name_map[key]: value for key, value in input_feed.items() if key in name_map}
def _run_decoder(self, session, input_feed, shape_group: int):
names = self._decoder_output_names(session, shape_group)
outputs = None
try:
outputs = session.run(list(names), input_feed, shape_group=shape_group)
except TypeError:
try:
outputs = session.run(list(names), input_feed, shape_group)
except TypeError:
outputs = session.run(None, input_feed, shape_group=shape_group)
if isinstance(outputs, dict):
return outputs[names[0]], outputs[names[1]], outputs[names[2]]
if isinstance(outputs, (list, tuple)):
if len(outputs) == 3:
return outputs[0], outputs[1], outputs[2]
offset = shape_group * 3
if len(outputs) >= offset + 3:
return outputs[offset], outputs[offset + 1], outputs[offset + 2]
return outputs[0], outputs[1], outputs[2]
return outputs[0], outputs[1], outputs[2]
def _run_post(self, hidden: np.ndarray) -> np.ndarray:
output_names = self._session_output_names(self.post_process_session)
outputs = self.post_process_session.run(None, {"input": hidden})
if isinstance(outputs, dict):
named_outputs = {name: np.array(value, copy=True) for name, value in outputs.items()}
else:
named_outputs = {
name: np.array(value, copy=True)
for name, value in zip(output_names, outputs)
}
if "output_norm" in named_outputs:
return named_outputs["output_norm"]
for value in named_outputs.values():
if value.ndim >= 2 and value.shape[-1] == self.hidden_size:
return value
raise ValueError(f"Cannot find output_norm-like tensor in post outputs: {list(named_outputs.keys())}")
def embed_text(self, token_ids: list[int], embed_data: np.ndarray, slice_len: int) -> np.ndarray:
seq_len = len(token_ids)
slice_indices = [i for i in range(seq_len // slice_len + 1)]
k_caches = [np.zeros((1, self.cache_len, self.kv_dim), dtype=bfloat16) for _ in range(self.num_hidden_layers)]
v_caches = [np.zeros((1, self.cache_len, self.kv_dim), dtype=bfloat16) for _ in range(self.num_hidden_layers)]
final_hidden = None
for slice_idx in slice_indices:
base_indices = np.arange(slice_idx * slice_len, (slice_idx + 1) * slice_len, dtype=np.uint32)
indices = base_indices.reshape(1, -1)
mask = np.zeros((1, slice_len, slice_len * (slice_idx + 1)), dtype=np.float32) - 65536
data = np.zeros((1, slice_len, self.hidden_size), dtype=bfloat16)
for i, token_pos in enumerate(range(slice_idx * slice_len, (slice_idx + 1) * slice_len)):
if token_pos < seq_len:
if self.mask_mode == "bidirectional":
mask[:, i, :seq_len] = 0
else:
mask[:, i, : slice_idx * slice_len + i + 1] = 0
token_embed = np.asarray(embed_data[token_pos], dtype=np.float32)
data[:, i : i + 1, :] = token_embed.reshape((1, 1, self.hidden_size)).astype(bfloat16)
remain_len = seq_len - slice_idx * slice_len if slice_idx == slice_indices[-1] else slice_len
mask = mask.astype(bfloat16)
for layer_idx in range(self.num_hidden_layers):
if slice_idx:
k_cache = k_caches[layer_idx][:, : slice_len * slice_idx, :]
v_cache = v_caches[layer_idx][:, : slice_len * slice_idx, :]
else:
k_cache = np.zeros((1, 1, self.kv_dim), dtype=bfloat16)
v_cache = np.zeros((1, 1, self.kv_dim), dtype=bfloat16)
input_feed = {
"K_cache": k_cache,
"V_cache": v_cache,
"indices": indices,
"input": data,
"mask": mask,
}
input_feed = self._prepare_decoder_input(self.decoder_sessions[layer_idx], input_feed, shape_group=slice_idx + 1)
k_out, v_out, data = self._run_decoder(self.decoder_sessions[layer_idx], input_feed, shape_group=slice_idx + 1)
k_caches[layer_idx][:, slice_idx * slice_len : slice_idx * slice_len + remain_len, :] = k_out[:, :remain_len, :]
v_caches[layer_idx][:, slice_idx * slice_len : slice_idx * slice_len + remain_len, :] = v_out[:, :remain_len, :]
if slice_idx == slice_indices[-1]:
last_local_index = (seq_len - 1) - (slice_idx * slice_len)
final_hidden = data[:, last_local_index : last_local_index + 1, :]
if final_hidden is None:
raise ValueError("No final hidden state produced during prefill")
post_hidden = self._run_post(final_hidden)
vector = np.asarray(post_hidden[:, 0, :], dtype=np.float32)
norm = np.linalg.norm(vector, axis=-1, keepdims=True) + 1e-12
return vector / norm