# Qwen4 cache integration for oMLX. # Licensed under the Apache License 2.0. See LICENSE. from enum import Enum from typing import Any import mlx.core as mx from omlx.cache.type_handlers import ( CacheStateAxisInfo, CacheTypeHandler, ) from omlx.cache.type_registry import CacheTypeRegistry class Qwen4CacheType(Enum): QSA_KV = "QSAKVCache" QSA_QUANTIZED_KV = "QSAQuantizedKVCache" def _offset_from_meta(meta_state, fallback): if isinstance(meta_state, (list, tuple)) and meta_state: return int(meta_state[0]) if meta_state not in (None, ""): return int(meta_state) return fallback class QSAKVCacheHandler(CacheTypeHandler): @property def cache_type(self): return Qwen4CacheType.QSA_KV @property def supports_block_slicing(self): return True def get_state_axis_info(self): return ( CacheStateAxisInfo("keys", 2, True), CacheStateAxisInfo("values", 2, True), CacheStateAxisInfo("index_keys", 1, True), ) def serialize_state(self, cache_obj): keys, values, index_keys = cache_obj.state return keys, values, index_keys def serialize_meta_state(self, cache_obj): return (int(cache_obj.offset),) def extract_state(self, cache_obj): elements = self.serialize_state(cache_obj) return { "keys": elements[0], "values": elements[1], "index_keys": elements[2], "states": elements, "cache_type": self.cache_type.value, } def get_seq_len(self, state): keys = state.get("keys") if keys is not None: return int(keys.shape[2]) index_keys = state.get("index_keys") return 0 if index_keys is None else int(index_keys.shape[1]) def slice_state(self, state, start_idx, end_idx): keys = state.get("keys") values = state.get("values") index_keys = state.get("index_keys") if keys is None or values is None: return None end_idx = min(end_idx, int(keys.shape[2])) if start_idx >= end_idx: return None index_end = min(end_idx, int(index_keys.shape[1])) elements = ( keys[:, :, start_idx:end_idx, :], values[:, :, start_idx:end_idx, :], index_keys[:, start_idx:index_end, :], ) return { "keys": elements[0], "values": elements[1], "index_keys": elements[2], "states": elements, "cache_type": self.cache_type.value, } def concatenate_states(self, states): elements = [state.get("states") for state in states] elements = [value for value in elements if value] if not elements: return {} combined = ( mx.concatenate([value[0] for value in elements], axis=2), mx.concatenate([value[1] for value in elements], axis=2), mx.concatenate([value[2] for value in elements], axis=1), ) return { "keys": combined[0], "values": combined[1], "index_keys": combined[2], "states": combined, "cache_type": self.cache_type.value, } def deserialize_state(self, elements, meta_state=None): from mlx_lm.models.qwen4_exp import QSAKVCache keys = elements[0] if len(elements) > 0 else None values = elements[1] if len(elements) > 1 else None index_keys = elements[2] if len(elements) > 2 else None fallback = 0 if keys is None else int(keys.shape[2]) cache = QSAKVCache() cache.keys = keys cache.values = values cache.index_keys = index_keys cache.offset = _offset_from_meta(meta_state, fallback) return cache def reconstruct_cache(self, state, meta_state=None): elements = state.get("states") if elements is None: elements = ( state.get("keys"), state.get("values"), state.get("index_keys"), ) return self.deserialize_state(tuple(elements), meta_state) class QSAQuantizedKVCacheHandler(CacheTypeHandler): @property def cache_type(self): return Qwen4CacheType.QSA_QUANTIZED_KV @property def supports_block_slicing(self): return True def get_state_axis_info(self): return ( CacheStateAxisInfo("key_weight", 2, True), CacheStateAxisInfo("key_scales", 2, True), CacheStateAxisInfo("key_biases", 2, True), CacheStateAxisInfo("value_weight", 2, True), CacheStateAxisInfo("value_scales", 2, True), CacheStateAxisInfo("value_biases", 2, True), CacheStateAxisInfo("index_keys", 1, True), ) def serialize_state(self, cache_obj): if cache_obj.keys is None: return (None,) * 7 offset = int(cache_obj.offset) keys = tuple(value[:, :, :offset, :] for value in cache_obj.keys) values = tuple(value[:, :, :offset, :] for value in cache_obj.values) index_keys = cache_obj.index_keys if index_keys is not None: index_keys = index_keys[:, :offset, :] return (*keys, *values, index_keys) def serialize_meta_state(self, cache_obj): return ( int(cache_obj.offset), int(cache_obj.group_size), int(cache_obj.bits), ) def extract_state(self, cache_obj): elements = self.serialize_state(cache_obj) return { "states": elements, "keys": elements[0], "values": elements[3], "index_keys": elements[6], "cache_type": self.cache_type.value, } def get_seq_len(self, state): keys = state.get("keys") if keys is not None: return int(keys.shape[2]) index_keys = state.get("index_keys") return 0 if index_keys is None else int(index_keys.shape[1]) def slice_state(self, state, start_idx, end_idx): elements = state.get("states") if not elements or elements[0] is None: return None end_idx = min(end_idx, int(elements[0].shape[2])) if start_idx >= end_idx: return None sliced = tuple( value[:, :, start_idx:end_idx, :] for value in elements[:6] ) + (elements[6][:, start_idx:end_idx, :],) return { "states": sliced, "keys": sliced[0], "values": sliced[3], "index_keys": sliced[6], "cache_type": self.cache_type.value, } def concatenate_states(self, states): elements = [state.get("states") for state in states] elements = [value for value in elements if value and value[0] is not None] if not elements: return {} combined = tuple( mx.concatenate([value[index] for value in elements], axis=2) for index in range(6) ) + (mx.concatenate([value[6] for value in elements], axis=1),) return { "states": combined, "keys": combined[0], "values": combined[3], "index_keys": combined[6], "cache_type": self.cache_type.value, } def deserialize_state(self, elements, meta_state=None): from mlx_lm.models.qwen4_exp import QSAQuantizedKVCache offset = _offset_from_meta( meta_state, 0 if not elements or elements[0] is None else int(elements[0].shape[2]), ) group_size = int(meta_state[1]) if meta_state and len(meta_state) > 1 else 64 bits = int(meta_state[2]) if meta_state and len(meta_state) > 2 else 4 cache = QSAQuantizedKVCache(group_size=group_size, bits=bits) if elements and elements[0] is not None: cache.keys = tuple(elements[:3]) cache.values = tuple(elements[3:6]) cache.index_keys = elements[6] if len(elements) > 6 else None cache.offset = offset return cache def reconstruct_cache(self, state, meta_state=None): return self.deserialize_state(tuple(state.get("states") or ()), meta_state) def _batch_indices(batch_indices): if hasattr(batch_indices, "tolist"): return [int(value) for value in batch_indices.tolist()] return [int(value) for value in batch_indices] def _install_single_cache_batch_methods(cache_class): def filter_rows(self, batch_indices): indices = _batch_indices(batch_indices) if not indices: self.keys = None self.values = None self.index_keys = None self.offset = 0 return self.keys = _map_cache_arrays(self.keys, lambda value: value[indices]) self.values = _map_cache_arrays(self.values, lambda value: value[indices]) if self.index_keys is not None: self.index_keys = self.index_keys[indices] def extract_row(self, index): result = type(self).__new__(type(self)) result.keys = _map_cache_arrays(self.keys, lambda value: value[index : index + 1]) result.values = _map_cache_arrays(self.values, lambda value: value[index : index + 1]) result.index_keys = ( None if self.index_keys is None else self.index_keys[index : index + 1] ) result.offset = self.offset if hasattr(self, "group_size"): result.group_size = self.group_size result.bits = self.bits return result def extend_rows(self, other): if int(self.offset) != int(other.offset): raise ValueError("QSA caches can only batch rows at the same offset") if hasattr(self, "group_size") and ( self.group_size != other.group_size or self.bits != other.bits ): raise ValueError("quantized QSA caches must use the same layout") offset = int(self.offset) self.keys = _merge_cache_arrays(self.keys, other.keys, offset, 2) self.values = _merge_cache_arrays(self.values, other.values, offset, 2) self.index_keys = _merge_cache_arrays( self.index_keys, other.index_keys, offset, 1, ) @classmethod def merge_rows(cls, caches): caches = list(caches) if not caches: return cls() result = caches[0].extract(0) for cache in caches[1:]: result.extend(cache) return result cache_class.filter = filter_rows cache_class.extract = extract_row cache_class.extend = extend_rows cache_class.merge = merge_rows def _install_quantized_state_layout(cache_class): def get_state(self): if self.keys is None: return (None,) * 7 offset = int(self.offset) keys = tuple(value[:, :, :offset, :] for value in self.keys) values = tuple(value[:, :, :offset, :] for value in self.values) index_keys = self.index_keys if index_keys is not None: index_keys = index_keys[:, :offset, :] return (*keys, *values, index_keys) def set_state(self, value): if len(value) == 2 and isinstance(value[0], (list, tuple)): quantized_state, self.index_keys = value self.keys, self.values = quantized_state else: self.keys = tuple(value[:3]) if value and value[0] is not None else None self.values = tuple(value[3:6]) if value and value[3] is not None else None self.index_keys = value[6] if len(value) > 6 else None self.offset = 0 if self.keys is None else int(self.keys[0].shape[2]) cache_class.state = property(get_state, set_state) def _map_cache_arrays(value, function): if value is None: return None if isinstance(value, (list, tuple)): return tuple(function(item) for item in value) return function(value) def _merge_cache_arrays(left, right, length, sequence_axis): if left is None: return right if right is None: return left if isinstance(left, (list, tuple)): return tuple( _merge_cache_arrays(a, b, length, sequence_axis) for a, b in zip(left, right) ) slices = [slice(None)] * left.ndim slices[sequence_axis] = slice(0, length) slices = tuple(slices) return mx.concatenate([left[slices], right[slices]], axis=0) def register_qwen4_cache_integration(): import mlx_lm.models.cache as mlx_cache from mlx_lm.models.qwen4_exp import QSAKVCache, QSAQuantizedKVCache handlers = (QSAKVCacheHandler(), QSAQuantizedKVCacheHandler()) for handler in handlers: CacheTypeRegistry.register(handler) CacheTypeRegistry._class_name_map.update( { QSAKVCache.__name__: Qwen4CacheType.QSA_KV, QSAQuantizedKVCache.__name__: Qwen4CacheType.QSA_QUANTIZED_KV, } ) _install_single_cache_batch_methods(QSAKVCache) _install_single_cache_batch_methods(QSAQuantizedKVCache) _install_quantized_state_layout(QSAQuantizedKVCache) mlx_cache.QSAKVCache = QSAKVCache mlx_cache.QSAQuantizedKVCache = QSAQuantizedKVCache