xiaol/gemma-4-e4B-hybrid-rnn-mem-rwkv-fable5-gpt5.5-v1
"One can have both the fish and the bear's paw."
๐๐พ Gemma4 E4B + RWKV-MS Online Memory
๐ง Keep the Gemma base, add a tiny recurrent memory path for local tool-agent behavior.
This is not a merged model and not a normal LoRA card. The original google/gemma-4-E4B-it weights stay frozen and untouched. This repo ships a compact RWKV multi-state online memory checkpoint that reads and writes recurrent state during inference. โกSource repo and canonical inference entry point:
https://github.com/xiaol/Multi-state-RWKV-online-memory๐ฏ What It Is
This checkpoint attaches RWKV-MS online RNN memory to the first six Gemma4 text-attention layers. The learned memory path is small: 797,808 trainable parameters, about 0.8M, while the Gemma4 E4B base checkpoint is still loaded separately and remains unchanged.
The practical idea is simple:
This is a research checkpoint, but the signal matters for local small models: the base checkpoint scored 4/20 on the accepted tau2 telecom screen, while this online-memory checkpoint reached 14/20.
๐ฆ What You Get
Plain AutoModelForCausalLM.from_pretrained() on this repo will not work. You need the base model plus the patched delta-Mem runtime.
๐งฌ Architecture Snapshot
delta-Mem provides attention wrapping, memory checkpoint loading, online state handling, KV-cache/session synchronization, and chat-template handling. Multi-state-RWKV-online-memory provides the RWKV-MS patch, benchmark docs, and recommended inference script. This Hub repo stores the memory weights/config and a convenience script.
๐ Early Tau2 Signal
Tau2 telecom 20-task screen, accepted no-rule setup:
Accepted benchmark setup: solo/dummy-user tau2 mode, greedy decoding, max_new_tokens=96, max_steps=40, max_errors=4, seed 300, infra_error_count=0.
The rule-based planner / float-format repair path is not included in this comparison because it is benchmark-specific control logic, not model behavior.
This is still a 20-task model-selection screen. A larger >=50 task run or full telecom split is needed before treating the gain as robust.
๐๏ธ Local Training Cost & Recipe
All runs behind this checkpoint were local experiments. The accepted checkpoint was trained/evaluated on a local RTX 4090 24 GB setup using CUDA bf16 with attn_implementation="sdpa".
๐ Training Data Provenance
Despite the repo name containing gpt5.5, the accepted checkpoint here was not trained on data generated by GPT-5.5. The selected RWKV-MS online-memory checkpoint used tau2 telecom mobile-data/action traces generated by a local deterministic tau2 rule-planner pipeline, replayed against the tau2 environment, then turn-sliced for next-action learning.
The training rows are synthetic/derived tau2 action traces, not human support logs and not private customer data. The rule planner is used to create local training targets; it is not included as a model-comparison row in the benchmark table because eval-time planner logic would be benchmark-specific control code rather than learned model behavior.
Leakage note: the generated training set was built with the reported 20-task benchmark screen held out (exclude_heldout=true in the local data summary), so the accepted checkpoint was not trained on those exact benchmark task IDs. This avoids exact task leakage. It is still the same tau2 telecom/mobile-data family, with synthetic traces from the same environment and tools, so the reported 14/20 should be read as an in-domain held-out screen, not a broad out-of-domain generalization result.
GPT-5.5-generated traces and Fable-5-style data are planned upgrade data, not a claim about this exact checkpoint. The next model iteration should test whether Fable-5 and GPT-5.5-generated multi-domain traces improve generalization beyond the narrow tau2 mobile-data screen.
The exact wall-clock cost is hardware- and cache-dependent, so treat this as a small local recipe, not a fixed training quote. VRAM use varies with base-model path, attention backend, sequence length, layer count, rank, tokenizer cache, and fragmentation. For your own domain, keep the base frozen and adjust:
๐ญ Next Upgrade Direction
The current checkpoint is a narrow first signal. The next upgrades should keep the frozen-base principle and test the memory path more systematically:
The selective-memory direction is connected to `xiaol/SelectingMemory`, which explores Raven-style top-k memory-slot routing and RWKV-7 mixer variants. The relevant idea for RWKV-MS is not a claim of solved long context yet; it is a research path where each token chooses which recurrent memory states to update, while unselected states are preserved for later recall.
๐ Quick Start
Install and patch the runtime:
git clone https://github.com/xiaol/Multi-state-RWKV-online-memory.git
git clone https://github.com/declare-lab/delta-Mem.git
cd delta-Mem
python -m venv .venv
source .venv/bin/activate
pip install -U pip setuptools wheel
pip install -r requirements.txt
pip install -U "huggingface_hub>=1.0.0"
git apply --unidiff-zero --whitespace=nowarn \
../Multi-state-RWKV-online-memory/integrations/delta_mem_rwkv_ms/delta_mem_rwkv_ms.patchRun the online-memory checkpoint:
cd ../Multi-state-RWKV-online-memory
../delta-Mem/.venv/bin/python integrations/delta_mem_rwkv_ms/inference.py \
--delta-mem-root ../delta-Mem \
--memory-repo xiaol/gemma-4-e4B-hybrid-rnn-mem-rwkv-fable5-gpt5.5-v1 \
--base-model google/gemma-4-E4B-it \
--device cuda:0 \
--dtype bfloat16 \
--attn-implementation sdpaUse --memory-dir /path/to/local/model-repo if you have already cloned this Hub repo. Use a local --base-model /path/to/gemma-4-E4B-it if your Gemma checkpoint is stored outside the Hub cache.
๐งช Tau2-Style Python API Smoke Test
This is a benchmark-like sanity check, not the tau2 benchmark harness. It does not execute tools, simulate a user, enforce max_steps, or compute pass/fail. It checks that the patched runtime loads and the model can produce the next tau2-style telecom tool action greedily.
from huggingface_hub import snapshot_download
from deltamem.runtime.session import DeltaMemChatSession, load_delta_mem_chat_model
prompt = """You are a telecom solo-mode tool agent. Return exactly one tool call in this format:
[ACTION]
tool_name(arg_name="value")
[/ACTION]
Available tools:
- get_customer_by_phone(phone_number: str)
- check_network_status(line_id: str)
- toggle_data(line_id: str, enabled: bool)
- run_speed_test(line_id: str)
- done()
Ticket: Customer phone number 555-123-2002 reports no usable mobile data.
First step: identify the customer account from the phone number. Return only the next tool call."""
memory_dir = snapshot_download(
"xiaol/gemma-4-e4B-hybrid-rnn-mem-rwkv-fable5-gpt5.5-v1"
)
model, tokenizer = load_delta_mem_chat_model(
model_path="google/gemma-4-E4B-it", # or your local base checkpoint path
adapter_dir=memory_dir, # delta-Mem API name for the memory repo
device="cuda:0",
dtype="bfloat16",
attn_implementation="sdpa",
)
session = DeltaMemChatSession(model=model, tokenizer=tokenizer, device="cuda:0")
out = session.generate_reply(
prompt,
max_new_tokens=64,
do_sample=False,
include_debug=True,
)
print(out["assistant_display"])
print(out["state_stats"])
print(out["turn_stats"])Real greedy response from the Python API:
[ACTION]
get_customer_by_phone(phone_number="555-123-2002")
[/ACTION]Observed debug summary on the local smoke run: all 6 memory modules had nonzero state; prompt ingest was 160 tokens; decode generated 37 tokens in about 980 ms after the model was already loaded.
๐งญ Practical Notes
- โ
Tested path uses CUDA, bf16, and
attn_implementation="sdpa". - โ The base model remains the dominant VRAM cost; this repo adds a tiny memory checkpoint, not another full model copy.
- โ The intended adaptation path is local: keep Gemma frozen, train the online memory on your own agent traces or domain data, then benchmark honestly.
- โ ๏ธ GGUF is a possible next step, but this release is not GGUF yet. A GGUF path needs a runtime representation for the online RWKV-MS state and read/write hooks, not only static quantized weights.
โ ๏ธ Limitations
- Tuned for a narrow telecom/tool-agent setting.
- The reported gain is from a 20-task screen.
- Requires the patched delta-Mem runtime; it is not a drop-in Transformers-only model.
- Safety behavior is inherited mostly from the base checkpoint and was not the focus of this run.
- Freezing the base helps preserve original behavior, but you should still run your own regression checks for any deployment domain.
- Context-length boost and long-context state selection are intended next directions, not solved claims in this checkpoint.
๐ License
Apache-2.0. This checkpoint requires separate access to and compliance with the google/gemma-4-E4B-it base model license.
