Introduction
MemSFT specializes modern large language models with an external parametric
memory. This checkpoint contains the Qwen3-8B memory trained on DISC-Law-SFT.
The memory learns to approximate retrieval-based teacher distributions over
domain SFT data. At each decoding step, a learned token-level router combines
the next-token distributions of the frozen base model and memory. This
checkpoint is an auxiliary memory, not a standalone chat model. Its key
advantages are:
- Plug-and-Play: Attaches to a frozen backbone without modifying its
parameters or architecture.
- Legal Specialization: Improves legal knowledge and reasoning on the
evaluated LawBench tasks.
- Capability Retention: Preserves the backbone's average performance on
the evaluated general benchmarks.
Quick Start
The evaluated 14B + 8B configuration is intended for a CUDA GPU with
sufficient memory to load both models in BF16.
1. Install
git clone https://github.com/LUMIA-Group/MemSFT.git
cd MemSFT
conda create -n memsft-generate python=3.10 pip -y
conda activate memsft-generate
python -m pip install -e .
python -m pip install \
"torch>=2.4,<2.7" \
"transformers==4.51.3" \
"huggingface-hub==0.35.3" \
"accelerate>=0.34,<2"
2. Load the base, memory, and router
from pathlib import Path
import torch
from huggingface_hub import snapshot_download
from transformers import AutoModelForCausalLM, AutoTokenizer, set_seed
from memsft.router.adaptive_memdec import AdaptiveMemoryDecoder
device = torch.device("cuda:0")
base_id = "Qwen/Qwen3-14B"
memory_id = "Jiarui-Wang/MemSFT-Qwen3-Law-Memory-8B"
router_repo = "Jiarui-Wang/MemSFT-Qwen3-Routers"
router_subdir = "Qwen3-14B-Law-M8B-Router"
router_root = snapshot_download(
repo_id=router_repo,
revision="v1.0.0",
allow_patterns=[f"{router_subdir}/*"],
)
router_path = str(Path(router_root) / router_subdir)
tokenizer = AutoTokenizer.from_pretrained(
base_id,
revision="40c069824f4251a91eefaf281ebe4c544efd3e18",
)
base = AutoModelForCausalLM.from_pretrained(
base_id,
revision="40c069824f4251a91eefaf281ebe4c544efd3e18",
torch_dtype=torch.bfloat16,
low_cpu_mem_usage=True,
).to(device).eval()
memory = AutoModelForCausalLM.from_pretrained(
memory_id,
revision="v1.0.0",
torch_dtype=torch.bfloat16,
low_cpu_mem_usage=True,
).to(device).eval()
vocab_size = len(tokenizer)
base.resize_token_embeddings(vocab_size)
memory.resize_token_embeddings(vocab_size)
base.requires_grad_(False)
memory.requires_grad_(False)
model = AdaptiveMemoryDecoder(
base_lm=base,
knn_generator=memory,
router_path=router_path,
router_device=device,
).eval()
model.set_tokenizer(tokenizer)
3. Generate
instruction = (
"依据给出的实体类型提取句子的实体信息,实体类型包括:犯罪嫌疑人、受害人、"
"被盗货币、物品价值、盗窃获利、被盗物品、作案工具、时间、地点、组织机构。"
"逐个列出实体信息。"
)
question = "句子:被告人周某甲被归案。"
prompt = f"{instruction}\n{question}"
messages = [{"role": "user", "content": prompt}]
prompt_text = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)
inputs = tokenizer(prompt_text, return_tensors="pt").to(device)
set_seed(42)
with torch.inference_mode():
output_ids = model.generate(
**inputs,
do_sample=True,
temperature=0.6,
top_p=0.95,
top_k=20,
max_new_tokens=256,
eos_token_id=tokenizer.eos_token_id,
pad_token_id=tokenizer.eos_token_id,
)
answer = tokenizer.decode(
output_ids[0, inputs["input_ids"].shape[1]:],
skip_special_tokens=True,
)
print(answer)
Example output:
The output exactly matches the reference entity set. With the same prompt,
seed, and generation settings, Qwen3-14B alone instead enumerates all ten
entity types and assigns the nine unmentioned types the value 无. Both runs
ended naturally at EOS in BF16 on NVIDIA A800 80GB GPUs.
LawBench is distributed under the Apache License 2.0.
The released Law memory is evaluated with Qwen3-14B and its matching router.
LawBench Avg. is computed over the 19 evaluated legal subtasks on a 0–100
scale, where higher is better. General is the average over MATH-500, C-Eval,
IFEval, MMLU-Redux, and INCLUDE.
Table with columns: Base model, LawBench Avg. ↑, General ↑| Base model | LawBench Avg. ↑ | General ↑ |
|---|
| Qwen3-14B | 49.83 | 83.22 |
| Qwen3-14B + MemSFT | 56.47 | 83.79 |
Compatible Pairing
The evaluated pairing is:
- base:
Qwen/Qwen3-14B
- memory:
Jiarui-Wang/MemSFT-Qwen3-Law-Memory-8B
- router:
Jiarui-Wang/MemSFT-Qwen3-Routers/Qwen3-14B-Law-M8B-Router
MemSFT prefers tensor-only .safetensors router checkpoints. Legacy .pt
checkpoints should be loaded only from trusted sources; the MemSFT loader uses
PyTorch's restricted weights_only=True mode for compatibility.
Intended Use and Limitations
This checkpoint is intended for reproducing MemSFT and for augmenting
Qwen3-14B on the evaluated legal-domain tasks. It should be used with the
matching Qwen3-14B/M8B router. Performance outside the evaluated model
combination and domain has not been established. Model outputs may be
incorrect or incomplete and should not be treated as legal advice. Qualified
professionals should review outputs before use in consequential legal
applications.
License
This MemSFT checkpoint is released under the Apache License 2.0. Upstream
models, software, and datasets remain subject to their respective licenses
and terms.
Citation
If you find MemSFT helpful in your research, please consider citing:
@misc{wang2026memsftmitigatingalignmenttax,
title={MemSFT: Mitigating Alignment Tax with an External Parametric Memory},
author={Jiarui Wang and Xiang Shi and Jiaqi Cao and Rubin Wei and Xiquan Wang and Hao Sun and Jingzhi Wang and Zhiqi Yang and Qipeng Guo and Bowen Zhou and Zhouhan Lin},
year={2026},
eprint={2607.25614},
archivePrefix={arXiv},
primaryClass={cs.LG},
url={https://arxiv.org/abs/2607.25614},
}
For questions and discussions, feel free to email
wangjiarui1@sjtu.edu.cn.