Fix cross-process recall: MEMB trailer + raw eval/sample in chat()
Two bugs were blocking memba's main promise (load .memb in a fresh process → model continues with full recalled context): 1. llama_state_set_data() restores the C-level KV-cache + SSM hidden state, but llama-cpp-python's Python wrapper still reports n_tokens=0. The next eval() then decodes new tokens at offset 0 and overwrites the loaded state. Fix: extend MEMB format with an optional 12-byte trailer appended after the CRC32. It carries the wrapper's n_tokens. The C library reads up to CRC and ignores anything past it, so files stay backward-compatible with libmemba; only the Python loader uses it. 2. Llama.__call__ / create_chat_completion / generate all re-tokenise the prompt on every call and clear the KV-cache when the new tokens don't prefix-match input_ids. That destroys any state we just loaded. Fix: rewrite Session.chat() to use raw tokenize → eval → sample. eval() appends tokens to the live state without resetting, and we handle stop-token detection ourselves. Verified end-to-end on Nemotron-3-Nano-4B (hybrid 21x Mamba-2 + 4x attention) — see experiments/README.md for the full findings log. diag_session_nemotron.py, mood_batch_poc.py, mood_stream_poc.py and recall_poc.py now all pass their cross-process tests; Falcon-Mamba still fails because the trained model itself can't do cross-turn recall — that was the original misdiagnosis. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
+52
-11
@@ -83,25 +83,66 @@ class Session:
|
||||
|
||||
# ── Public methods ─────────────────────────────────────────────
|
||||
|
||||
def chat(self, prompt: str, max_tokens: int = 512) -> str:
|
||||
def chat(self, prompt: str, max_tokens: int = 512,
|
||||
temperature: float = 0.1, top_k: int = 1) -> str:
|
||||
"""
|
||||
Feed *prompt* to the model (wrapped in the active chat format) and
|
||||
return the generated text.
|
||||
|
||||
For SSM models the hidden state accumulates in llama_context across
|
||||
calls — there is no explicit message history list, the recurrent
|
||||
state IS the memory. Call save() at any checkpoint you want to
|
||||
resume from later.
|
||||
This uses raw tokenize+eval+sample rather than ``Llama.__call__`` or
|
||||
``create_chat_completion`` because those high-level helpers re-tokenise
|
||||
the entire prompt and then reset the KV-cache when the prefix does not
|
||||
match — which destroys any state loaded from a .memb file. Going
|
||||
through ``llama.eval()`` directly *appends* new tokens to the live
|
||||
state, which is exactly what memba needs.
|
||||
|
||||
For SSM/hybrid models the hidden state accumulates in llama_context
|
||||
across calls — there is no explicit message history list, the recurrent
|
||||
state IS the memory. Call save() at any checkpoint you want to resume
|
||||
from later.
|
||||
"""
|
||||
template, stops = CHAT_FORMATS[self._chat_format]
|
||||
wrapped = template.format(prompt=prompt)
|
||||
result = self._llama(
|
||||
wrapped,
|
||||
max_tokens=max_tokens,
|
||||
echo=False,
|
||||
stop=stops,
|
||||
|
||||
# Tokenise the wrapped turn. Add BOS only when the context is fresh
|
||||
# (no prior tokens in either the live conversation or a loaded state).
|
||||
is_fresh = self._llama.n_tokens == 0
|
||||
new_tokens = self._llama.tokenize(
|
||||
wrapped.encode("utf-8"),
|
||||
add_bos=is_fresh,
|
||||
special=True,
|
||||
)
|
||||
return result["choices"][0]["text"].strip()
|
||||
|
||||
# Append to the live state. eval() does NOT reset the KV-cache.
|
||||
self._llama.eval(new_tokens)
|
||||
|
||||
# Build the set of single-token stops + add EOS.
|
||||
stop_token_ids = {self._llama.token_eos()}
|
||||
for s in stops:
|
||||
for t in self._llama.tokenize(s.encode("utf-8"), add_bos=False, special=True):
|
||||
stop_token_ids.add(t)
|
||||
|
||||
out_tokens: list[int] = []
|
||||
for _ in range(max_tokens):
|
||||
tok = self._llama.sample(top_k=top_k, temp=temperature)
|
||||
out_tokens.append(tok)
|
||||
# Always eval the sampled token so it lives in the state too —
|
||||
# that way the next chat() turn sees the assistant reply as
|
||||
# part of the conversation.
|
||||
self._llama.eval([tok])
|
||||
if tok in stop_token_ids:
|
||||
break
|
||||
# Multi-token stop-string check (some stops span several BPE pieces)
|
||||
if stops:
|
||||
snippet = self._llama.detokenize(out_tokens).decode("utf-8", errors="ignore")
|
||||
if any(s in snippet for s in stops):
|
||||
break
|
||||
|
||||
text = self._llama.detokenize(out_tokens).decode("utf-8", errors="ignore")
|
||||
for s in stops:
|
||||
if s in text:
|
||||
text = text.split(s)[0]
|
||||
return text.strip()
|
||||
|
||||
def save(self, session_id: Optional[str] = None) -> Path:
|
||||
"""Persist the current state. Returns the path written."""
|
||||
|
||||
Reference in New Issue
Block a user