CoolFace
Modelpublic

xiaol/gemma-4-e4B-hybrid-rnn-mem-rwkv-fable5-gpt5.5-v1

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
6likes
Model Card
"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:

text
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:

GoalHow this release approaches it
๐ŸงŠ Keep original model abilityFreeze the base Gemma weights; learn only the online memory path.
๐Ÿง  Add stateful behaviorRWKV-MS memory keeps recurrent state across prompt ingestion and decoding.
๐Ÿ’ป Stay local/smallThe learned weights are tiny; the base model is still the main VRAM cost.
๐Ÿ”ญ Next directionUse multi-state memory for long-context state selection and context extension.

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

FilePurpose
delta_mem_adapter.ptRWKV-MS online-memory weights. The filename comes from delta-Mem.
delta_mem_config.jsonMemory configuration consumed by the patched delta-Mem runtime.
inference.pyMinimal CLI inference script. Canonical copy lives in the source repo.
adapter_metadata.jsonMachine-readable memory config, training summary, and benchmark notes.
requirements.txtMinimal package list for the runtime environment.

Plain AutoModelForCausalLM.from_pretrained() on this repo will not work. You need the base model plus the patched delta-Mem runtime.

๐Ÿงฌ Architecture Snapshot

FieldValue
Base checkpointgoogle/gemma-4-E4B-it
Base weightsfrozen, not included
Memory typeRWKV-MS online recurrent memory
Runtime wrapperdelta-Mem attention/session runtime
Wrapped layersGemma4 text attention layers 0-5
Delta headsq,o
Rank / alpha8 / 16
RWKV-MS states4
Chunk size1024
Trainable memory params797,808

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:

Model / conditionResult
๐Ÿงฑ Base google/gemma-4-E4B-it, focused tools + line verify + autostop4/20, pass_hat_1=0.20
๐Ÿ“ Base google/gemma-4-E4B-it, checklist prompt7/20, pass_hat_1=0.35
๐Ÿง  RWKV-MS online memory, checkpoint step-10014/20, pass_hat_1=0.70
โฑ๏ธ Same run final checkpoint step-20012/20, pass_hat_1=0.60
๐Ÿงช Generated-action SFT, 6 layers, len25610/20, pass_hat_1=0.50
๐Ÿงช Generated-action SFT, 2 layers, len2569/20, pass_hat_1=0.45

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".

StageLocal dataBudget
Generated mobile-data action SFT3,519 turn rows656 optimizer steps
Format-refresh continuation5,027 turn rows200 optimizer steps
Selected checkpointcontinuation step-100best 20-task screen

๐Ÿ“š 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.

Dataset stageRowsRole
tau2_telecom_mobile_data_rule_planner_train_focusedtools_nopolicy_turns.jsonl3,519generated mobile-data action SFT
tau2_telecom_mobile_data_rule_planner_train_focusedtools_nopolicy_turns_formatrefresh_balanced.jsonl5,027balanced format/action continuation
original converted tau2_telecom_all_valid.jsonl82tested and rejected; loss moved but benchmark transfer failed

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:

KnobWhy change it
max_lengthFirst lever for VRAM. Shorter context was used here to fit safely.
wrapped layersMore layers add capacity and online state; 6 layers beat 2 here.
rank / alphaControls memory-path size and strength.
local data formatThe original 82-row tau2 data moved loss but did not transfer; aligned action data worked better.

๐Ÿ”ญ 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:

DirectionWhy it matters
๐Ÿงช Fable-5 / GPT-5.5 dataTest whether richer generated traces improve generalization beyond tau2 telecom.
๐Ÿงฑ Layer sweepsCompare 2-layer, 6-layer, and deeper selective bands instead of assuming one layer budget.
๐ŸŽš๏ธ Rank/state sweepsMeasure memory capacity, VRAM, and benchmark behavior at different ranks and state counts.
๐Ÿงญ Selective memoryRoute each token to a small subset of memory states to reduce interference and support longer contexts.

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:

bash
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.patch

Run the online-memory checkpoint:

bash
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 sdpa

Use --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.

python
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:

text
[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.