Views
No views yet
1git clone https://github.com/LUMIA-Group/MemSFT.git
2cd MemSFT
3
4conda create -n memsft-generate python=3.10 pip -y
5conda activate memsft-generate
6python -m pip install -e .
7python -m pip install \
8 "torch>=2.4,<2.7" \
9 "transformers==4.51.3" \
10 "huggingface-hub==0.35.3" \
11 "accelerate>=0.34,<2"1from pathlib import Path
2
3import torch
4from huggingface_hub import snapshot_download
5from transformers import AutoModelForCausalLM, AutoTokenizer, set_seed
6
7from memsft.router.adaptive_memdec import AdaptiveMemoryDecoder
8
9device = torch.device("cuda:0")
10base_id = "Qwen/Qwen3-14B"
11memory_id = "Jiarui-Wang/MemSFT-Qwen3-Law-Memory-8B"
12router_repo = "Jiarui-Wang/MemSFT-Qwen3-Routers"
13router_subdir = "Qwen3-14B-Law-M8B-Router"
14
15router_root = snapshot_download(
16 repo_id=router_repo,
17 revision="v1.0.0",
18 allow_patterns=[f"{router_subdir}/*"],
19)
20router_path = str(Path(router_root) / router_subdir)
21
22tokenizer = AutoTokenizer.from_pretrained(
23 base_id,
24 revision="40c069824f4251a91eefaf281ebe4c544efd3e18",
25)
26base = AutoModelForCausalLM.from_pretrained(
27 base_id,
28 revision="40c069824f4251a91eefaf281ebe4c544efd3e18",
29 torch_dtype=torch.bfloat16,
30 low_cpu_mem_usage=True,
31).to(device).eval()
32memory = AutoModelForCausalLM.from_pretrained(
33 memory_id,
34 revision="v1.0.0",
35 torch_dtype=torch.bfloat16,
36 low_cpu_mem_usage=True,
37).to(device).eval()
38
39vocab_size = len(tokenizer)
40base.resize_token_embeddings(vocab_size)
41memory.resize_token_embeddings(vocab_size)
42base.requires_grad_(False)
43memory.requires_grad_(False)
44model = AdaptiveMemoryDecoder(
45 base_lm=base,
46 knn_generator=memory,
47 router_path=router_path,
48 router_device=device,
49).eval()
50model.set_tokenizer(tokenizer)1instruction = (
2 "依据给出的实体类型提取句子的实体信息,实体类型包括:犯罪嫌疑人、受害人、"
3 "被盗货币、物品价值、盗窃获利、被盗物品、作案工具、时间、地点、组织机构。"
4 "逐个列出实体信息。"
5)
6question = "句子:被告人周某甲被归案。"
7prompt = f"{instruction}\n{question}"
8
9messages = [{"role": "user", "content": prompt}]
10prompt_text = tokenizer.apply_chat_template(
11 messages,
12 tokenize=False,
13 add_generation_prompt=True,
14 enable_thinking=False,
15)
16inputs = tokenizer(prompt_text, return_tensors="pt").to(device)
17
18set_seed(42)
19with torch.inference_mode():
20 output_ids = model.generate(
21 **inputs,
22 do_sample=True,
23 temperature=0.6,
24 top_p=0.95,
25 top_k=20,
26 max_new_tokens=256,
27 eos_token_id=tokenizer.eos_token_id,
28 pad_token_id=tokenizer.eos_token_id,
29 )
30
31answer = tokenizer.decode(
32 output_ids[0, inputs["input_ids"].shape[1]:],
33 skip_special_tokens=True,
34)
35print(answer)犯罪嫌疑人:周某甲无. Both runs
ended naturally at EOS in BF16 on NVIDIA A800 80GB GPUs.| Base model | LawBench Avg. ↑ | General ↑ |
|---|---|---|
| Qwen3-14B | 49.83 | 83.22 |
| Qwen3-14B + MemSFT | 56.47 | 83.79 |
Qwen/Qwen3-14BJiarui-Wang/MemSFT-Qwen3-Law-Memory-8BJiarui-Wang/MemSFT-Qwen3-Routers/Qwen3-14B-Law-M8B-Router.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.1@misc{wang2026memsftmitigatingalignmenttax,
2 title={MemSFT: Mitigating Alignment Tax with an External Parametric Memory},
3 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},
4 year={2026},
5 eprint={2607.25614},
6 archivePrefix={arXiv},
7 primaryClass={cs.LG},
8 url={https://arxiv.org/abs/2607.25614},
9}