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\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.*)_p(?P\d+)_l(?P\d+)_together\.axmodel$") post_pattern = re.compile(r"^(?P.*)_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