Text Generation
Transformers
Safetensors
livemem
qwen3
custom-code
long-context
reinforcement-learning
conversational
custom_code
Instructions to use chen-l/LiveMem-RL with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use chen-l/LiveMem-RL with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="chen-l/LiveMem-RL", trust_remote_code=True) messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("chen-l/LiveMem-RL", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use chen-l/LiveMem-RL with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "chen-l/LiveMem-RL" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "chen-l/LiveMem-RL", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/chen-l/LiveMem-RL
- SGLang
How to use chen-l/LiveMem-RL with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "chen-l/LiveMem-RL" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "chen-l/LiveMem-RL", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "chen-l/LiveMem-RL" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "chen-l/LiveMem-RL", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use chen-l/LiveMem-RL with Docker Model Runner:
docker model run hf.co/chen-l/LiveMem-RL
| """Memory-augmented Qwen3: 主路 Qwen3Attention ‖ 边路 GDN2, o = o_main + o_side. | |
| Design X (continuous scan) and Design Y (gated read/write) share one skeleton | |
| and the same attention eviction mask; they differ only in the RNN `write_mask`: | |
| - X: write_mask = None (RNN scans every token; state = compression of all) | |
| - Y: write_mask = is_evicted (open gate only on the evict/compress region) | |
| """ | |
| from __future__ import annotations | |
| import torch | |
| import torch.nn as nn | |
| from transformers.cache_utils import Cache, DynamicCache | |
| from transformers.masking_utils import create_causal_mask, create_sliding_window_causal_mask | |
| from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast | |
| from transformers.models.qwen3.modeling_qwen3 import ( | |
| Qwen3Attention, | |
| Qwen3ForCausalLM, | |
| Qwen3Model, | |
| Qwen3PreTrainedModel, | |
| ) | |
| from .configuration_livemem import LiveMemConfig | |
| from .modeling_livemem_gdn2 import LiveMemGatedDeltaNet2 | |
| def make_memory_and_mask( | |
| is_evicted: torch.Tensor | None = None, | |
| segment_ids: torch.Tensor | None = None, | |
| seq_ids: torch.Tensor | None = None, | |
| chunk_id: torch.Tensor | None = None, | |
| evict_step: torch.Tensor | None = None, | |
| ): | |
| """Build an `and_mask_function` for `create_causal_mask` from up to four | |
| constraints, AND-combined with causal by `create_causal_mask`: | |
| - dynamic eviction (`chunk_id [B,T]`, `evict_step [B,T]` int): the real | |
| training path. Each token belongs to a chunk; `evict_step[t]` is the chunk | |
| step at which that token's chunk is evicted to the RNN state. keep(q,kv) = | |
| `evict_step[kv] > chunk_id[q]` — key kv's chunk must still be live when | |
| query q's chunk is processed (Design X: RNN scans all, attention evicts). | |
| - static eviction (`is_evicted [B,T]` bool): the simple synthetic variant. | |
| keep(q,kv) = (kv live) OR (q evicted). | |
| - segments (`segment_ids [B,T]` int): keep(q,kv) = (kv shared, id==0) OR | |
| (same segment). PACK block-diagonal QA. | |
| - documents (`seq_ids [B,T]` int): keep(q,kv) = same document. cu_seqlens | |
| packing isolation. | |
| Returns None if no constraint is given. Works for flex_attention (BlockMask) | |
| and sdpa/eager (vmapped) backends. | |
| """ | |
| preds = [] | |
| if chunk_id is not None and evict_step is not None: | |
| preds.append(lambda b, q, kv: evict_step[b, kv] > chunk_id[b, q]) | |
| if is_evicted is not None: | |
| preds.append(lambda b, q, kv: (~is_evicted[b, kv]) | is_evicted[b, q]) | |
| if segment_ids is not None: | |
| preds.append(lambda b, q, kv: (segment_ids[b, kv] == 0) | (segment_ids[b, kv] == segment_ids[b, q])) | |
| if seq_ids is not None: | |
| preds.append(lambda b, q, kv: seq_ids[b, kv] == seq_ids[b, q]) | |
| if not preds: | |
| return None | |
| def and_mask(b, h, q, kv): | |
| out = preds[0](b, q, kv) | |
| for p in preds[1:]: | |
| out = out & p(b, q, kv) | |
| return out | |
| return and_mask | |
| # Backwards-compatible alias. | |
| def make_evict_and_mask(is_evicted: torch.Tensor): | |
| return make_memory_and_mask(is_evicted) | |
| class LiveMemAttention(nn.Module): | |
| """Wraps the original Qwen3Attention (main path) and adds a GDN2 side branch. | |
| Per-forward memory control (`write_mask`, side-branch cache) is set as | |
| attributes by the model loop rather than threaded through kwargs, so the | |
| base attention path and HF decorators never see custom kwargs. | |
| """ | |
| def __init__(self, base_attn: Qwen3Attention, config: LiveMemConfig) -> None: | |
| super().__init__() | |
| self.layer_idx = base_attn.layer_idx | |
| self.attn = base_attn # main path: untouched Qwen3Attention | |
| self.mem = LiveMemGatedDeltaNet2( | |
| hidden_size=config.hidden_size, | |
| expand_v=config.mem_expand_v, | |
| head_dim=config.mem_head_dim, | |
| num_heads=config.mem_num_heads, | |
| num_v_heads=config.mem_num_v_heads, | |
| mode="chunk", | |
| use_short_conv=True, | |
| conv_size=config.mem_conv_size, | |
| conv_bias=config.mem_conv_bias, | |
| layer_idx=base_attn.layer_idx, | |
| norm_eps=config.mem_norm_eps, | |
| ) | |
| # per-forward control, set by LiveMemModel.forward | |
| self._mem_write_mask: torch.Tensor | None = None | |
| self._mem_cache = None | |
| self._mem_use_cache: bool = False | |
| self._mem_cu_seqlens: torch.Tensor | None = None | |
| # Training diagnostics. Disabled by default; train/sft/loop.py enables | |
| # this on one layer so normal forward/inference pays no reduction cost. | |
| self._record_o_stats: bool = False | |
| self._last_o_stats: dict[str, torch.Tensor] = {} | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| position_embeddings, | |
| attention_mask: torch.Tensor | None, | |
| past_key_values: Cache | None = None, | |
| **kwargs, | |
| ): | |
| o_main, attn_weights = self.attn( | |
| hidden_states, | |
| position_embeddings=position_embeddings, | |
| attention_mask=attention_mask, | |
| past_key_values=past_key_values, | |
| **kwargs, | |
| ) | |
| o_side, _, _ = self.mem( | |
| hidden_states, | |
| write_mask=self._mem_write_mask, | |
| past_key_values=self._mem_cache, | |
| use_cache=self._mem_use_cache, | |
| cu_seqlens=self._mem_cu_seqlens, | |
| ) | |
| o_total = o_main + o_side | |
| if self._record_o_stats: | |
| with torch.no_grad(): | |
| main_abs = o_main.detach().float().abs().mean() | |
| side_abs = o_side.detach().float().abs().mean() | |
| total_abs = o_total.detach().float().abs().mean() | |
| eps = torch.tensor(1e-12, device=total_abs.device, dtype=total_abs.dtype) | |
| self._last_o_stats = { | |
| "main_out_abs": main_abs, | |
| "side_out_abs": side_abs, | |
| "total_out_abs": total_abs, | |
| "side_out_ratio": side_abs / torch.maximum(total_abs, eps), | |
| "side_main_ratio": side_abs / torch.maximum(main_abs, eps), | |
| } | |
| return o_total, attn_weights | |
| class LiveMemPreTrainedModel(Qwen3PreTrainedModel): | |
| config: LiveMemConfig | |
| _no_split_modules = ["Qwen3DecoderLayer"] | |
| class LiveMemModel(LiveMemPreTrainedModel, Qwen3Model): | |
| config_class = LiveMemConfig | |
| def __init__(self, config: LiveMemConfig) -> None: | |
| Qwen3Model.__init__(self, config) | |
| # Replace self_attn with LiveMemAttention on the selected layers. | |
| mem_layers = set(config.memory_layer_indices) | |
| for idx in mem_layers: | |
| layer = self.layers[idx] | |
| layer.self_attn = LiveMemAttention(layer.self_attn, config) | |
| self._mem_layer_indices = sorted(mem_layers) | |
| self.post_init() | |
| def _set_mem_control(self, write_mask, mem_cache, mem_use_cache, cu_seqlens=None) -> None: | |
| for idx in self._mem_layer_indices: | |
| m = self.layers[idx].self_attn | |
| m._mem_write_mask = write_mask | |
| m._mem_cache = mem_cache | |
| m._mem_use_cache = mem_use_cache | |
| m._mem_cu_seqlens = cu_seqlens | |
| def _clear_mem_control(self) -> None: | |
| self._set_mem_control(None, None, False, None) | |
| def forward( | |
| self, | |
| input_ids: torch.LongTensor | None = None, | |
| attention_mask: torch.Tensor | None = None, | |
| position_ids: torch.LongTensor | None = None, | |
| past_key_values: Cache | None = None, | |
| inputs_embeds: torch.FloatTensor | None = None, | |
| use_cache: bool | None = None, | |
| is_evicted: torch.Tensor | None = None, | |
| write_mask: torch.Tensor | None = None, | |
| segment_ids: torch.Tensor | None = None, | |
| seq_ids: torch.Tensor | None = None, | |
| chunk_id: torch.Tensor | None = None, | |
| evict_step: torch.Tensor | None = None, | |
| cu_seqlens: torch.Tensor | None = None, | |
| mem_cache=None, | |
| **kwargs, | |
| ) -> BaseModelOutputWithPast: | |
| if (input_ids is None) ^ (inputs_embeds is not None): | |
| raise ValueError("You must specify exactly one of input_ids or inputs_embeds") | |
| if inputs_embeds is None: | |
| inputs_embeds = self.embed_tokens(input_ids) | |
| if use_cache and past_key_values is None: | |
| past_key_values = DynamicCache(config=self.config) | |
| if position_ids is None: | |
| past_seen = past_key_values.get_seq_length() if past_key_values is not None else 0 | |
| position_ids = torch.arange( | |
| inputs_embeds.shape[1], device=inputs_embeds.device | |
| ).unsqueeze(0) + past_seen | |
| # RNN write gate (per token). An explicit `write_mask` always wins (PACK | |
| # freezes QA segments, segment-write uses 双位置); otherwise Design Y | |
| # derives it from the eviction layout, and Design X scans continuously. | |
| if write_mask is not None: | |
| write_mask = write_mask.to(inputs_embeds.dtype) | |
| elif is_evicted is not None and self.config.memory_design == "Y": | |
| write_mask = is_evicted.to(inputs_embeds.dtype) | |
| # Build the (eviction / segment / document-aware) causal mask, reused for | |
| # all layers. `seq_ids` isolates packed sequences (cu_seqlens path); | |
| # `chunk_id`/`evict_step` drive dynamic chunk eviction (real training). | |
| if not isinstance(attention_mask, dict): | |
| and_mask = make_memory_and_mask(is_evicted, segment_ids, seq_ids, chunk_id, evict_step) | |
| mask_kwargs = { | |
| "config": self.config, | |
| "inputs_embeds": inputs_embeds, | |
| "attention_mask": attention_mask, | |
| "past_key_values": past_key_values, | |
| "position_ids": position_ids, | |
| "and_mask_function": and_mask, | |
| } | |
| causal_mask_mapping = {"full_attention": create_causal_mask(**mask_kwargs)} | |
| if self.has_sliding_layers: | |
| causal_mask_mapping["sliding_attention"] = create_sliding_window_causal_mask(**mask_kwargs) | |
| else: | |
| causal_mask_mapping = attention_mask | |
| hidden_states = inputs_embeds | |
| position_embeddings = self.rotary_emb(hidden_states, position_ids) | |
| # Set control on every forward (incl. None when no eviction), so there is | |
| # no stale state. We deliberately do NOT clear afterwards: gradient | |
| # checkpointing recomputes this forward during backward and must see the | |
| # same write_mask. Training is sequential (forward→backward→next forward), | |
| # so the values stay valid until the next forward overwrites them. | |
| self._set_mem_control(write_mask, mem_cache, bool(use_cache), cu_seqlens) | |
| for i, decoder_layer in enumerate(self.layers[: self.config.num_hidden_layers]): | |
| hidden_states = decoder_layer( | |
| hidden_states, | |
| attention_mask=causal_mask_mapping[self.config.layer_types[i]], | |
| position_embeddings=position_embeddings, | |
| position_ids=position_ids, | |
| past_key_values=past_key_values, | |
| use_cache=use_cache, | |
| **kwargs, | |
| ) | |
| hidden_states = self.norm(hidden_states) | |
| return BaseModelOutputWithPast( | |
| last_hidden_state=hidden_states, | |
| past_key_values=past_key_values if use_cache else None, | |
| ) | |
| class LiveMemForCausalLM(LiveMemPreTrainedModel, Qwen3ForCausalLM): | |
| config_class = LiveMemConfig | |
| _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"} | |
| def __init__(self, config: LiveMemConfig) -> None: | |
| # Build directly (don't call Qwen3ForCausalLM.__init__, which would | |
| # construct a throwaway base Qwen3Model first). | |
| Qwen3PreTrainedModel.__init__(self, config) | |
| self.model = LiveMemModel(config) | |
| self.vocab_size = config.vocab_size | |
| self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) | |
| self.post_init() | |
| # Honor zero-init of the side o_proj for *both* construction paths | |
| # (post_init randomizes it, so this must run last). For from_qwen3 the | |
| # copy-init also zeros it; here it covers from-scratch construction. | |
| if config.mem_o_proj_zero_init: | |
| for idx in config.memory_layer_indices: | |
| nn.init.zeros_(self.model.layers[idx].self_attn.mem.o_proj.weight) | |
| def forward( | |
| self, | |
| input_ids: torch.LongTensor | None = None, | |
| attention_mask: torch.Tensor | None = None, | |
| position_ids: torch.LongTensor | None = None, | |
| past_key_values: Cache | None = None, | |
| inputs_embeds: torch.FloatTensor | None = None, | |
| labels: torch.LongTensor | None = None, | |
| use_cache: bool | None = None, | |
| is_evicted: torch.Tensor | None = None, | |
| write_mask: torch.Tensor | None = None, | |
| segment_ids: torch.Tensor | None = None, | |
| seq_ids: torch.Tensor | None = None, | |
| chunk_id: torch.Tensor | None = None, | |
| evict_step: torch.Tensor | None = None, | |
| cu_seqlens: torch.Tensor | None = None, | |
| mem_cache=None, | |
| logits_to_keep: int | torch.Tensor = 0, | |
| **kwargs, | |
| ) -> CausalLMOutputWithPast: | |
| outputs = self.model( | |
| input_ids=input_ids, | |
| attention_mask=attention_mask, | |
| position_ids=position_ids, | |
| past_key_values=past_key_values, | |
| inputs_embeds=inputs_embeds, | |
| use_cache=use_cache, | |
| is_evicted=is_evicted, | |
| write_mask=write_mask, | |
| segment_ids=segment_ids, | |
| seq_ids=seq_ids, | |
| chunk_id=chunk_id, | |
| evict_step=evict_step, | |
| cu_seqlens=cu_seqlens, | |
| mem_cache=mem_cache, | |
| **kwargs, | |
| ) | |
| hidden_states = outputs.last_hidden_state | |
| if labels is not None: | |
| # Answer-only logits: gather just the supervised positions and run | |
| # lm_head on those, so we never materialize [B, L, vocab] (≈40GB at | |
| # L=128k). Mathematically identical to full-seq CE with ignore_index | |
| # = -100 (mean over answer tokens -> 按 answer token 归一). | |
| shift_hidden = hidden_states[:, :-1, :] | |
| shift_labels = labels[:, 1:].to(hidden_states.device) | |
| sel = shift_labels != -100 | |
| sel_hidden = shift_hidden[sel] # [n_answer, H] bf16 | |
| sel_lab = shift_labels[sel] # [n_answer] | |
| n = sel_hidden.shape[0] | |
| # Chunked lm_head + CE over the answer tokens: never materialize the | |
| # full [n_answer, vocab] fp32 logits (≈30GB when a 64k pack is mostly | |
| # answer, e.g. long open-ended replies -> OOM). sum/n == mean CE. | |
| if n == 0: | |
| loss = hidden_states.sum() * 0.0 # keep graph; no supervised token | |
| else: | |
| CH = 8192 | |
| tot = hidden_states.new_zeros((), dtype=torch.float32) | |
| for s in range(0, n, CH): | |
| lg = self.lm_head(sel_hidden[s:s + CH]).float() | |
| tot = tot + nn.functional.cross_entropy( | |
| lg, sel_lab[s:s + CH], reduction="sum") | |
| loss = tot / n | |
| return CausalLMOutputWithPast(loss=loss, logits=None, | |
| past_key_values=outputs.past_key_values) | |
| slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep | |
| logits = self.lm_head(hidden_states[:, slice_indices, :]) | |
| return CausalLMOutputWithPast( | |
| loss=None, | |
| logits=logits, | |
| past_key_values=outputs.past_key_values, | |
| hidden_states=outputs.last_hidden_state, | |
| ) | |
| # ------------------------------------------------------------------ init | |
| def from_qwen3( | |
| cls, | |
| qwen3_path: str, | |
| memory_design: str = "Y", | |
| mem_layers: list[int] | None = None, | |
| mem_o_proj_zero_init: bool = True, | |
| dtype: torch.dtype | None = torch.bfloat16, | |
| device_map: str | None = None, | |
| attn_implementation: str | None = None, | |
| **config_overrides, | |
| ) -> "LiveMemForCausalLM": | |
| """Build a LiveMem from a pretrained Qwen3: load base weights, then | |
| copy-init each side branch from the backbone attention geometry.""" | |
| base = Qwen3ForCausalLM.from_pretrained( | |
| qwen3_path, dtype=dtype, attn_implementation=attn_implementation | |
| ) | |
| config = LiveMemConfig( | |
| memory_design=memory_design, | |
| mem_layers=mem_layers, | |
| mem_o_proj_zero_init=mem_o_proj_zero_init, | |
| **{**base.config.to_dict(), **config_overrides}, | |
| ) | |
| if attn_implementation is not None: | |
| config._attn_implementation = attn_implementation | |
| model = cls(config) | |
| if dtype is not None: | |
| model = model.to(dtype=dtype) | |
| # 1) load all base weights that map directly (embed/mlp/norms/lm_head and, | |
| # for wrapped layers, the main attention under `.self_attn.attn.*`). | |
| sd = _remap_base_state_dict(base.state_dict(), config.memory_layer_indices) | |
| missing, unexpected = model.load_state_dict(sd, strict=False) | |
| # side-branch params (model.layers.*.self_attn.mem.*) are expected-missing | |
| leftover = [k for k in missing if ".self_attn.mem." not in k] | |
| if leftover: | |
| raise RuntimeError(f"Unexpected missing keys after base load: {leftover[:8]} ...") | |
| if unexpected: | |
| raise RuntimeError(f"Unexpected keys when loading base: {unexpected[:8]} ...") | |
| # 2) copy-init each side branch from its (now-loaded) main attention. | |
| for idx in config.memory_layer_indices: | |
| mattn = model.model.layers[idx].self_attn | |
| copy_init_side_branch(mattn.mem, mattn.attn, config.mem_o_proj_zero_init) | |
| del base | |
| if device_map is not None: | |
| model = model.to(device_map) | |
| return model | |
| def _remap_base_state_dict(state_dict: dict, mem_layers: list[int]) -> dict: | |
| """Insert `.attn` into self_attn keys for wrapped layers so base attention | |
| weights land on LiveMemAttention.attn.*; all other keys pass through.""" | |
| mem_set = set(mem_layers) | |
| sub = ("q_proj", "k_proj", "v_proj", "o_proj", "q_norm", "k_norm") | |
| out = {} | |
| for k, v in state_dict.items(): | |
| nk = k | |
| if ".self_attn." in k: | |
| parts = k.split(".") | |
| try: | |
| li = parts.index("layers") | |
| layer_idx = int(parts[li + 1]) | |
| except (ValueError, IndexError): | |
| layer_idx = None | |
| if layer_idx in mem_set and any(f".self_attn.{s}." in k for s in sub): | |
| nk = k.replace(".self_attn.", ".self_attn.attn.", 1) | |
| out[nk] = v | |
| return out | |
| def copy_init_side_branch( | |
| side: LiveMemGatedDeltaNet2, attn: Qwen3Attention, zero_o: bool | |
| ) -> None: | |
| """Copy Qwen3 QKVO into the GDN2 side branch. | |
| Supports both the legacy full-MHA side branch (32 Q/K/V heads for Qwen3-4B) | |
| and the compact KV-head branch (8 Q/K/V heads + expanded V): | |
| - Q: direct copy if head counts match; if target heads match backbone KV | |
| heads, average the corresponding GQA Q group. | |
| - K: copy/adapt from backbone KV heads. | |
| - V: copy/adapt from backbone KV heads, then block-repeat each V head along | |
| its channel dimension when `expand_v > 1`. | |
| - O: copied only when shapes match; normally zero-initialized for training. | |
| """ | |
| hd = attn.head_dim | |
| dt = side.q_proj.weight.dtype | |
| qh = attn.q_proj.weight.shape[0] // hd | |
| kvh = attn.k_proj.weight.shape[0] // hd | |
| def adapt_heads(heads: torch.Tensor, target_heads: int, name: str) -> torch.Tensor: | |
| src_heads = heads.shape[0] | |
| if target_heads == src_heads: | |
| return heads | |
| if target_heads > src_heads and target_heads % src_heads == 0: | |
| return heads.repeat_interleave(target_heads // src_heads, dim=0) | |
| if src_heads > target_heads and src_heads % target_heads == 0: | |
| return heads.view(target_heads, src_heads // target_heads, hd, -1).mean(dim=1) | |
| raise ValueError(f"cannot adapt {name} heads from {src_heads} to {target_heads}") | |
| q_heads = attn.q_proj.weight.view(qh, hd, -1) | |
| if side.num_heads == qh: | |
| q_init = q_heads | |
| elif qh % kvh == 0 and side.num_heads == kvh: | |
| q_init = q_heads.view(kvh, qh // kvh, hd, -1).mean(dim=1) | |
| else: | |
| q_init = adapt_heads(q_heads, side.num_heads, "q") | |
| side.q_proj.weight.copy_(q_init.reshape(side.q_proj.weight.shape).to(dt)) | |
| k_heads = attn.k_proj.weight.view(kvh, hd, -1) | |
| k_init = adapt_heads(k_heads, side.num_heads, "k") | |
| side.k_proj.weight.copy_(k_init.reshape(side.k_proj.weight.shape).to(dt)) | |
| v_heads = adapt_heads(attn.v_proj.weight.view(kvh, hd, -1), side.num_v_heads, "v") | |
| if side.head_v_dim % hd != 0: | |
| raise ValueError( | |
| f"side.head_v_dim={side.head_v_dim} must be a multiple of backbone head_dim={hd} " | |
| "for copy initialization" | |
| ) | |
| v_expand = side.head_v_dim // hd | |
| v_init = v_heads.repeat(1, v_expand, 1) | |
| side.v_proj.weight.copy_(v_init.reshape(side.v_proj.weight.shape).to(dt)) | |
| if zero_o: | |
| side.o_proj.weight.zero_() | |
| else: | |
| if side.o_proj.weight.shape != attn.o_proj.weight.shape: | |
| raise ValueError( | |
| f"cannot copy-init o_proj with shape {tuple(side.o_proj.weight.shape)} " | |
| f"from backbone shape {tuple(attn.o_proj.weight.shape)}; use zero_o=True" | |
| ) | |
| side.o_proj.weight.copy_(attn.o_proj.weight.to(side.o_proj.weight.dtype)) | |