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:
@@ -0,0 +1,91 @@
|
||||
"""
|
||||
diag_nemotron3.py — test hypothesis that llama-cpp-python's built-in
|
||||
save_state/load_state preserves Python-side trackers (n_tokens, input_ids)
|
||||
which memba's raw C-level save/load is missing.
|
||||
|
||||
If built-in works → memba's MEMB format needs to be extended to include
|
||||
those trackers. If built-in also fails → the issue is elsewhere.
|
||||
"""
|
||||
import sys, time, pickle
|
||||
from pathlib import Path
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "python"))
|
||||
from llama_cpp import Llama
|
||||
|
||||
MODEL = "/home/emil/Desktop/Coding/AI/Memba/NVIDIA-Nemotron3-Nano-4B-Q4_K_M.gguf"
|
||||
STATE = "/tmp/diag_nemotron_native.pickle"
|
||||
|
||||
INGEST = ("I'm going to tell you a fact about my pet. My pet hamster is named "
|
||||
"Bartholomew. He is 4 years old. Reply with just 'noted'.")
|
||||
QUERY = "What is the name of my pet?"
|
||||
|
||||
|
||||
def make_llama():
|
||||
return Llama(model_path=MODEL, n_ctx=4096, n_gpu_layers=-1, verbose=False)
|
||||
|
||||
|
||||
def build():
|
||||
m = make_llama()
|
||||
print(f"[build] initial n_tokens={m.n_tokens}")
|
||||
|
||||
msgs = [{"role": "user", "content": INGEST}]
|
||||
out = m.create_chat_completion(messages=msgs, max_tokens=80, temperature=0.1)
|
||||
ack = out["choices"][0]["message"]["content"]
|
||||
msgs.append({"role": "assistant", "content": ack})
|
||||
print(f"[build] after ingest: n_tokens={m.n_tokens}")
|
||||
print(f"[build] ack snippet: {ack[-100:]!r}")
|
||||
|
||||
# Use llama-cpp-python's NATIVE save_state — captures Python trackers too
|
||||
state = m.save_state()
|
||||
with open(STATE, "wb") as f:
|
||||
pickle.dump(state, f)
|
||||
print(f"[build] saved native state to {STATE} ({Path(STATE).stat().st_size:,} B)")
|
||||
|
||||
|
||||
def query():
|
||||
if not Path(STATE).exists():
|
||||
print("[query] no state file"); return
|
||||
m = make_llama()
|
||||
print(f"[query] before load: n_tokens={m.n_tokens}")
|
||||
|
||||
with open(STATE, "rb") as f:
|
||||
state = pickle.load(f)
|
||||
m.load_state(state)
|
||||
print(f"[query] after load: n_tokens={m.n_tokens}")
|
||||
|
||||
# Now ask via create_chat_completion. Pass ONLY the new question
|
||||
# (history is already in KV-cache+n_tokens).
|
||||
# The template will format this as if it's turn 1 — see what happens.
|
||||
out = m.create_chat_completion(
|
||||
messages=[{"role": "user", "content": QUERY}],
|
||||
max_tokens=150, temperature=0.1,
|
||||
)
|
||||
print(f"\n[query] answer via chat_completion (resets context):\n{out['choices'][0]['message']['content']}")
|
||||
|
||||
# Alternative: try to continue manually, using raw eval
|
||||
print(f"\n[query] re-loading state for raw continuation test")
|
||||
with open(STATE, "rb") as f:
|
||||
state = pickle.load(f)
|
||||
m.load_state(state)
|
||||
print(f"[query] re-loaded: n_tokens={m.n_tokens}")
|
||||
|
||||
# Append a new user turn via raw tokens, then generate
|
||||
continuation = "<|im_end|>\n<|im_start|>user\nWhat is the name of my pet?<|im_end|>\n<|im_start|>assistant\n"
|
||||
tokens = m.tokenize(continuation.encode(), add_bos=False, special=True)
|
||||
|
||||
out_tokens = []
|
||||
eos = m.token_eos()
|
||||
for tok in m.generate(tokens, top_k=1, temp=0.0):
|
||||
if tok == eos or len(out_tokens) >= 150:
|
||||
break
|
||||
out_tokens.append(tok)
|
||||
snippet = m.detokenize(out_tokens).decode("utf-8", errors="ignore")
|
||||
if "<|im_end|>" in snippet:
|
||||
break
|
||||
|
||||
text = m.detokenize(out_tokens).decode("utf-8", errors="ignore")
|
||||
print(f"\n[query] raw continuation answer:\n{text}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
cmd = sys.argv[1] if len(sys.argv) > 1 else "build"
|
||||
{"build": build, "query": query}[cmd]()
|
||||
Reference in New Issue
Block a user