| 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 |
|
|