""" Low-level state I/O — reads/writes the MEMB file format using llama-cpp-python's exposed C functions directly. No ABI conflict: we reuse the libllama.so that llama-cpp-python already loaded instead of linking our own copy. """ from __future__ import annotations import ctypes import hashlib import struct import zlib from pathlib import Path from typing import TYPE_CHECKING if TYPE_CHECKING: from llama_cpp import Llama # ── File format ─────────────────────────────────────────────────── MAGIC = b"MEMB" VERSION = 1 MODEL_ID_LEN = 64 # little-endian: 4s magic | I version | 64s model_id | I n_ctx | I llama_ver | Q data_size _HEADER_FMT = "<4sI64sIIQ" _HEADER_SIZE = struct.calcsize(_HEADER_FMT) # == 88 # Optional Python trailer (appended after the V1 CRC32). The C library reads # only up to data_size+CRC and ignores anything past it, so adding a trailer # keeps the file C-compatible. The trailer carries llama-cpp-python's # wrapper state (n_tokens) which is needed for `Llama.eval()` to append new # tokens at the correct position instead of overwriting from offset 0. _TRAILER_MAGIC = b"MTRL" _TRAILER_FMT = "<4sII" # magic | version | n_tokens _TRAILER_SIZE = struct.calcsize(_TRAILER_FMT) # == 12 _TRAILER_VERSION = 1 # ── Helpers ─────────────────────────────────────────────────────── def compute_model_id(model_path: str) -> bytes: """SHA-256 of first 1 KiB of the GGUF file, hex-encoded and zero-padded to 64 bytes.""" h = hashlib.sha256() try: with open(model_path, "rb") as f: h.update(f.read(1024)) except OSError: h.update(model_path.encode("utf-8", errors="replace")) digest = h.hexdigest().encode("ascii") # 64 chars exactly for SHA-256 hex return digest.ljust(MODEL_ID_LEN, b"\x00")[:MODEL_ID_LEN] def _lib_and_ctx(llama_model: "Llama"): """Return (llama_cpp low-level module, raw llama_context_p pointer). Modern llama-cpp-python (≥0.3) wraps the raw `llama_context*` in a `_LlamaContext` helper object — the actual ctypes pointer lives one attribute deeper. Older versions exposed the pointer directly. """ try: from llama_cpp import llama_cpp as lib except ImportError as e: raise ImportError("llama-cpp-python is not installed: pip install llama-cpp-python") from e wrapper = getattr(llama_model, "_ctx", None) or getattr(llama_model, "ctx", None) if wrapper is None: raise RuntimeError( "Cannot find _ctx attribute on Llama instance — unsupported llama-cpp-python version" ) # If wrapper is already a ctypes pointer (older llama-cpp-python), use directly if isinstance(wrapper, (int, ctypes.c_void_p)) or hasattr(wrapper, "_type_"): return lib, wrapper # Otherwise unwrap one level: `_LlamaContext.ctx` holds the raw pointer # (may be an int address, ctypes pointer, or c_void_p depending on version) for attr in ("ctx", "context", "_ctx"): raw = getattr(wrapper, attr, None) if raw is not None and (isinstance(raw, (int, ctypes.c_void_p)) or hasattr(raw, "_type_")): return lib, raw raise RuntimeError( f"Could not extract raw llama_context from {type(wrapper).__name__}; " "your llama-cpp-python version may have changed its internals." ) def _state_get_size(lib, ctx) -> int: try: return lib.llama_state_get_size(ctx) # post-2024 API except AttributeError: return lib.llama_get_state_size(ctx) # VERIFY: older API fallback def _state_get_data(lib, ctx, buf: ctypes.Array, size: int) -> int: try: return lib.llama_state_get_data(ctx, buf, size) # post-2024 API except AttributeError: return lib.llama_copy_state_data(ctx, buf) # VERIFY: older API fallback def _state_set_data(lib, ctx, buf: ctypes.Array, size: int) -> int: try: return lib.llama_state_set_data(ctx, buf, size) # post-2024 API except AttributeError: return lib.llama_set_state_data(ctx, buf) # VERIFY: older API fallback def _n_ctx(lib, ctx) -> int: try: return lib.llama_n_ctx(ctx) except AttributeError: return 0 # ── Public API ──────────────────────────────────────────────────── def get_state_size(llama_model: "Llama") -> int: """Return the byte size of the current SSM hidden state.""" lib, ctx = _lib_and_ctx(llama_model) return _state_get_size(lib, ctx) def save_state(llama_model: "Llama", model_path: str, file_path: str) -> None: """Serialise the current SSM state to *file_path* in MEMB format.""" lib, ctx = _lib_and_ctx(llama_model) size = _state_get_size(lib, ctx) if size == 0: raise RuntimeError("llama context returned state size 0 — is it initialised?") buf = (ctypes.c_uint8 * size)() written = _state_get_data(lib, ctx, buf, size) if written == 0: raise RuntimeError("llama_state_get_data returned 0 bytes") data = bytes(buf[:written]) crc = zlib.crc32(data) & 0xFFFFFFFF model_id = compute_model_id(model_path) n_ctx = _n_ctx(lib, ctx) header = struct.pack( _HEADER_FMT, MAGIC, VERSION, model_id, n_ctx, 0, # llama_ver — reserved len(data), ) # Python-only trailer with Llama wrapper position so eval() resumes # at the right offset instead of overwriting from 0 after load_state. n_tokens = getattr(llama_model, "n_tokens", 0) trailer = struct.pack(_TRAILER_FMT, _TRAILER_MAGIC, _TRAILER_VERSION, n_tokens) out = Path(file_path) out.parent.mkdir(parents=True, exist_ok=True) out.write_bytes(header + data + struct.pack(" None: """Restore SSM state from *file_path* into *llama_model*.""" lib, ctx = _lib_and_ctx(llama_model) expected_model_id = compute_model_id(model_path) raw = Path(file_path).read_bytes() if len(raw) < _HEADER_SIZE + 4: raise ValueError(f"State file too small: {file_path}") magic, version, file_model_id, n_ctx, llama_ver, data_size = struct.unpack_from( _HEADER_FMT, raw, 0 ) if magic != MAGIC: raise ValueError(f"Bad magic: expected {MAGIC!r}, got {magic!r}") if version != VERSION: raise ValueError(f"Unsupported state version {version} (expected {VERSION})") if file_model_id != expected_model_id: raise ValueError( "Model identity mismatch — state file was created with a different GGUF.\n" f" file model_id : {file_model_id.rstrip(b'\\x00').decode()}\n" f" current model : {expected_model_id.rstrip(b'\\x00').decode()}" ) offset = _HEADER_SIZE if len(raw) < offset + data_size + 4: raise ValueError("Truncated state file") data = raw[offset : offset + data_size] stored_crc = struct.unpack_from("= trailer_off + _TRAILER_SIZE: magic, version, n_tokens = struct.unpack_from(_TRAILER_FMT, raw, trailer_off) if magic == _TRAILER_MAGIC and version == _TRAILER_VERSION: # Restore Llama wrapper state so eval() appends at the right # offset instead of overwriting the loaded KV-cache from 0. try: llama_model.n_tokens = n_tokens except AttributeError: pass # wrapper version doesn't expose .n_tokens — eval may misbehave