--- license: apache-2.0 library_name: transformers pipeline_tag: text-generation base_model: - Qwen/Qwen3-8B tags: - memsft - memory-decoder - biology - biology-instructions - qwen3 --- # MemSFT-Qwen3-Bio-Memory-8B

📄 Paper • 💻 GitHub • 🤗 HF Collection

## Introduction MemSFT specializes modern large language models with an external parametric memory. This checkpoint contains the Qwen3-8B memory trained on Biology-Instructions. 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. - **Strong Specialization:** Improves domain performance with negligible degradation in general capabilities. - **Cross-Scale Reuse:** Works with Qwen3 backbones from 8B to 235B-A22B without retraining the memory for each backbone. ## Quick Start The 14B + 8B example is intended for a CUDA GPU with sufficient memory to load both models in BF16. ### 1. Install ```bash 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 ```python from pathlib import Path import torch from huggingface_hub import snapshot_download from transformers import AutoModelForCausalLM, AutoTokenizer from memsft.router.adaptive_memdec import AdaptiveMemoryDecoder device = torch.device("cuda:0") base_id = "Qwen/Qwen3-14B" memory_id = "Jiarui-Wang/MemSFT-Qwen3-Bio-Memory-8B" router_repo = "Jiarui-Wang/MemSFT-Qwen3-Routers" router_subdir = "Qwen3-14B-Bio-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 ```python sequence = ( "MKSILIEKPNQLAIVEREIPTPSAGEVRVKVKLAGICGSDSHIYRGHNPFAKYPRVIGHEFFGVIDAV" "GEGVESARVGERVAVDPVVSCGHCYPCSIGKPNVCTTLAVLGVHADGGFSEYAVVPAKNAWKIPEAVA" "DQYAVMIEPFTIAANVTGHGQPTENDTVLVYGAGPIGLTIVQVLKGVYNVKNVIVADRIDERLEKAKE" "SGADWAINNSQTPLGEIFTEKGIKPTLIIDAACHPSILKEAVTLASPAARIVLMGFSSEPSEVIQQGI" "TGKELSIFSSRLNANKFPIVIDWLSKGLIKPEKLITHTFDFQHVADAISLFEQDQKHCCKVLLTFSE" ) prompt = ( r" " + sequence + r" What is the EC number associated with the enzymatic " r"function of this protein? Please put the final enzyme within \boxed{} " r"using an EC number such as ECx.x.x.x, and separate multiple entries " r"with commas." ) 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) with torch.inference_mode(): output_ids = model.generate( **inputs, do_sample=False, max_new_tokens=32, 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:** ```text \boxed{EC1.1.1.-} ``` UniProtKB annotates *Escherichia coli* K-12 RspB ([P38105](https://www.uniprot.org/uniprotkb/P38105/entry)) with EC `1.1.1.-`. For comparison, using the same prompt and deterministic generation configuration, Qwen3-14B alone predicts `EC 4.2.1.22`. The outputs were reproduced in BF16 on NVIDIA A800 80GB GPUs. ## Performance The same Qwen3-8B Biology-Instructions memory is reused with each base model; each pairing uses its corresponding router. | Base model | Biology-Instructions ↑ | General ↑ | |---|---:|---:| | Qwen3-8B + MemSFT | 42.84 | 81.15 | | Qwen3-14B + MemSFT | 42.92 | 83.62 | | Qwen3-32B + MemSFT | 42.82 | 85.33 | | Qwen3-235B-A22B + MemSFT | 42.05 | 87.11 | ## Compatible Pairing The example above uses: - base: [`Qwen/Qwen3-14B`](https://huggingface.co/Qwen/Qwen3-14B) - memory: `Jiarui-Wang/MemSFT-Qwen3-Bio-Memory-8B` - router: `Jiarui-Wang/MemSFT-Qwen3-Routers/Qwen3-14B-Bio-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 compatible Qwen3 base models on Biology-Instructions tasks. It should be used with the router matching the selected base/memory pair. Performance outside the evaluated model combinations and domains has not been established. ## 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: ```bibtex @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}, } ``` ## Contact For questions and discussions, feel free to email [wangjiarui1@sjtu.edu.cn](mailto:wangjiarui1@sjtu.edu.cn).